diff --git a/pkg/agent/instance.go b/pkg/agent/instance.go index df6b8ba0d..f13518366 100644 --- a/pkg/agent/instance.go +++ b/pkg/agent/instance.go @@ -15,22 +15,23 @@ import ( // AgentInstance represents a fully configured agent with its own workspace, // session manager, context builder, and tool registry. type AgentInstance struct { - ID string - Name string - Model string - Fallbacks []string - Workspace string - MaxIterations int - MaxTokens int - Temperature float64 - ContextWindow int - Provider providers.LLMProvider - Sessions *session.SessionManager - ContextBuilder *ContextBuilder - Tools *tools.ToolRegistry - Subagents *config.SubagentsConfig - SkillsFilter []string - Candidates []providers.FallbackCandidate + ID string + Name string + Model string + Fallbacks []string + Workspace string + MaxIterations int + MaxTokens int + Temperature float64 + ContextWindow int + Provider providers.LLMProvider // Default provider for backward compatibility + ProviderRegistry *providers.ProviderRegistry // Registry for multi-provider fallback + Sessions *session.SessionManager + ContextBuilder *ContextBuilder + Tools *tools.ToolRegistry + Subagents *config.SubagentsConfig + SkillsFilter []string + Candidates []providers.FallbackCandidate } // NewAgentInstance creates an agent instance from config. @@ -39,6 +40,7 @@ func NewAgentInstance( defaults *config.AgentDefaults, cfg *config.Config, provider providers.LLMProvider, + providerRegistry *providers.ProviderRegistry, ) *AgentInstance { workspace := resolveAgentWorkspace(agentCfg, defaults) os.MkdirAll(workspace, 0o755) @@ -99,22 +101,23 @@ func NewAgentInstance( candidates := providers.ResolveCandidates(modelCfg, defaults.Provider) return &AgentInstance{ - ID: agentID, - Name: agentName, - Model: model, - Fallbacks: fallbacks, - Workspace: workspace, - MaxIterations: maxIter, - MaxTokens: maxTokens, - Temperature: temperature, - ContextWindow: maxTokens, - Provider: provider, - Sessions: sessionsManager, - ContextBuilder: contextBuilder, - Tools: toolsRegistry, - Subagents: subagents, - SkillsFilter: skillsFilter, - Candidates: candidates, + ID: agentID, + Name: agentName, + Model: model, + Fallbacks: fallbacks, + Workspace: workspace, + MaxIterations: maxIter, + MaxTokens: maxTokens, + Temperature: temperature, + ContextWindow: maxTokens, + Provider: provider, + ProviderRegistry: providerRegistry, + Sessions: sessionsManager, + ContextBuilder: contextBuilder, + Tools: toolsRegistry, + Subagents: subagents, + SkillsFilter: skillsFilter, + Candidates: candidates, } } diff --git a/pkg/agent/instance_resolution_test.go b/pkg/agent/instance_resolution_test.go index e61ffdc2c..e66bbea39 100644 --- a/pkg/agent/instance_resolution_test.go +++ b/pkg/agent/instance_resolution_test.go @@ -133,7 +133,7 @@ func TestNewAgentInstance_ModelResolution(t *testing.T) { // Use the existing mockProvider from mock_provider_test.go provider := &mockProvider{} - instance := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) + instance := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider, nil) require.NotNil(t, instance) diff --git a/pkg/agent/instance_test.go b/pkg/agent/instance_test.go index fcc8e9bea..fa8635206 100644 --- a/pkg/agent/instance_test.go +++ b/pkg/agent/instance_test.go @@ -29,7 +29,7 @@ func TestNewAgentInstance_UsesDefaultsTemperatureAndMaxTokens(t *testing.T) { cfg.Agents.Defaults.Temperature = &configuredTemp provider := &mockProvider{} - agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) + agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider, nil) if agent.MaxTokens != 1234 { t.Fatalf("MaxTokens = %d, want %d", agent.MaxTokens, 1234) @@ -61,7 +61,7 @@ func TestNewAgentInstance_DefaultsTemperatureWhenZero(t *testing.T) { cfg.Agents.Defaults.Temperature = &configuredTemp provider := &mockProvider{} - agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) + agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider, nil) if agent.Temperature != 0.0 { t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.0) @@ -87,7 +87,7 @@ func TestNewAgentInstance_DefaultsTemperatureWhenUnset(t *testing.T) { } provider := &mockProvider{} - agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider) + agent := NewAgentInstance(nil, &cfg.Agents.Defaults, cfg, provider, nil) if agent.Temperature != 0.7 { t.Fatalf("Temperature = %f, want %f", agent.Temperature, 0.7) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index ef31223d3..38c3de52e 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -519,11 +519,26 @@ func (al *AgentLoop) runLLMIteration( callLLM := func() (*providers.LLMResponse, error) { if len(agent.Candidates) > 1 && al.fallback != nil { fbResult, fbErr := al.fallback.Execute(ctx, agent.Candidates, - func(ctx context.Context, provider, model string) (*providers.LLMResponse, error) { - return agent.Provider.Chat(ctx, messages, providerToolDefs, model, map[string]any{ - "max_tokens": agent.MaxTokens, - "temperature": agent.Temperature, - "prompt_cache_key": agent.ID, + func(ctx context.Context, providerName, model string) (*providers.LLMResponse, error) { + // Get the correct provider for this candidate + // This is critical: different candidates may use different providers + // (e.g., cerebras, ollama, openai), and each requires its own + // provider instance with the correct API endpoint configuration. + var provider providers.LLMProvider + if agent.ProviderRegistry != nil { + p, err := agent.ProviderRegistry.GetProvider(providerName) + if err != nil { + return nil, fmt.Errorf("failed to get provider for %q: %w", providerName, err) + } + provider = p + } else { + // Fallback to default provider if registry not available + // (this shouldn't happen in normal operation) + provider = agent.Provider + } + return provider.Chat(ctx, messages, providerToolDefs, model, map[string]any{ + "max_tokens": agent.MaxTokens, + "temperature": agent.Temperature, }) }, ) @@ -538,9 +553,8 @@ func (al *AgentLoop) runLLMIteration( return fbResult.Response, nil } return agent.Provider.Chat(ctx, messages, providerToolDefs, agent.Model, map[string]any{ - "max_tokens": agent.MaxTokens, - "temperature": agent.Temperature, - "prompt_cache_key": agent.ID, + "max_tokens": agent.MaxTokens, + "temperature": agent.Temperature, }) } @@ -1004,9 +1018,8 @@ func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string) { nil, agent.Model, map[string]any{ - "max_tokens": 1024, - "temperature": 0.3, - "prompt_cache_key": agent.ID, + "max_tokens": 1024, + "temperature": 0.3, }, ) if err == nil { @@ -1055,9 +1068,8 @@ func (al *AgentLoop) summarizeBatch( nil, agent.Model, map[string]any{ - "max_tokens": 1024, - "temperature": 0.3, - "prompt_cache_key": agent.ID, + "max_tokens": 1024, + "temperature": 0.3, }, ) if err != nil { diff --git a/pkg/agent/registry.go b/pkg/agent/registry.go index 77b846832..b8915cc8a 100644 --- a/pkg/agent/registry.go +++ b/pkg/agent/registry.go @@ -26,20 +26,23 @@ func NewAgentRegistry( resolver: routing.NewRouteResolver(cfg), } + // Create provider registry from config for proper fallback support + providerRegistry := providers.NewProviderRegistry(cfg) + agentConfigs := cfg.Agents.List if len(agentConfigs) == 0 { implicitAgent := &config.AgentConfig{ ID: "main", Default: true, } - instance := NewAgentInstance(implicitAgent, &cfg.Agents.Defaults, cfg, provider) + instance := NewAgentInstance(implicitAgent, &cfg.Agents.Defaults, cfg, provider, providerRegistry) registry.agents["main"] = instance logger.InfoCF("agent", "Created implicit main agent (no agents.list configured)", nil) } else { for i := range agentConfigs { ac := &agentConfigs[i] id := routing.NormalizeAgentID(ac.ID) - instance := NewAgentInstance(ac, &cfg.Agents.Defaults, cfg, provider) + instance := NewAgentInstance(ac, &cfg.Agents.Defaults, cfg, provider, providerRegistry) registry.agents[id] = instance logger.InfoCF("agent", "Registered agent", map[string]any{ diff --git a/pkg/providers/openai_compat/provider.go b/pkg/providers/openai_compat/provider.go index 7dace71f2..232af2758 100644 --- a/pkg/providers/openai_compat/provider.go +++ b/pkg/providers/openai_compat/provider.go @@ -153,10 +153,11 @@ func (p *Provider) Chat( // with the same key and reuse prefix KV cache across calls. // The key is typically the agent ID — stable per agent, shared across requests. // See: https://platform.openai.com/docs/guides/prompt-caching - // Prompt caching is only supported by OpenAI-native endpoints. - // Gemini and other providers reject unknown fields, so skip for non-OpenAI APIs. + // IMPORTANT: Prompt caching is ONLY supported by OpenAI-native endpoints (api.openai.com). + // Gemini, Cerebras, and other providers reject unknown fields, so we must skip them. if cacheKey, ok := options["prompt_cache_key"].(string); ok && cacheKey != "" { - if !strings.Contains(p.apiBase, "generativelanguage.googleapis.com") { + // Only send prompt_cache_key to actual OpenAI endpoints + if strings.Contains(p.apiBase, "api.openai.com") { requestBody["prompt_cache_key"] = cacheKey } } diff --git a/pkg/providers/registry.go b/pkg/providers/registry.go new file mode 100644 index 000000000..0d1be8976 --- /dev/null +++ b/pkg/providers/registry.go @@ -0,0 +1,100 @@ +// PicoClaw - Ultra-lightweight personal AI agent +// License: MIT +// +// Copyright (c) 2026 PicoClaw contributors + +package providers + +import ( + "fmt" + "strings" + "sync" + + "github.com/sipeed/picoclaw/pkg/config" +) + +// ProviderRegistry manages provider instances for different provider names. +// It lazily creates providers on-demand and caches them for reuse. +// This is essential for proper fallback behavior where different candidates +// may use different providers (e.g., cerebras, ollama, openai). +type ProviderRegistry struct { + cfg *config.Config + modelList []config.ModelConfig + providers map[string]LLMProvider + mu sync.RWMutex +} + +// NewProviderRegistry creates a new provider registry with the given config. +func NewProviderRegistry(cfg *config.Config) *ProviderRegistry { + var modelList []config.ModelConfig + if cfg != nil { + modelList = cfg.ModelList + } + return &ProviderRegistry{ + cfg: cfg, + modelList: modelList, + providers: make(map[string]LLMProvider), + } +} + +// GetProvider returns a provider for the given provider name (e.g., "openai", "cerebras", "ollama"). +// If the provider has already been created, it returns the cached instance. +// Otherwise, it creates a new provider from the model list config. +func (pr *ProviderRegistry) GetProvider(providerName string) (LLMProvider, error) { + // Normalize provider name (lowercase) + providerName = strings.ToLower(providerName) + + // Check cache first + pr.mu.RLock() + if provider, ok := pr.providers[providerName]; ok { + pr.mu.RUnlock() + return provider, nil + } + pr.mu.RUnlock() + + // Find the model config for this provider + var modelCfg *config.ModelConfig + for i := range pr.modelList { + cfg := &pr.modelList[i] + // Extract provider from model string (e.g., "cerebras/gpt-oss-120b" -> "cerebras") + protocol, _ := ExtractProtocol(cfg.Model) + if strings.EqualFold(protocol, providerName) { + modelCfg = cfg + break + } + } + + // If not found in model list, try creating with protocol-only config + if modelCfg == nil { + modelCfg = &config.ModelConfig{ + Model: providerName + "/dummy", // Protocol is what matters + } + } + + // Create the provider + provider, _, err := CreateProviderFromConfig(modelCfg) + if err != nil { + return nil, fmt.Errorf("failed to create provider for %q: %w", providerName, err) + } + + // Cache the provider + pr.mu.Lock() + pr.providers[providerName] = provider + pr.mu.Unlock() + + return provider, nil +} + +// GetDefaultProvider returns the provider for the default/primary protocol. +// This is used for backward compatibility with code that expects a single provider. +func (pr *ProviderRegistry) GetDefaultProvider() (LLMProvider, error) { + if len(pr.modelList) == 0 { + // No model list configured, return error + // Caller should handle this by using a default provider + return nil, fmt.Errorf("no model list configured") + } + + // Use the first model's provider as default + protocol, _ := ExtractProtocol(pr.modelList[0].Model) + return pr.GetProvider(protocol) +}