feat(load): initialize and reload LLM Provider and MCP Client Registries

- Added initialization for LLM Provider and MCP Client Registries during the Load process.
- Implemented reload functionality for both registries to ensure they are properly refreshed when needed.
- Enhanced error handling to capture and report issues during initialization and reloading of the registries.
This commit is contained in:
Max 2026-04-28 14:33:10 +08:00
parent 9326f4b747
commit 654e7ee567
13 changed files with 2765 additions and 0 deletions

View file

@ -33,7 +33,9 @@ import (
"github.com/yaoapp/yao/i18n" "github.com/yaoapp/yao/i18n"
"github.com/yaoapp/yao/job" "github.com/yaoapp/yao/job"
"github.com/yaoapp/yao/kb" "github.com/yaoapp/yao/kb"
"github.com/yaoapp/yao/llmprovider"
"github.com/yaoapp/yao/mcp" "github.com/yaoapp/yao/mcp"
"github.com/yaoapp/yao/mcpclient"
"github.com/yaoapp/yao/messenger" "github.com/yaoapp/yao/messenger"
"github.com/yaoapp/yao/model" "github.com/yaoapp/yao/model"
"github.com/yaoapp/yao/monitor" "github.com/yaoapp/yao/monitor"
@ -421,6 +423,22 @@ func Load(cfg config.Config, options LoadOption, progressCallback ...func(string
} }
}() }()
// Initialize LLM Provider Registry
err = loadStep("LLM Provider", func() error {
return llmprovider.Init()
}, callback)
if err != nil {
warnings = append(warnings, Warning{Widget: "LLM Provider", Error: err})
}
// Initialize MCP Client Registry
err = loadStep("MCP Client Registry", func() error {
return mcpclient.Init()
}, callback)
if err != nil {
warnings = append(warnings, Warning{Widget: "MCP Client Registry", Error: err})
}
for name, hook := range LoadHooks { for name, hook := range LoadHooks {
err = hook(cfg) err = hook(cfg)
if err != nil { if err != nil {
@ -655,6 +673,32 @@ func Reload(cfg config.Config, options LoadOption) (err error) {
printErr(cfg.Mode, "Agent", err) printErr(cfg.Mode, "Agent", err)
} }
// Reload LLM Provider Registry
if llmprovider.Global != nil {
err = llmprovider.Global.Reload()
if err != nil {
printErr(cfg.Mode, "LLM Provider", err)
}
} else {
err = llmprovider.Init()
if err != nil {
printErr(cfg.Mode, "LLM Provider", err)
}
}
// Reload MCP Client Registry
if mcpclient.Global != nil {
err = mcpclient.Global.Reload()
if err != nil {
printErr(cfg.Mode, "MCP Client Registry", err)
}
} else {
err = mcpclient.Init()
if err != nil {
printErr(cfg.Mode, "MCP Client Registry", err)
}
}
// Load OpenAPI // Load OpenAPI
_, err = openapi.Load(cfg) _, err = openapi.Load(cfg)
if err != nil { if err != nil {

42
llmprovider/presets.go Normal file
View file

@ -0,0 +1,42 @@
package llmprovider
import (
_ "embed"
"gopkg.in/yaml.v3"
)
//go:embed presets.yml
var presetsYAML []byte
var presets []ProviderPreset
func init() {
presets = loadPresets()
}
func loadPresets() []ProviderPreset {
var list []ProviderPreset
if err := yaml.Unmarshal(presetsYAML, &list); err != nil {
panic("llmprovider: failed to parse presets.yml: " + err.Error())
}
return list
}
// GetPresets returns a copy of the embedded preset list.
func GetPresets() []ProviderPreset {
out := make([]ProviderPreset, len(presets))
copy(out, presets)
return out
}
// GetPreset returns the preset for the given key, or nil if not found.
func GetPreset(key string) *ProviderPreset {
for i := range presets {
if presets[i].Key == key {
cp := presets[i]
return &cp
}
}
return nil
}

61
llmprovider/presets.yml Normal file
View file

@ -0,0 +1,61 @@
- key: openai
name: OpenAI
type: openai
api_url: https://api.openai.com
require_key: true
default_models:
- id: gpt-4o
name: GPT-4o
capabilities: [vision, tool_calls, streaming, json]
enabled: true
- id: gpt-4o-mini
name: GPT-4o Mini
capabilities: [tool_calls, streaming, json]
enabled: true
- id: o3-mini
name: o3-mini
capabilities: [tool_calls, streaming, reasoning]
enabled: false
- key: anthropic
name: Anthropic
type: anthropic
api_url: https://api.anthropic.com
require_key: true
default_models:
- id: claude-sonnet-4-20250514
name: Claude Sonnet 4
capabilities: [vision, tool_calls, streaming, reasoning]
enabled: true
- id: claude-haiku-3-5-20241022
name: Claude Haiku 3.5
capabilities: [tool_calls, streaming]
enabled: true
- key: ollama
name: Ollama
type: openai
api_url: http://localhost:11434
require_key: false
url_editable: true
default_models: []
- key: azure
name: Azure OpenAI
type: openai
api_url: ""
require_key: true
url_editable: true
default_models: []
- key: yaoagents
name: Yao Agents
type: openai
api_url: https://api.yaoagents.com
require_key: false
is_cloud: true
default_models:
- id: default
name: Default
capabilities: [vision, tool_calls, streaming]
enabled: true

312
llmprovider/registry.go Normal file
View file

@ -0,0 +1,312 @@
package llmprovider
import (
"fmt"
"strings"
"sync"
"github.com/yaoapp/gou/connector"
"github.com/yaoapp/gou/store"
)
// Global is the singleton LLM Provider Registry.
var Global *Registry
// Registry manages LLM providers with CRUD, persistence, cache and runtime sync.
type Registry struct {
store store.Store
cache store.Store
encKey string
mu sync.RWMutex
}
// Init initializes the global Registry.
// Must be called after store.Load (so __yao.store and __yao.cache are available).
func Init() error {
s, err := store.Get("__yao.store")
if err != nil {
return fmt.Errorf("llmprovider.Init: %w", err)
}
c, _ := store.Get("__yao.cache")
r := &Registry{store: s, cache: c}
Global = r
if err := importFromConnectors(r); err != nil {
return fmt.Errorf("llmprovider.Init importFromConnectors: %w", err)
}
return nil
}
// SetEncryptionKey sets the key used for API key encryption at rest.
// Should be called right after Init if encryption is desired.
func (r *Registry) SetEncryptionKey(key string) {
r.mu.Lock()
defer r.mu.Unlock()
r.encKey = key
}
// Get retrieves a provider by key. Lazily ensures its connector is registered.
func (r *Registry) Get(key string) (*Provider, error) {
r.mu.RLock()
defer r.mu.RUnlock()
p, err := storeGet(r.store, r.cache, key, r.encKey)
if err != nil {
return nil, err
}
_ = ensureConnector(p)
return p, nil
}
// GetMasked retrieves a provider with the API key masked for display.
func (r *Registry) GetMasked(key string) (*Provider, error) {
p, err := r.Get(key)
if err != nil {
return nil, err
}
cp := *p
cp.APIKey = maskAPIKey(cp.APIKey)
return &cp, nil
}
// Create adds a new provider. Persists, caches, registers connector, and updates index.
func (r *Registry) Create(p *Provider) (*Provider, error) {
r.mu.Lock()
defer r.mu.Unlock()
if p.Key == "" {
return nil, fmt.Errorf("provider key is required")
}
if r.store.Has(storeKey(p.Key)) {
return nil, fmt.Errorf("provider %s already exists", p.Key)
}
if p.Source == "" {
p.Source = ProviderSourceDynamic
}
if p.ConnectorID == "" {
p.ConnectorID = connectorID(p)
}
if p.Status == "" {
p.Status = "unconfigured"
}
if err := storeSet(r.store, r.cache, p, r.encKey); err != nil {
return nil, err
}
if err := indexAdd(r.store, r.cache, p.Key); err != nil {
return nil, err
}
if p.Enabled {
_ = ensureConnector(p)
}
return p, nil
}
// Update modifies an existing provider. Hot-replaces the connector if needed.
func (r *Registry) Update(key string, p *Provider) (*Provider, error) {
r.mu.Lock()
defer r.mu.Unlock()
old, err := storeGet(r.store, r.cache, key, r.encKey)
if err != nil {
return nil, err
}
p.Key = key
if p.Source == "" {
p.Source = old.Source
}
if p.ConnectorID == "" {
p.ConnectorID = old.ConnectorID
}
if p.Owner == (ProviderOwner{}) {
p.Owner = old.Owner
}
_ = unregisterConnector(old)
if err := storeSet(r.store, r.cache, p, r.encKey); err != nil {
return nil, err
}
if p.Enabled {
_ = ensureConnector(p)
}
return p, nil
}
// Delete removes a provider by key. Unregisters connector, deletes store/cache/index.
func (r *Registry) Delete(key string) error {
r.mu.Lock()
defer r.mu.Unlock()
p, err := storeGet(r.store, r.cache, key, r.encKey)
if err != nil {
return err
}
_ = unregisterConnector(p)
if err := storeDel(r.store, r.cache, key); err != nil {
return err
}
return indexRemove(r.store, r.cache, key)
}
// List returns providers matching the filter.
func (r *Registry) List(filter *ProviderFilter) ([]Provider, error) {
r.mu.RLock()
defer r.mu.RUnlock()
keys, err := indexGet(r.store, r.cache)
if err != nil {
return nil, err
}
var result []Provider
for _, key := range keys {
p, err := storeGet(r.store, r.cache, key, r.encKey)
if err != nil {
continue
}
if filter != nil && !matchFilter(p, filter) {
continue
}
cp := *p
cp.APIKey = maskAPIKey(cp.APIKey)
result = append(result, cp)
}
return result, nil
}
// Reload re-reads all providers from persistent store and rebuilds cache + connectors.
func (r *Registry) Reload() error {
r.mu.Lock()
defer r.mu.Unlock()
keys, err := indexGet(r.store, nil)
if err != nil {
return err
}
for _, key := range keys {
p, err := storeGet(r.store, nil, key, r.encKey)
if err != nil {
continue
}
m, err := providerToMap(p, r.encKey)
if err != nil {
continue
}
if r.cache != nil {
r.cache.Set(storeKey(key), m, 0)
}
if p.Source == ProviderSourceDynamic && p.Enabled {
_ = ensureConnector(p)
}
}
return nil
}
// GetConnector returns the runtime connector for a given provider key.
func (r *Registry) GetConnector(key string) (connector.Connector, error) {
p, err := r.Get(key)
if err != nil {
return nil, err
}
cid := p.ConnectorID
if cid == "" {
cid = connectorID(p)
}
return connector.Select(cid)
}
// GetSetting returns the runtime connector setting map for a given provider key.
func (r *Registry) GetSetting(key string) (map[string]interface{}, error) {
conn, err := r.GetConnector(key)
if err != nil {
return nil, err
}
return conn.Setting(), nil
}
// matchFilter checks if a provider matches the given filter.
func matchFilter(p *Provider, f *ProviderFilter) bool {
src := f.Source
if src == "" {
src = ProviderSourceDynamic
}
if src != ProviderSourceAll && p.Source != src {
return false
}
if f.Owner != nil {
if f.Owner.Type != "" && p.Owner.Type != f.Owner.Type {
return false
}
if f.Owner.UserID != "" && p.Owner.UserID != f.Owner.UserID {
return false
}
if f.Owner.TeamID != "" && p.Owner.TeamID != f.Owner.TeamID {
return false
}
}
if f.Enabled != nil && p.Enabled != *f.Enabled {
return false
}
if f.Type != nil && p.Type != *f.Type {
return false
}
if f.PresetKey != nil && p.PresetKey != *f.PresetKey {
return false
}
if len(f.Capabilities) > 0 && !matchCapabilities(p, f.Capabilities) {
return false
}
if f.Keyword != "" {
kw := strings.ToLower(f.Keyword)
if !strings.Contains(strings.ToLower(p.Name), kw) &&
!strings.Contains(strings.ToLower(p.Key), kw) {
return false
}
}
return true
}
// matchCapabilities returns true if at least one model in the provider
// satisfies ALL of the required capabilities (AND logic).
func matchCapabilities(p *Provider, required []string) bool {
for _, m := range p.Models {
if !m.Enabled {
continue
}
capSet := make(map[string]bool, len(m.Capabilities))
for _, c := range m.Capabilities {
capSet[c] = true
}
allMatch := true
for _, req := range required {
if !capSet[req] {
allMatch = false
break
}
}
if allMatch {
return true
}
}
return false
}

View file

@ -0,0 +1,645 @@
package llmprovider_test
import (
"fmt"
"os"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/yaoapp/gou/connector"
"github.com/yaoapp/gou/store"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/llmprovider"
"github.com/yaoapp/yao/test"
)
func TestMain(m *testing.M) {
test.Prepare(nil, config.Conf)
defer test.Clean()
os.Exit(m.Run())
}
func setupRegistry(t *testing.T) *llmprovider.Registry {
t.Helper()
test.Prepare(t, config.Conf)
err := llmprovider.Init()
require.NoError(t, err)
t.Cleanup(func() {
s, _ := store.Get("__yao.store")
if s != nil {
s.Del("llmprovider:*")
}
c, _ := store.Get("__yao.cache")
if c != nil {
c.Del("llmprovider:*")
}
test.Clean()
})
return llmprovider.Global
}
var testProvider = llmprovider.Provider{
Key: "test-openai",
Name: "Test OpenAI",
Type: "openai",
APIURL: "https://api.openai.com",
APIKey: "sk-test-xxxxx",
Models: []llmprovider.ModelInfo{{ID: "gpt-4o", Name: "GPT-4o", Capabilities: []string{"vision", "tool_calls", "streaming"}, Enabled: true}},
Enabled: true,
RequireKey: true,
Owner: llmprovider.ProviderOwner{Type: "system"},
}
func TestCreate(t *testing.T) {
r := setupRegistry(t)
p := testProvider
created, err := r.Create(&p)
require.NoError(t, err)
assert.Equal(t, "test-openai", created.Key)
assert.Equal(t, llmprovider.ProviderSourceDynamic, created.Source)
assert.NotEmpty(t, created.ConnectorID)
// Verify store persistence
s, _ := store.Get("__yao.store")
assert.True(t, s.Has("llmprovider:p:test-openai"))
// Verify connector registered
_, err = connector.Select(created.ConnectorID)
assert.NoError(t, err)
}
func TestCreateDuplicate(t *testing.T) {
r := setupRegistry(t)
p := testProvider
_, err := r.Create(&p)
require.NoError(t, err)
dup := testProvider
_, err = r.Create(&dup)
assert.Error(t, err)
assert.Contains(t, err.Error(), "already exists")
}
func TestGet(t *testing.T) {
r := setupRegistry(t)
p := testProvider
_, err := r.Create(&p)
require.NoError(t, err)
got, err := r.Get("test-openai")
require.NoError(t, err)
assert.Equal(t, "Test OpenAI", got.Name)
assert.Equal(t, "openai", got.Type)
assert.Equal(t, "https://api.openai.com", got.APIURL)
assert.Len(t, got.Models, 1)
assert.Equal(t, "gpt-4o", got.Models[0].ID)
}
func TestGetNotFound(t *testing.T) {
r := setupRegistry(t)
_, err := r.Get("nonexistent")
assert.Error(t, err)
}
func TestGetMasked(t *testing.T) {
r := setupRegistry(t)
p := testProvider
_, err := r.Create(&p)
require.NoError(t, err)
got, err := r.GetMasked("test-openai")
require.NoError(t, err)
assert.NotEqual(t, "sk-test-xxxxx", got.APIKey)
assert.True(t, len(got.APIKey) > 0)
// Last 4 chars should be visible
assert.Contains(t, got.APIKey, "xxxx")
}
func TestGetLazy(t *testing.T) {
r := setupRegistry(t)
p := testProvider
created, err := r.Create(&p)
require.NoError(t, err)
// Manually unregister the connector
err = connector.Unregister(created.ConnectorID)
require.NoError(t, err)
// Verify it's gone
_, err = connector.Select(created.ConnectorID)
assert.Error(t, err)
// Get should lazily re-register
got, err := r.Get("test-openai")
require.NoError(t, err)
assert.Equal(t, "test-openai", got.Key)
// Connector should be back
_, err = connector.Select(got.ConnectorID)
assert.NoError(t, err)
}
func TestList(t *testing.T) {
r := setupRegistry(t)
providers := []llmprovider.Provider{
{Key: "p1", Name: "Provider 1", Type: "openai", Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "gpt-4o", Name: "GPT-4o", Capabilities: []string{"vision", "tool_calls"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "system"}},
{Key: "p2", Name: "Provider 2", Type: "anthropic", Enabled: false,
Models: []llmprovider.ModelInfo{{ID: "claude-3", Name: "Claude 3", Capabilities: []string{"tool_calls"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "user", UserID: "123"}},
{Key: "p3", Name: "Provider 3", Type: "openai", Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "gpt-4o-mini", Name: "GPT-4o Mini", Capabilities: []string{"streaming"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "system"}},
}
for i := range providers {
_, err := r.Create(&providers[i])
require.NoError(t, err)
}
t.Run("AllDynamic", func(t *testing.T) {
list, err := r.List(&llmprovider.ProviderFilter{Source: llmprovider.ProviderSourceDynamic})
require.NoError(t, err)
assert.GreaterOrEqual(t, len(list), 3)
})
t.Run("FilterByType", func(t *testing.T) {
typ := "openai"
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceDynamic,
Type: &typ,
})
require.NoError(t, err)
for _, p := range list {
assert.Equal(t, "openai", p.Type)
}
})
t.Run("FilterByEnabled", func(t *testing.T) {
enabled := true
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceDynamic,
Enabled: &enabled,
})
require.NoError(t, err)
for _, p := range list {
assert.True(t, p.Enabled)
}
})
t.Run("FilterByOwner", func(t *testing.T) {
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceDynamic,
Owner: &llmprovider.ProviderOwner{Type: "user", UserID: "123"},
})
require.NoError(t, err)
for _, p := range list {
assert.Equal(t, "user", p.Owner.Type)
assert.Equal(t, "123", p.Owner.UserID)
}
})
t.Run("FilterByCapabilities", func(t *testing.T) {
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceDynamic,
Capabilities: []string{"vision", "tool_calls"},
})
require.NoError(t, err)
for _, p := range list {
found := false
for _, m := range p.Models {
capSet := map[string]bool{}
for _, c := range m.Capabilities {
capSet[c] = true
}
if capSet["vision"] && capSet["tool_calls"] {
found = true
break
}
}
assert.True(t, found, "provider %s should have model matching vision+tool_calls", p.Key)
}
})
t.Run("FilterByKeyword", func(t *testing.T) {
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceDynamic,
Keyword: "Provider 2",
})
require.NoError(t, err)
found := false
for _, p := range list {
if p.Key == "p2" {
found = true
}
}
assert.True(t, found)
})
}
func TestUpdate(t *testing.T) {
r := setupRegistry(t)
p := testProvider
created, err := r.Create(&p)
require.NoError(t, err)
updated := *created
updated.APIURL = "https://custom.openai.com"
updated.APIKey = "sk-new-key"
result, err := r.Update("test-openai", &updated)
require.NoError(t, err)
assert.Equal(t, "https://custom.openai.com", result.APIURL)
// Verify store updated
got, err := r.Get("test-openai")
require.NoError(t, err)
assert.Equal(t, "https://custom.openai.com", got.APIURL)
}
func TestDelete(t *testing.T) {
r := setupRegistry(t)
p := testProvider
created, err := r.Create(&p)
require.NoError(t, err)
cid := created.ConnectorID
err = r.Delete("test-openai")
require.NoError(t, err)
// Verify removed from store
_, err = r.Get("test-openai")
assert.Error(t, err)
// Verify connector unregistered
_, err = connector.Select(cid)
assert.Error(t, err)
}
func TestReload(t *testing.T) {
r := setupRegistry(t)
p := testProvider
_, err := r.Create(&p)
require.NoError(t, err)
// Clear cache to simulate stale state
c, _ := store.Get("__yao.cache")
if c != nil {
c.Del("llmprovider:*")
}
err = r.Reload()
require.NoError(t, err)
// Should still be able to get the provider
got, err := r.Get("test-openai")
require.NoError(t, err)
assert.Equal(t, "Test OpenAI", got.Name)
}
func TestImportFromConnectors(t *testing.T) {
r := setupRegistry(t)
// After Init, builtin connectors should be imported
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceAll,
})
require.NoError(t, err)
builtinCount := 0
for _, p := range list {
if p.Source == llmprovider.ProviderSourceBuiltIn {
builtinCount++
}
}
// Should have imported some from connector.AIConnectors (if test app has connectors)
t.Logf("Imported %d builtin providers from connector.AIConnectors (total AIConnectors: %d)", builtinCount, len(connector.AIConnectors))
}
func TestGetPresets(t *testing.T) {
presets := llmprovider.GetPresets()
assert.Greater(t, len(presets), 0, "should have at least one preset")
// Verify openai preset exists
var openai *llmprovider.ProviderPreset
for i := range presets {
if presets[i].Key == "openai" {
openai = &presets[i]
break
}
}
require.NotNil(t, openai, "openai preset should exist")
assert.Equal(t, "OpenAI", openai.Name)
assert.Equal(t, "openai", openai.Type)
assert.True(t, openai.RequireKey)
assert.Greater(t, len(openai.DefaultModels), 0)
}
func TestGetPreset(t *testing.T) {
p := llmprovider.GetPreset("anthropic")
require.NotNil(t, p)
assert.Equal(t, "Anthropic", p.Name)
none := llmprovider.GetPreset("nonexistent")
assert.Nil(t, none)
}
func TestEncryptionRoundTrip(t *testing.T) {
r := setupRegistry(t)
r.SetEncryptionKey("my-super-secret-key-for-tests")
p := testProvider
p.Key = "test-encrypted"
_, err := r.Create(&p)
require.NoError(t, err)
got, err := r.Get("test-encrypted")
require.NoError(t, err)
assert.Equal(t, "sk-test-xxxxx", got.APIKey, "APIKey should be decrypted on read")
masked, err := r.GetMasked("test-encrypted")
require.NoError(t, err)
assert.NotEqual(t, "sk-test-xxxxx", masked.APIKey)
assert.Contains(t, masked.APIKey, "xxxx")
// Verify raw store value is encrypted
s, _ := store.Get("__yao.store")
raw, ok := s.Get("llmprovider:p:test-encrypted")
require.True(t, ok)
m := raw.(map[string]interface{})
storedKey, _ := m["api_key"].(string)
assert.True(t, len(storedKey) > 0)
assert.NotEqual(t, "sk-test-xxxxx", storedKey, "raw stored value should be encrypted")
}
func TestGetConnector(t *testing.T) {
r := setupRegistry(t)
p := testProvider
p.Key = "test-getconn"
_, err := r.Create(&p)
require.NoError(t, err)
conn, err := r.GetConnector("test-getconn")
require.NoError(t, err)
assert.NotNil(t, conn)
setting := conn.Setting()
assert.NotNil(t, setting)
host, _ := setting["host"].(string)
assert.Equal(t, "https://api.openai.com", host)
}
func TestGetSetting(t *testing.T) {
r := setupRegistry(t)
p := testProvider
p.Key = "test-getsetting"
_, err := r.Create(&p)
require.NoError(t, err)
setting, err := r.GetSetting("test-getsetting")
require.NoError(t, err)
assert.NotNil(t, setting)
host, _ := setting["host"].(string)
assert.Equal(t, "https://api.openai.com", host)
}
func TestGetConnectorNotFound(t *testing.T) {
r := setupRegistry(t)
_, err := r.GetConnector("not-exist")
assert.Error(t, err)
}
func TestGetSettingNotFound(t *testing.T) {
r := setupRegistry(t)
_, err := r.GetSetting("not-exist")
assert.Error(t, err)
}
func TestCreateEmptyKey(t *testing.T) {
r := setupRegistry(t)
p := llmprovider.Provider{Name: "No Key"}
_, err := r.Create(&p)
assert.Error(t, err)
assert.Contains(t, err.Error(), "key is required")
}
func TestCreateDisabled(t *testing.T) {
r := setupRegistry(t)
p := llmprovider.Provider{
Key: "test-disabled",
Name: "Disabled Provider",
Type: "openai",
APIURL: "https://api.openai.com",
Enabled: false,
Models: []llmprovider.ModelInfo{{ID: "gpt-4o", Name: "GPT-4o", Capabilities: []string{"streaming"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "system"},
}
created, err := r.Create(&p)
require.NoError(t, err)
assert.Equal(t, "unconfigured", created.Status)
// Disabled provider should not have its connector registered
_, err = connector.Select(created.ConnectorID)
assert.Error(t, err, "disabled provider should not register connector")
}
func TestOwnerPrefixedIDs(t *testing.T) {
r := setupRegistry(t)
cases := []struct {
key string
owner llmprovider.ProviderOwner
prefix string
}{
{"owner-sys", llmprovider.ProviderOwner{Type: "system"}, "s."},
{"owner-user", llmprovider.ProviderOwner{Type: "user", UserID: "42"}, "u42."},
{"owner-team", llmprovider.ProviderOwner{Type: "team", TeamID: "99"}, "t99."},
}
for _, tc := range cases {
t.Run(tc.key, func(t *testing.T) {
p := llmprovider.Provider{
Key: tc.key,
Name: tc.key,
Type: "openai",
APIURL: "https://api.openai.com",
Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "m1", Name: "M1", Capabilities: []string{"streaming"}, Enabled: true}},
Owner: tc.owner,
}
created, err := r.Create(&p)
require.NoError(t, err)
assert.Contains(t, created.ConnectorID, tc.prefix,
"ConnectorID for %s owner should contain prefix %s", tc.owner.Type, tc.prefix)
// Verify connector is registered with the prefixed ID
_, err = connector.Select(created.ConnectorID)
assert.NoError(t, err)
})
}
}
func TestListBuiltInFilter(t *testing.T) {
r := setupRegistry(t)
builtinList, err := r.List(&llmprovider.ProviderFilter{Source: llmprovider.ProviderSourceBuiltIn})
require.NoError(t, err)
for _, p := range builtinList {
assert.Equal(t, llmprovider.ProviderSourceBuiltIn, p.Source)
}
}
func TestListPresetKeyFilter(t *testing.T) {
r := setupRegistry(t)
p := llmprovider.Provider{
Key: "from-preset",
Name: "From Preset",
Type: "openai",
PresetKey: "openai",
Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "gpt-4o", Name: "GPT-4o", Capabilities: []string{"streaming"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "system"},
}
_, err := r.Create(&p)
require.NoError(t, err)
pk := "openai"
list, err := r.List(&llmprovider.ProviderFilter{
Source: llmprovider.ProviderSourceDynamic,
PresetKey: &pk,
})
require.NoError(t, err)
found := false
for _, item := range list {
if item.Key == "from-preset" {
found = true
assert.Equal(t, "openai", item.PresetKey)
}
}
assert.True(t, found)
}
func TestDefaultModelFallback(t *testing.T) {
r := setupRegistry(t)
// Provider with no enabled models — should use first model ID as default
p := llmprovider.Provider{
Key: "test-fallback",
Name: "Fallback",
Type: "openai",
APIURL: "https://api.openai.com",
Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "only-model", Name: "Only", Capabilities: []string{"streaming"}, Enabled: false}},
Owner: llmprovider.ProviderOwner{Type: "system"},
}
created, err := r.Create(&p)
require.NoError(t, err)
// Connector should still be registered using the fallback model
conn, cerr := connector.Select(created.ConnectorID)
require.NoError(t, cerr)
setting := conn.Setting()
model, _ := setting["model"].(string)
assert.Equal(t, "only-model", model)
}
func TestMaskShortKey(t *testing.T) {
r := setupRegistry(t)
p := llmprovider.Provider{
Key: "test-shortkey",
Name: "Short",
Type: "openai",
APIKey: "ab",
Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "m", Name: "M", Capabilities: []string{"streaming"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "system"},
}
_, err := r.Create(&p)
require.NoError(t, err)
masked, err := r.GetMasked("test-shortkey")
require.NoError(t, err)
// Short keys should be fully masked
assert.Equal(t, "**", masked.APIKey)
}
func TestConcurrency(t *testing.T) {
r := setupRegistry(t)
var wg sync.WaitGroup
errCh := make(chan error, 30)
// Concurrent creates with unique keys
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
p := llmprovider.Provider{
Key: fmt.Sprintf("conc-%d", idx),
Name: fmt.Sprintf("Concurrent %d", idx),
Type: "openai",
APIURL: "https://api.openai.com",
Enabled: true,
Models: []llmprovider.ModelInfo{{ID: "gpt-4o", Name: "GPT-4o", Capabilities: []string{"streaming"}, Enabled: true}},
Owner: llmprovider.ProviderOwner{Type: "system"},
}
if _, err := r.Create(&p); err != nil {
errCh <- err
}
}(i)
}
wg.Wait()
// Concurrent reads
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
_, err := r.Get(fmt.Sprintf("conc-%d", idx))
if err != nil {
errCh <- err
}
}(i)
}
wg.Wait()
// Concurrent deletes
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
if err := r.Delete(fmt.Sprintf("conc-%d", idx)); err != nil {
errCh <- err
}
}(i)
}
wg.Wait()
close(errCh)
for err := range errCh {
t.Errorf("concurrent operation error: %v", err)
}
}

274
llmprovider/store.go Normal file
View file

@ -0,0 +1,274 @@
package llmprovider
import (
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/json"
"fmt"
"io"
"strings"
"github.com/yaoapp/gou/store"
)
const (
keyPrefix = "llmprovider:p:"
indexKey = "llmprovider:index"
maskChars = 4
encPrefix = "enc:"
)
func storeKey(key string) string { return keyPrefix + key }
// providerToMap converts Provider to map[string]interface{} for store.Set.
// Encrypts APIKey before writing.
func providerToMap(p *Provider, encKey string) (map[string]interface{}, error) {
cp := *p
if cp.APIKey != "" && encKey != "" {
encrypted, err := encryptString(cp.APIKey, encKey)
if err != nil {
return nil, fmt.Errorf("encrypt api_key: %w", err)
}
cp.APIKey = encPrefix + encrypted
}
raw, err := json.Marshal(cp)
if err != nil {
return nil, err
}
var m map[string]interface{}
if err := json.Unmarshal(raw, &m); err != nil {
return nil, err
}
return m, nil
}
// mapToProvider converts map[string]interface{} from store.Get back to Provider.
// Decrypts APIKey after reading.
func mapToProvider(m map[string]interface{}, encKey string) (*Provider, error) {
raw, err := json.Marshal(m)
if err != nil {
return nil, err
}
var p Provider
if err := json.Unmarshal(raw, &p); err != nil {
return nil, err
}
if strings.HasPrefix(p.APIKey, encPrefix) && encKey != "" {
decrypted, err := decryptString(strings.TrimPrefix(p.APIKey, encPrefix), encKey)
if err != nil {
return nil, fmt.Errorf("decrypt api_key: %w", err)
}
p.APIKey = decrypted
}
return &p, nil
}
// maskAPIKey returns a masked version of the API key for display.
func maskAPIKey(key string) string {
if len(key) <= maskChars {
return strings.Repeat("*", len(key))
}
return strings.Repeat("*", len(key)-maskChars) + key[len(key)-maskChars:]
}
// storeGet reads a provider from cache first, then persistent store.
func storeGet(s, c store.Store, key, encKey string) (*Provider, error) {
sk := storeKey(key)
if c != nil {
if val, ok := c.Get(sk); ok {
if m, ok := val.(map[string]interface{}); ok {
return mapToProvider(m, encKey)
}
}
}
val, ok := s.Get(sk)
if !ok {
return nil, fmt.Errorf("provider %s not found", key)
}
m, ok := val.(map[string]interface{})
if !ok {
return nil, fmt.Errorf("provider %s: unexpected store type %T", key, val)
}
p, err := mapToProvider(m, encKey)
if err != nil {
return nil, err
}
if c != nil {
c.Set(sk, m, 0)
}
return p, nil
}
// storeSet writes a provider to both persistent store and cache.
func storeSet(s, c store.Store, p *Provider, encKey string) error {
m, err := providerToMap(p, encKey)
if err != nil {
return err
}
sk := storeKey(p.Key)
if err := s.Set(sk, m, 0); err != nil {
return err
}
if c != nil {
c.Set(sk, m, 0)
}
return nil
}
// storeDel removes a provider from both persistent store and cache.
func storeDel(s, c store.Store, key string) error {
sk := storeKey(key)
if err := s.Del(sk); err != nil {
return err
}
if c != nil {
c.Del(sk)
}
return nil
}
// indexGet returns all provider keys from the index.
func indexGet(s, c store.Store) ([]string, error) {
var raw interface{}
var ok bool
if c != nil {
raw, ok = c.Get(indexKey)
}
if !ok {
raw, ok = s.Get(indexKey)
if !ok {
return nil, nil
}
if c != nil {
c.Set(indexKey, raw, 0)
}
}
switch v := raw.(type) {
case []interface{}:
keys := make([]string, 0, len(v))
for _, item := range v {
if str, ok := item.(string); ok {
keys = append(keys, str)
}
}
return keys, nil
case []string:
return v, nil
default:
return nil, fmt.Errorf("unexpected index type %T", raw)
}
}
// indexSet writes the full index to both stores.
func indexSet(s, c store.Store, keys []string) error {
iface := make([]interface{}, len(keys))
for i, k := range keys {
iface[i] = k
}
if err := s.Set(indexKey, iface, 0); err != nil {
return err
}
if c != nil {
c.Set(indexKey, iface, 0)
}
return nil
}
// indexAdd appends a key to the index if not present.
func indexAdd(s, c store.Store, key string) error {
keys, err := indexGet(s, c)
if err != nil {
return err
}
for _, k := range keys {
if k == key {
return nil
}
}
return indexSet(s, c, append(keys, key))
}
// indexRemove removes a key from the index.
func indexRemove(s, c store.Store, key string) error {
keys, err := indexGet(s, c)
if err != nil {
return err
}
filtered := make([]string, 0, len(keys))
for _, k := range keys {
if k != key {
filtered = append(filtered, k)
}
}
return indexSet(s, c, filtered)
}
// --- AES-256-GCM encryption helpers ---
func deriveKey(secret string) []byte {
h := sha256.Sum256([]byte(secret))
return h[:]
}
func encryptString(plaintext, secret string) (string, error) {
key := deriveKey(secret)
block, err := aes.NewCipher(key)
if err != nil {
return "", err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", err
}
nonce := make([]byte, gcm.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return "", err
}
ciphertext := gcm.Seal(nonce, nonce, []byte(plaintext), nil)
return base64.StdEncoding.EncodeToString(ciphertext), nil
}
func decryptString(encoded, secret string) (string, error) {
key := deriveKey(secret)
data, err := base64.StdEncoding.DecodeString(encoded)
if err != nil {
return "", err
}
block, err := aes.NewCipher(key)
if err != nil {
return "", err
}
gcm, err := cipher.NewGCM(block)
if err != nil {
return "", err
}
nonceSize := gcm.NonceSize()
if len(data) < nonceSize {
return "", fmt.Errorf("ciphertext too short")
}
plaintext, err := gcm.Open(nil, data[:nonceSize], data[nonceSize:], nil)
if err != nil {
return "", err
}
return string(plaintext), nil
}
// storeCleanAll removes all llmprovider keys (for testing cleanup).
func storeCleanAll(s, c store.Store) {
_ = s.Del(keyPrefix + "*")
_ = s.Del(indexKey)
if c != nil {
_ = c.Del(keyPrefix + "*")
_ = c.Del(indexKey)
}
}

199
llmprovider/sync.go Normal file
View file

@ -0,0 +1,199 @@
package llmprovider
import (
"encoding/json"
"fmt"
"github.com/yaoapp/gou/connector"
)
// connectorID builds the runtime ID for registering into connector.Connectors.
// Dynamic providers get an owner prefix to avoid collision with builtin IDs.
func connectorID(p *Provider) string {
switch p.Owner.Type {
case "user":
return "u" + p.Owner.UserID + "." + p.Key
case "team":
return "t" + p.Owner.TeamID + "." + p.Key
default:
return "s." + p.Key
}
}
// defaultModel returns the first enabled model ID, or empty string.
func defaultModel(p *Provider) string {
for _, m := range p.Models {
if m.Enabled {
return m.ID
}
}
if len(p.Models) > 0 {
return p.Models[0].ID
}
return ""
}
// marshalDSL builds a connector DSL JSON from the flat Provider fields.
func marshalDSL(p *Provider) ([]byte, error) {
dsl := map[string]interface{}{
"type": p.Type,
"name": p.Name,
"label": p.Name,
"options": map[string]interface{}{
"host": p.APIURL,
"key": p.APIKey,
"model": defaultModel(p),
},
}
return json.Marshal(dsl)
}
// ensureConnector makes sure the provider's connector is registered in the runtime.
// Builtin providers are managed by engine.Load and skipped here.
func ensureConnector(p *Provider) error {
if p.Source == ProviderSourceBuiltIn {
return nil
}
if !p.Enabled {
return nil
}
cid := p.ConnectorID
if cid == "" {
cid = connectorID(p)
}
if _, err := connector.Select(cid); err == nil {
return nil
}
dslJSON, err := marshalDSL(p)
if err != nil {
return fmt.Errorf("ensureConnector %s: marshal DSL: %w", p.Key, err)
}
_, err = connector.LoadSourceSync(dslJSON, cid, "__registry/"+cid+".conn.yao")
if err != nil {
return fmt.Errorf("ensureConnector %s: LoadSourceSync: %w", p.Key, err)
}
return nil
}
// unregisterConnector removes the provider's connector from the runtime.
func unregisterConnector(p *Provider) error {
if p.Source == ProviderSourceBuiltIn {
return nil
}
cid := p.ConnectorID
if cid == "" {
cid = connectorID(p)
}
return connector.Unregister(cid)
}
// importFromConnectors scans existing AI connectors loaded by engine.Load
// and imports them as builtin providers into the Registry store.
// If a store record with the same key already exists (dynamic), it is not overwritten.
func importFromConnectors(r *Registry) error {
for _, opt := range connector.AIConnectors {
id := opt.Value
if r.store.Has(storeKey(id)) {
continue
}
conn, err := connector.Select(id)
if err != nil {
continue
}
p := providerFromConnector(id, conn)
m, err := providerToMap(&p, r.encKey)
if err != nil {
continue
}
sk := storeKey(id)
_ = r.store.Set(sk, m, 0)
if r.cache != nil {
_ = r.cache.Set(sk, m, 0)
}
_ = indexAdd(r.store, r.cache, id)
}
return nil
}
// providerFromConnector builds a Provider from a runtime Connector interface.
func providerFromConnector(id string, conn connector.Connector) Provider {
meta := conn.GetMetaInfo()
setting := conn.Setting()
name := meta.Label
if name == "" {
name = id
}
typ := connectorType(conn)
apiURL, _ := setting["host"].(string)
apiKey, _ := setting["key"].(string)
model, _ := setting["model"].(string)
var models []ModelInfo
if model != "" {
caps := capabilitiesFromSetting(setting)
models = []ModelInfo{{
ID: model,
Name: model,
Capabilities: caps,
Enabled: true,
}}
}
return Provider{
Key: id,
ConnectorID: id,
Name: name,
Type: typ,
APIURL: apiURL,
APIKey: apiKey,
Models: models,
Enabled: true,
Status: "connected",
Source: ProviderSourceBuiltIn,
Owner: ProviderOwner{Type: "system"},
}
}
func connectorType(conn connector.Connector) string {
switch {
case conn.Is(6): // OPENAI
return "openai"
case conn.Is(11): // ANTHROPIC
return "anthropic"
case conn.Is(9): // FASTEMBED
return "fastembed"
case conn.Is(8): // MOAPI
return "moapi"
default:
return "custom"
}
}
func capabilitiesFromSetting(setting map[string]interface{}) []string {
raw, ok := setting["capabilities"]
if !ok {
return nil
}
switch caps := raw.(type) {
case map[string]interface{}:
var out []string
for k, v := range caps {
if b, ok := v.(bool); ok && b {
out = append(out, k)
}
}
return out
default:
return nil
}
}

