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:
parent
9326f4b747
commit
654e7ee567
13 changed files with 2765 additions and 0 deletions
|
|
@ -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
42
llmprovider/presets.go
Normal 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
61
llmprovider/presets.yml
Normal 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
312
llmprovider/registry.go
Normal 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
|
||||||
|
}
|
||||||
645
llmprovider/registry_test.go
Normal file
645
llmprovider/registry_test.go
Normal 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
274
llmprovider/store.go
Normal 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
199
llmprovider/sync.go
Normal 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
85
llmprovider/types.go
Normal 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
255
mcpclient/registry.go
Normal 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
483
mcpclient/registry_test.go
Normal 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
179
mcpclient/store.go
Normal 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
136
mcpclient/sync.go
Normal 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
50
mcpclient/types.go
Normal 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"`
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue