From 654e7ee567f2556c856f62e3c153097cff8c819e Mon Sep 17 00:00:00 2001 From: Max Date: Tue, 28 Apr 2026 14:33:10 +0800 Subject: [PATCH] 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. --- engine/load.go | 44 +++ llmprovider/presets.go | 42 +++ llmprovider/presets.yml | 61 ++++ llmprovider/registry.go | 312 +++++++++++++++++ llmprovider/registry_test.go | 645 +++++++++++++++++++++++++++++++++++ llmprovider/store.go | 274 +++++++++++++++ llmprovider/sync.go | 199 +++++++++++ llmprovider/types.go | 85 +++++ mcpclient/registry.go | 255 ++++++++++++++ mcpclient/registry_test.go | 483 ++++++++++++++++++++++++++ mcpclient/store.go | 179 ++++++++++ mcpclient/sync.go | 136 ++++++++ mcpclient/types.go | 50 +++ 13 files changed, 2765 insertions(+) create mode 100644 llmprovider/presets.go create mode 100644 llmprovider/presets.yml create mode 100644 llmprovider/registry.go create mode 100644 llmprovider/registry_test.go create mode 100644 llmprovider/store.go create mode 100644 llmprovider/sync.go create mode 100644 llmprovider/types.go create mode 100644 mcpclient/registry.go create mode 100644 mcpclient/registry_test.go create mode 100644 mcpclient/store.go create mode 100644 mcpclient/sync.go create mode 100644 mcpclient/types.go diff --git a/engine/load.go b/engine/load.go index 823da786..f80e5445 100644 --- a/engine/load.go +++ b/engine/load.go @@ -33,7 +33,9 @@ import ( "github.com/yaoapp/yao/i18n" "github.com/yaoapp/yao/job" "github.com/yaoapp/yao/kb" + "github.com/yaoapp/yao/llmprovider" "github.com/yaoapp/yao/mcp" + "github.com/yaoapp/yao/mcpclient" "github.com/yaoapp/yao/messenger" "github.com/yaoapp/yao/model" "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 { err = hook(cfg) if err != nil { @@ -655,6 +673,32 @@ func Reload(cfg config.Config, options LoadOption) (err error) { 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 _, err = openapi.Load(cfg) if err != nil { diff --git a/llmprovider/presets.go b/llmprovider/presets.go new file mode 100644 index 00000000..f05648ec --- /dev/null +++ b/llmprovider/presets.go @@ -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 +} diff --git a/llmprovider/presets.yml b/llmprovider/presets.yml new file mode 100644 index 00000000..65552448 --- /dev/null +++ b/llmprovider/presets.yml @@ -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 diff --git a/llmprovider/registry.go b/llmprovider/registry.go new file mode 100644 index 00000000..5e2e39af --- /dev/null +++ b/llmprovider/registry.go @@ -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 +} diff --git a/llmprovider/registry_test.go b/llmprovider/registry_test.go new file mode 100644 index 00000000..5b4cdf24 --- /dev/null +++ b/llmprovider/registry_test.go @@ -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) + } +} diff --git a/llmprovider/store.go b/llmprovider/store.go new file mode 100644 index 00000000..dccd9df2 --- /dev/null +++ b/llmprovider/store.go @@ -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) + } +} diff --git a/llmprovider/sync.go b/llmprovider/sync.go new file mode 100644 index 00000000..089920a5 --- /dev/null +++ b/llmprovider/sync.go @@ -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 + } +} diff --git a/llmprovider/types.go b/llmprovider/types.go new file mode 100644 index 00000000..1b187770 --- /dev/null +++ b/llmprovider/types.go @@ -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"` +} diff --git a/mcpclient/registry.go b/mcpclient/registry.go new file mode 100644 index 00000000..a7a24c10 --- /dev/null +++ b/mcpclient/registry.go @@ -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 +} diff --git a/mcpclient/registry_test.go b/mcpclient/registry_test.go new file mode 100644 index 00000000..96d67a4e --- /dev/null +++ b/mcpclient/registry_test.go @@ -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) + } +} diff --git a/mcpclient/store.go b/mcpclient/store.go new file mode 100644 index 00000000..93cebbbc --- /dev/null +++ b/mcpclient/store.go @@ -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) + } +} diff --git a/mcpclient/sync.go b/mcpclient/sync.go new file mode 100644 index 00000000..0056d951 --- /dev/null +++ b/mcpclient/sync.go @@ -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"}, + } +} diff --git a/mcpclient/types.go b/mcpclient/types.go new file mode 100644 index 00000000..61e6a3cf --- /dev/null +++ b/mcpclient/types.go @@ -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"` +}