85
llmprovider/types.go Normal file
View file

@ -0,0 +1,85 @@
package llmprovider
// Provider represents a configured LLM provider (one vendor connection with multiple models).
// Fields align with the frontend ProviderConfig interface.
type Provider struct {
Key string `json:"key"`
ConnectorID string `json:"connector_id"`
Name string `json:"name"`
Type string `json:"type"`
APIURL string `json:"api_url"`
APIKey string `json:"api_key"`
Models []ModelInfo `json:"models"`
Enabled bool `json:"enabled"`
Status string `json:"status"`
IsCustom bool `json:"is_custom,omitempty"`
PresetKey string `json:"preset_key,omitempty"`
RequireKey bool `json:"require_key"`
Source ProviderSource `json:"source"`
Owner ProviderOwner `json:"owner"`
}
// ModelInfo describes a single model within a provider.
// Fields align with the frontend ModelInfo interface.
type ModelInfo struct {
ID string `json:"id" yaml:"id"`
Name string `json:"name" yaml:"name"`
Capabilities []string `json:"capabilities" yaml:"capabilities"`
Enabled bool `json:"enabled" yaml:"enabled"`
}
// ProviderOwner identifies who owns a provider.
type ProviderOwner struct {
Type string `json:"type"`
TeamID string `json:"team_id,omitempty"`
UserID string `json:"user_id,omitempty"`
}
// ProviderSource distinguishes dynamic (registry-created) from builtin (DSL-loaded) providers.
type ProviderSource string
const (
ProviderSourceDynamic ProviderSource = "dynamic"
ProviderSourceBuiltIn ProviderSource = "builtin"
ProviderSourceAll ProviderSource = "all"
)
// ProviderFilter specifies criteria for listing providers.
type ProviderFilter struct {
Owner *ProviderOwner
Enabled *bool
Source ProviderSource // defaults to "dynamic" when zero-value
Type *string
PresetKey *string
Capabilities []string // AND filter: provider matches if any model satisfies all
Keyword string
}
// ProviderPreset is a static UI-only template for creating providers.
// Fields align with the frontend ProviderPreset interface.
type ProviderPreset struct {
Key string `json:"key" yaml:"key"`
Name string `json:"name" yaml:"name"`
Type string `json:"type" yaml:"type"`
APIURL string `json:"api_url" yaml:"api_url"`
RequireKey bool `json:"require_key" yaml:"require_key"`
IsCloud bool `json:"is_cloud,omitempty" yaml:"is_cloud,omitempty"`
URLEditable bool `json:"url_editable,omitempty" yaml:"url_editable,omitempty"`
DefaultModels []ModelInfo `json:"default_models" yaml:"default_models"`
}
// ProviderTestResult holds the outcome of a provider connectivity test.
type ProviderTestResult struct {
Success bool `json:"success"`
Message string `json:"message"`
LatencyMs int64 `json:"latency_ms,omitempty"`
}
// RoleAssignment maps model roles to specific provider+model pairs.
type RoleAssignment map[string]RoleTarget
// RoleTarget identifies a provider and model for a given role.
type RoleTarget struct {
Provider string `json:"provider"`
Model string `json:"model"`
}

255
mcpclient/registry.go Normal file
View file

@ -0,0 +1,255 @@
package mcpclient
import (
"fmt"
"strings"
"sync"
"github.com/yaoapp/gou/mcp"
"github.com/yaoapp/gou/store"
)
// Global is the singleton MCP Client Registry.
var Global *Registry
// Registry manages MCP clients with CRUD, persistence, cache and runtime sync.
type Registry struct {
store store.Store
cache store.Store
mu sync.RWMutex
}
// Init initializes the global Registry.
// Must be called after store.Load and mcp.Load.
func Init() error {
s, err := store.Get("__yao.store")
if err != nil {
return fmt.Errorf("mcpclient.Init: %w", err)
}
c, _ := store.Get("__yao.cache")
r := &Registry{store: s, cache: c}
Global = r
if err := importFromClients(r); err != nil {
return fmt.Errorf("mcpclient.Init importFromClients: %w", err)
}
return nil
}
// Get retrieves a client by ID. Lazily ensures its runtime client is registered.
func (r *Registry) Get(id string) (*Client, error) {
r.mu.RLock()
defer r.mu.RUnlock()
c, err := storeGet(r.store, r.cache, id)
if err != nil {
return nil, err
}
_ = ensureClient(c)
return c, nil
}
// Create adds a new client. Persists, caches, registers runtime, and updates index.
func (r *Registry) Create(c *Client) (*Client, error) {
r.mu.Lock()
defer r.mu.Unlock()
if c.ID == "" {
return nil, fmt.Errorf("client id is required")
}
if r.store.Has(storeKey(c.ID)) {
return nil, fmt.Errorf("client %s already exists", c.ID)
}
if c.Source == "" {
c.Source = ClientSourceDynamic
}
if c.RuntimeID == "" {
c.RuntimeID = runtimeID(c)
}
if c.Status == "" {
c.Status = "unconfigured"
}
if err := storeSet(r.store, r.cache, c); err != nil {
return nil, err
}
if err := indexAdd(r.store, r.cache, c.ID); err != nil {
return nil, err
}
if c.Enabled {
_ = ensureClient(c)
}
return c, nil
}
// Update modifies an existing client. Hot-replaces the runtime client.
func (r *Registry) Update(id string, c *Client) (*Client, error) {
r.mu.Lock()
defer r.mu.Unlock()
old, err := storeGet(r.store, r.cache, id)
if err != nil {
return nil, err
}
c.ID = id
if c.Source == "" {
c.Source = old.Source
}
if c.RuntimeID == "" {
c.RuntimeID = old.RuntimeID
}
if c.Owner == (ClientOwner{}) {
c.Owner = old.Owner
}
unloadClient(old)
if err := storeSet(r.store, r.cache, c); err != nil {
return nil, err
}
if c.Enabled {
_ = ensureClient(c)
}
return c, nil
}
// Delete removes a client by ID. Unloads runtime, deletes store/cache/index.
func (r *Registry) Delete(id string) error {
r.mu.Lock()
defer r.mu.Unlock()
c, err := storeGet(r.store, r.cache, id)
if err != nil {
return err
}
unloadClient(c)
if err := storeDel(r.store, r.cache, id); err != nil {
return err
}
return indexRemove(r.store, r.cache, id)
}
// List returns clients matching the filter.
func (r *Registry) List(filter *ClientFilter) ([]Client, error) {
r.mu.RLock()
defer r.mu.RUnlock()
ids, err := indexGet(r.store, r.cache)
if err != nil {
return nil, err
}
var result []Client
for _, id := range ids {
c, err := storeGet(r.store, r.cache, id)
if err != nil {
continue
}
if filter != nil && !matchFilter(c, filter) {
continue
}
result = append(result, *c)
}
return result, nil
}
// Reload re-reads all clients from persistent store and rebuilds cache + runtime.
func (r *Registry) Reload() error {
r.mu.Lock()
defer r.mu.Unlock()
ids, err := indexGet(r.store, nil)
if err != nil {
return err
}
for _, id := range ids {
c, err := storeGet(r.store, nil, id)
if err != nil {
continue
}
m, err := clientToMap(c)
if err != nil {
continue
}
if r.cache != nil {
r.cache.Set(storeKey(id), m, 0)
}
if c.Source == ClientSourceDynamic && c.Enabled {
_ = ensureClient(c)
}
}
return nil
}
// GetMCPClient returns the runtime mcp.Client for a given registry ID.
func (r *Registry) GetMCPClient(id string) (mcp.Client, error) {
c, err := r.Get(id)
if err != nil {
return nil, err
}
rid := c.RuntimeID
if rid == "" {
rid = runtimeID(c)
}
defer func() { recover() }()
client := mcp.GetClient(rid)
if client == nil {
return nil, fmt.Errorf("runtime mcp client %s not found", rid)
}
return client, nil
}
func matchFilter(c *Client, f *ClientFilter) bool {
src := f.Source
if src == "" {
src = ClientSourceDynamic
}
if src != ClientSourceAll && c.Source != src {
return false
}
if f.Owner != nil {
if f.Owner.Type != "" && c.Owner.Type != f.Owner.Type {
return false
}
if f.Owner.ID != "" && c.Owner.ID != f.Owner.ID {
return false
}
}
if f.Enabled != nil && c.Enabled != *f.Enabled {
return false
}
if f.Transport != nil && c.ClientDSL.Transport != *f.Transport {
return false
}
if f.Type != nil && c.ClientDSL.Type != *f.Type {
return false
}
if f.Keyword != "" {
kw := strings.ToLower(f.Keyword)
if !strings.Contains(strings.ToLower(c.ClientDSL.Name), kw) &&
!strings.Contains(strings.ToLower(c.ID), kw) &&
!strings.Contains(strings.ToLower(c.ClientDSL.Label), kw) {
return false
}
}
return true
}

483
mcpclient/registry_test.go Normal file
View file

@ -0,0 +1,483 @@
package mcpclient_test
import (
"fmt"
"os"
"sync"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/yaoapp/gou/mcp"
mcpTypes "github.com/yaoapp/gou/mcp/types"
"github.com/yaoapp/gou/store"
"github.com/yaoapp/yao/config"
"github.com/yaoapp/yao/mcpclient"
"github.com/yaoapp/yao/test"
)
func TestMain(m *testing.M) {
test.Prepare(nil, config.Conf)
defer test.Clean()
os.Exit(m.Run())
}
func setupRegistry(t *testing.T) *mcpclient.Registry {
t.Helper()
test.Prepare(t, config.Conf)
err := mcpclient.Init()
require.NoError(t, err)
t.Cleanup(func() {
s, _ := store.Get("__yao.store")
if s != nil {
s.Del("mcpclient:*")
}
c, _ := store.Get("__yao.cache")
if c != nil {
c.Del("mcpclient:*")
}
test.Clean()
})
return mcpclient.Global
}
func newTestClient(id string) mcpclient.Client {
return mcpclient.Client{
ClientDSL: mcpTypes.ClientDSL{
ID: id,
Name: "Test " + id,
Type: "standard",
Transport: mcpTypes.TransportStdio,
Command: "echo",
Arguments: []string{"hello"},
},
Enabled: true,
Owner: mcpclient.ClientOwner{Type: "system"},
}
}
func TestCreate(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-stdio")
created, err := r.Create(&c)
require.NoError(t, err)
assert.Equal(t, "test-stdio", created.ID)
assert.Equal(t, mcpclient.ClientSourceDynamic, created.Source)
assert.NotEmpty(t, created.RuntimeID)
s, _ := store.Get("__yao.store")
assert.True(t, s.Has("mcpclient:c:test-stdio"))
}
func TestCreateDuplicate(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-dup")
_, err := r.Create(&c)
require.NoError(t, err)
dup := newTestClient("test-dup")
_, err = r.Create(&dup)
assert.Error(t, err)
assert.Contains(t, err.Error(), "already exists")
}
func TestGet(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-get")
_, err := r.Create(&c)
require.NoError(t, err)
got, err := r.Get("test-get")
require.NoError(t, err)
assert.Equal(t, "Test test-get", got.Name)
assert.Equal(t, mcpTypes.TransportStdio, got.Transport)
assert.Equal(t, "echo", got.Command)
}
func TestGetNotFound(t *testing.T) {
r := setupRegistry(t)
_, err := r.Get("nonexistent")
assert.Error(t, err)
}
func TestGetLazy(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-lazy")
created, err := r.Create(&c)
require.NoError(t, err)
// Manually unload the client
mcp.UnloadClient(created.RuntimeID)
assert.False(t, mcp.Exists(created.RuntimeID))
// Get should lazily re-register
got, err := r.Get("test-lazy")
require.NoError(t, err)
assert.Equal(t, "test-lazy", got.ID)
}
func TestList(t *testing.T) {
r := setupRegistry(t)
clients := []mcpclient.Client{
{
ClientDSL: mcpTypes.ClientDSL{ID: "c1", Name: "Client 1", Type: "standard", Transport: mcpTypes.TransportStdio, Command: "echo"},
Enabled: true,
Owner: mcpclient.ClientOwner{Type: "system"},
},
{
ClientDSL: mcpTypes.ClientDSL{ID: "c2", Name: "Client 2", Type: "agent", Transport: mcpTypes.TransportSSE, URL: "http://localhost:3001"},
Enabled: false,
Owner: mcpclient.ClientOwner{Type: "user", ID: "123"},
},
{
ClientDSL: mcpTypes.ClientDSL{ID: "c3", Name: "Client 3", Type: "standard", Transport: mcpTypes.TransportStdio, Command: "cat"},
Enabled: true,
Owner: mcpclient.ClientOwner{Type: "system"},
},
}
for i := range clients {
_, err := r.Create(&clients[i])
require.NoError(t, err)
}
t.Run("AllDynamic", func(t *testing.T) {
list, err := r.List(&mcpclient.ClientFilter{Source: mcpclient.ClientSourceDynamic})
require.NoError(t, err)
assert.GreaterOrEqual(t, len(list), 3)
})
t.Run("FilterByTransport", func(t *testing.T) {
tp := mcpTypes.TransportSSE
list, err := r.List(&mcpclient.ClientFilter{
Source: mcpclient.ClientSourceDynamic,
Transport: &tp,
})
require.NoError(t, err)
for _, c := range list {
assert.Equal(t, mcpTypes.TransportSSE, c.Transport)
}
})
t.Run("FilterByEnabled", func(t *testing.T) {
enabled := true
list, err := r.List(&mcpclient.ClientFilter{
Source: mcpclient.ClientSourceDynamic,
Enabled: &enabled,
})
require.NoError(t, err)
for _, c := range list {
assert.True(t, c.Enabled)
}
})
t.Run("FilterByOwner", func(t *testing.T) {
list, err := r.List(&mcpclient.ClientFilter{
Source: mcpclient.ClientSourceDynamic,
Owner: &mcpclient.ClientOwner{Type: "user", ID: "123"},
})
require.NoError(t, err)
for _, c := range list {
assert.Equal(t, "user", c.Owner.Type)
}
})
t.Run("FilterByType", func(t *testing.T) {
typ := "agent"
list, err := r.List(&mcpclient.ClientFilter{
Source: mcpclient.ClientSourceDynamic,
Type: &typ,
})
require.NoError(t, err)
for _, c := range list {
assert.Equal(t, "agent", c.ClientDSL.Type)
}
})
t.Run("FilterByKeyword", func(t *testing.T) {
list, err := r.List(&mcpclient.ClientFilter{
Source: mcpclient.ClientSourceDynamic,
Keyword: "Client 2",
})
require.NoError(t, err)
found := false
for _, c := range list {
if c.ID == "c2" {
found = true
}
}
assert.True(t, found)
})
}
func TestUpdate(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-update")
_, err := r.Create(&c)
require.NoError(t, err)
got, err := r.Get("test-update")
require.NoError(t, err)
updated := *got
updated.ClientDSL.Name = "Updated Name"
updated.ClientDSL.Command = "cat"
result, err := r.Update("test-update", &updated)
require.NoError(t, err)
assert.Equal(t, "Updated Name", result.Name)
got2, err := r.Get("test-update")
require.NoError(t, err)
assert.Equal(t, "cat", got2.Command)
}
func TestDelete(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-delete")
_, err := r.Create(&c)
require.NoError(t, err)
err = r.Delete("test-delete")
require.NoError(t, err)
_, err = r.Get("test-delete")
assert.Error(t, err)
}
func TestReload(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-reload")
_, err := r.Create(&c)
require.NoError(t, err)
// Clear cache
cache, _ := store.Get("__yao.cache")
if cache != nil {
cache.Del("mcpclient:*")
}
err = r.Reload()
require.NoError(t, err)
got, err := r.Get("test-reload")
require.NoError(t, err)
assert.Equal(t, "Test test-reload", got.Name)
}
func TestImportFromClients(t *testing.T) {
r := setupRegistry(t)
list, err := r.List(&mcpclient.ClientFilter{Source: mcpclient.ClientSourceAll})
require.NoError(t, err)
builtinCount := 0
for _, c := range list {
if c.Source == mcpclient.ClientSourceBuiltIn {
builtinCount++
}
}
loadedClients := mcp.ListClients()
t.Logf("Imported %d builtin clients from mcp.ListClients (total loaded: %d)", builtinCount, len(loadedClients))
}
func TestToolListField(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-toollist")
c.ClientDSL.Tools = map[string]string{"my-tool": "scripts.MyTool"}
c.ToolList = []mcpTypes.Tool{
{Name: "discovered-tool", Description: "A tool discovered at runtime"},
}
created, err := r.Create(&c)
require.NoError(t, err)
got, err := r.Get(created.ID)
require.NoError(t, err)
assert.Len(t, got.ToolList, 1)
assert.Equal(t, "discovered-tool", got.ToolList[0].Name)
assert.Equal(t, "scripts.MyTool", got.ClientDSL.Tools["my-tool"])
}
func TestGetMCPClient(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-getmcp")
_, err := r.Create(&c)
require.NoError(t, err)
// The MCP client may or may not actually start (depends on whether "echo" is a valid MCP server),
// but we should at least exercise the code path.
_, err = r.GetMCPClient("test-getmcp")
// Either it works or returns a "not found" — both are valid for this test fixture
t.Logf("GetMCPClient result: err=%v", err)
}
func TestGetMCPClientNotFound(t *testing.T) {
r := setupRegistry(t)
_, err := r.GetMCPClient("no-such-client")
assert.Error(t, err)
}
func TestCreateEmptyID(t *testing.T) {
r := setupRegistry(t)
c := mcpclient.Client{ClientDSL: mcpTypes.ClientDSL{Name: "No ID"}}
_, err := r.Create(&c)
assert.Error(t, err)
assert.Contains(t, err.Error(), "id is required")
}
func TestCreateDisabled(t *testing.T) {
r := setupRegistry(t)
c := mcpclient.Client{
ClientDSL: mcpTypes.ClientDSL{ID: "test-disabled", Name: "Disabled", Type: "standard", Transport: mcpTypes.TransportStdio, Command: "echo"},
Enabled: false,
Owner: mcpclient.ClientOwner{Type: "system"},
}
created, err := r.Create(&c)
require.NoError(t, err)
assert.Equal(t, "unconfigured", created.Status)
// Disabled client should not be registered at runtime
assert.False(t, mcp.Exists(created.RuntimeID), "disabled client should not be registered")
}
func TestOwnerPrefixedRuntimeIDs(t *testing.T) {
r := setupRegistry(t)
cases := []struct {
id string
owner mcpclient.ClientOwner
prefix string
}{
{"owner-sys", mcpclient.ClientOwner{Type: "system"}, "s."},
{"owner-usr", mcpclient.ClientOwner{Type: "user", ID: "42"}, "u42."},
{"owner-team", mcpclient.ClientOwner{Type: "team", ID: "99"}, "t99."},
{"owner-asst", mcpclient.ClientOwner{Type: "assistant", ID: "a1"}, "aa1."},
}
for _, tc := range cases {
t.Run(tc.id, func(t *testing.T) {
c := mcpclient.Client{
ClientDSL: mcpTypes.ClientDSL{ID: tc.id, Name: tc.id, Type: "standard", Transport: mcpTypes.TransportStdio, Command: "echo"},
Enabled: true,
Owner: tc.owner,
}
created, err := r.Create(&c)
require.NoError(t, err)
assert.Contains(t, created.RuntimeID, tc.prefix,
"RuntimeID for %s owner should contain prefix %s", tc.owner.Type, tc.prefix)
})
}
}
func TestListBuiltInFilter(t *testing.T) {
r := setupRegistry(t)
builtinList, err := r.List(&mcpclient.ClientFilter{Source: mcpclient.ClientSourceBuiltIn})
require.NoError(t, err)
for _, c := range builtinList {
assert.Equal(t, mcpclient.ClientSourceBuiltIn, c.Source)
}
}
func TestListAllSources(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("test-all-src")
_, err := r.Create(&c)
require.NoError(t, err)
all, err := r.List(&mcpclient.ClientFilter{Source: mcpclient.ClientSourceAll})
require.NoError(t, err)
hasDynamic := false
for _, item := range all {
if item.Source == mcpclient.ClientSourceDynamic {
hasDynamic = true
}
}
assert.True(t, hasDynamic)
}
func TestUpdateNotFound(t *testing.T) {
r := setupRegistry(t)
c := newTestClient("not-exist")
_, err := r.Update("not-exist", &c)
assert.Error(t, err)
}
func TestDeleteNotFound(t *testing.T) {
r := setupRegistry(t)
err := r.Delete("not-exist")
assert.Error(t, err)
}
func TestConcurrency(t *testing.T) {
r := setupRegistry(t)
var wg sync.WaitGroup
errCh := make(chan error, 30)
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
c := mcpclient.Client{
ClientDSL: mcpTypes.ClientDSL{
ID: fmt.Sprintf("conc-%d", idx),
Name: fmt.Sprintf("Concurrent %d", idx),
Type: "standard",
Transport: mcpTypes.TransportStdio,
Command: "echo",
},
Enabled: true,
Owner: mcpclient.ClientOwner{Type: "system"},
}
if _, err := r.Create(&c); err != nil {
errCh <- err
}
}(i)
}
wg.Wait()
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
_, err := r.Get(fmt.Sprintf("conc-%d", idx))
if err != nil {
errCh <- err
}
}(i)
}
wg.Wait()
for i := 0; i < 10; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
if err := r.Delete(fmt.Sprintf("conc-%d", idx)); err != nil {
errCh <- err
}
}(i)
}
wg.Wait()
close(errCh)
for err := range errCh {
t.Errorf("concurrent operation error: %v", err)
}
}

179
mcpclient/store.go Normal file
View file

@ -0,0 +1,179 @@
package mcpclient
import (
"encoding/json"
"fmt"
"github.com/yaoapp/gou/store"
)
const (
keyPrefix = "mcpclient:c:"
indexKey = "mcpclient:index"
)
func storeKey(id string) string { return keyPrefix + id }
func clientToMap(c *Client) (map[string]interface{}, error) {
raw, err := json.Marshal(c)
if err != nil {
return nil, err
}
var m map[string]interface{}
if err := json.Unmarshal(raw, &m); err != nil {
return nil, err
}
return m, nil
}
func mapToClient(m map[string]interface{}) (*Client, error) {
raw, err := json.Marshal(m)
if err != nil {
return nil, err
}
var c Client
if err := json.Unmarshal(raw, &c); err != nil {
return nil, err
}
return &c, nil
}
func storeGet(s, c store.Store, id string) (*Client, error) {
sk := storeKey(id)
if c != nil {
if val, ok := c.Get(sk); ok {
if m, ok := val.(map[string]interface{}); ok {
return mapToClient(m)
}
}
}
val, ok := s.Get(sk)
if !ok {
return nil, fmt.Errorf("client %s not found", id)
}
m, ok := val.(map[string]interface{})
if !ok {
return nil, fmt.Errorf("client %s: unexpected store type %T", id, val)
}
cl, err := mapToClient(m)
if err != nil {
return nil, err
}
if c != nil {
c.Set(sk, m, 0)
}
return cl, nil
}
func storeSet(s, c store.Store, cl *Client) error {
m, err := clientToMap(cl)
if err != nil {
return err
}
sk := storeKey(cl.ID)
if err := s.Set(sk, m, 0); err != nil {
return err
}
if c != nil {
c.Set(sk, m, 0)
}
return nil
}
func storeDel(s, c store.Store, id string) error {
sk := storeKey(id)
if err := s.Del(sk); err != nil {
return err
}
if c != nil {
c.Del(sk)
}
return nil
}
func indexGet(s, c store.Store) ([]string, error) {
var raw interface{}
var ok bool
if c != nil {
raw, ok = c.Get(indexKey)
}
if !ok {
raw, ok = s.Get(indexKey)
if !ok {
return nil, nil
}
if c != nil {
c.Set(indexKey, raw, 0)
}
}
switch v := raw.(type) {
case []interface{}:
keys := make([]string, 0, len(v))
for _, item := range v {
if str, ok := item.(string); ok {
keys = append(keys, str)
}
}
return keys, nil
case []string:
return v, nil
default:
return nil, fmt.Errorf("unexpected index type %T", raw)
}
}
func indexSet(s, c store.Store, ids []string) error {
iface := make([]interface{}, len(ids))
for i, k := range ids {
iface[i] = k
}
if err := s.Set(indexKey, iface, 0); err != nil {
return err
}
if c != nil {
c.Set(indexKey, iface, 0)
}
return nil
}
func indexAdd(s, c store.Store, id string) error {
ids, err := indexGet(s, c)
if err != nil {
return err
}
for _, k := range ids {
if k == id {
return nil
}
}
return indexSet(s, c, append(ids, id))
}
func indexRemove(s, c store.Store, id string) error {
ids, err := indexGet(s, c)
if err != nil {
return err
}
filtered := make([]string, 0, len(ids))
for _, k := range ids {
if k != id {
filtered = append(filtered, k)
}
}
return indexSet(s, c, filtered)
}
func storeCleanAll(s, c store.Store) {
_ = s.Del(keyPrefix + "*")
_ = s.Del(indexKey)
if c != nil {
_ = c.Del(keyPrefix + "*")
_ = c.Del(indexKey)
}
}

136
mcpclient/sync.go Normal file
View file

@ -0,0 +1,136 @@
package mcpclient
import (
"encoding/json"
"fmt"
"github.com/yaoapp/gou/mcp"
mcpTypes "github.com/yaoapp/gou/mcp/types"
)
// runtimeID builds the runtime ID for registering into mcp.clients.
// Dynamic clients get an owner prefix to avoid collision with builtin IDs.
func runtimeID(c *Client) string {
switch c.Owner.Type {
case "user":
return "u" + c.Owner.ID + "." + c.ID
case "team":
return "t" + c.Owner.ID + "." + c.ID
case "assistant":
return "a" + c.Owner.ID + "." + c.ID
default:
return "s." + c.ID
}
}
// ensureClient makes sure the MCP client is registered in the runtime.
// Builtin clients are managed by engine.Load and skipped here.
func ensureClient(c *Client) error {
if c.Source == ClientSourceBuiltIn {
return nil
}
if !c.Enabled {
return nil
}
rid := c.RuntimeID
if rid == "" {
rid = runtimeID(c)
}
if mcp.Exists(rid) {
return nil
}
dslJSON, err := json.Marshal(c.ClientDSL)
if err != nil {
return fmt.Errorf("ensureClient %s: marshal DSL: %w", c.ID, err)
}
clientType := c.ClientDSL.Type
_, err = mcp.LoadClientSourceWithType(string(dslJSON), rid, clientType)
if err != nil {
return fmt.Errorf("ensureClient %s: LoadClientSourceWithType: %w", c.ID, err)
}
return nil
}
// unloadClient removes the client from the runtime.
func unloadClient(c *Client) {
if c.Source == ClientSourceBuiltIn {
return
}
rid := c.RuntimeID
if rid == "" {
rid = runtimeID(c)
}
mcp.UnloadClient(rid)
}
// importFromClients scans existing MCP clients loaded by engine.Load
// and imports them as builtin entries into the Registry store.
// If a store record with the same ID already exists (dynamic), it is not overwritten.
func importFromClients(r *Registry) error {
ids := mcp.ListClients()
for _, id := range ids {
if r.store.Has(storeKey(id)) {
continue
}
cl := clientFromRuntime(id)
if cl == nil {
continue
}
m, err := clientToMap(cl)
if err != nil {
continue
}
sk := storeKey(id)
_ = r.store.Set(sk, m, 0)
if r.cache != nil {
_ = r.cache.Set(sk, m, 0)
}
_ = indexAdd(r.store, r.cache, id)
}
return nil
}
// clientFromRuntime builds a Client from a runtime mcp.Client interface.
// Uses Info() and GetMetaInfo() since full ClientDSL is not exposed.
func clientFromRuntime(id string) *Client {
defer func() { recover() }()
mcpClient := mcp.GetClient(id)
if mcpClient == nil {
return nil
}
info := mcpClient.Info()
if info == nil {
return nil
}
meta := mcpClient.GetMetaInfo()
name := info.Name
if name == "" {
name = id
}
return &Client{
ClientDSL: mcpTypes.ClientDSL{
ID: id,
Name: name,
Type: info.Type,
Transport: info.Transport,
MetaInfo: meta,
},
RuntimeID: id,
Enabled: true,
Status: "connected",
Source: ClientSourceBuiltIn,
Owner: ClientOwner{Type: "system"},
}
}

50
mcpclient/types.go Normal file
View file

@ -0,0 +1,50 @@
package mcpclient
import (
mcpTypes "github.com/yaoapp/gou/mcp/types"
)
// Client wraps mcpTypes.ClientDSL with Registry management fields.
// Uses ClientDSL.ID as the registry key.
type Client struct {
mcpTypes.ClientDSL
RuntimeID string `json:"runtime_id"`
Enabled bool `json:"enabled"`
Status string `json:"status"`
Source ClientSource `json:"source"`
ToolList []mcpTypes.Tool `json:"tool_list,omitempty"`
Owner ClientOwner `json:"owner"`
}
// ClientOwner identifies who owns a client entry.
type ClientOwner struct {
Type string `json:"type"`
ID string `json:"id,omitempty"`
}
// ClientSource distinguishes registry-created from DSL-loaded clients.
type ClientSource string
const (
ClientSourceDynamic ClientSource = "dynamic"
ClientSourceBuiltIn ClientSource = "builtin"
ClientSourceAll ClientSource = "all"
)
// ClientFilter specifies criteria for listing clients.
type ClientFilter struct {
Owner *ClientOwner
Enabled *bool
Source ClientSource
Transport *mcpTypes.TransportType
Type *string
Keyword string
}
// ClientTestResult holds the outcome of a client connectivity test.
type ClientTestResult struct {
Success bool `json:"success"`
Message string `json:"message"`
LatencyMs int64 `json:"latency_ms,omitempty"`
}