diff --git a/pkg/agent/constants.go b/pkg/agent/constants.go new file mode 100644 index 000000000..d84a77e7d --- /dev/null +++ b/pkg/agent/constants.go @@ -0,0 +1,43 @@ +package agent + +import "time" + +// Agent loop constants — extracted from inline magic numbers for maintainability. +const ( + // DefaultMaxLLMRetries is the maximum number of retries for LLM calls + // when a context window error is detected. + DefaultMaxLLMRetries = 2 + + // SummarizeMessageThreshold is the minimum number of messages in session + // history before summarization is triggered. + SummarizeMessageThreshold = 20 + + // ContextWindowUsagePercent is the percentage of the context window that, + // when exceeded by the token estimate, triggers summarization. + ContextWindowUsagePercent = 75 + + // MessagesKeptAfterSummary is the number of most recent messages kept + // after summarization to maintain conversational continuity. + MessagesKeptAfterSummary = 4 + + // SummarizationTimeout is the maximum time allowed for the summarization + // LLM call(s) before they are cancelled. + SummarizationTimeout = 120 * time.Second + + // SummarizeMaxTokens is the max tokens requested from the LLM for + // summarization responses. + SummarizeMaxTokens = 1024 + + // SummarizeTemperature is the temperature used for summarization calls + // (low for deterministic output). + SummarizeTemperature = 0.3 + + // MinHistoryForCompression is the minimum number of messages in history + // before force compression will attempt to reduce it. + MinHistoryForCompression = 4 + + // MultiPartSummarizationThreshold is the number of messages above which + // summarization splits messages into two batches for separate summarization + // before merging. + MultiPartSummarizationThreshold = 10 +) diff --git a/pkg/agent/context_compressor.go b/pkg/agent/context_compressor.go new file mode 100644 index 000000000..d9c3bd517 --- /dev/null +++ b/pkg/agent/context_compressor.go @@ -0,0 +1,237 @@ +package agent + +import ( + "context" + "fmt" + "strings" + "sync" + "unicode/utf8" + + "github.com/sipeed/picoclaw/pkg/bus" + "github.com/sipeed/picoclaw/pkg/constants" + "github.com/sipeed/picoclaw/pkg/logger" + "github.com/sipeed/picoclaw/pkg/providers" +) + +// ContextCompressor handles context window management: force compression, +// summarization triggers, and session history optimization. +// Extracted from AgentLoop to improve separation of concerns. +type ContextCompressor struct { + bus *bus.MessageBus + summarizing *sync.Map +} + +// NewContextCompressor creates a new ContextCompressor. +func NewContextCompressor(msgBus *bus.MessageBus, summarizing *sync.Map) *ContextCompressor { + return &ContextCompressor{ + bus: msgBus, + summarizing: summarizing, + } +} + +// MaybeSummarize triggers summarization if the session history exceeds thresholds. +func (cc *ContextCompressor) MaybeSummarize(agent *AgentInstance, sessionKey, channel, chatID string) { + newHistory := agent.Sessions.GetHistory(sessionKey) + tokenEstimate := cc.EstimateTokens(newHistory) + threshold := agent.ContextWindow * ContextWindowUsagePercent / 100 + + if len(newHistory) > SummarizeMessageThreshold || tokenEstimate > threshold { + summarizeKey := agent.ID + ":" + sessionKey + if _, loading := cc.summarizing.LoadOrStore(summarizeKey, true); !loading { + go func() { + defer cc.summarizing.Delete(summarizeKey) + if !constants.IsInternalChannel(channel) { + cc.bus.PublishOutbound(bus.OutboundMessage{ + Channel: channel, + ChatID: chatID, + Content: "Memory threshold reached. Optimizing conversation history...", + }) + } + cc.SummarizeSession(agent, sessionKey) + }() + } + } +} + +// ForceCompression aggressively reduces context when the limit is hit. +// It drops the oldest 50% of messages (keeping system prompt and last user message). +func (cc *ContextCompressor) ForceCompression(agent *AgentInstance, sessionKey string) { + history := agent.Sessions.GetHistory(sessionKey) + if len(history) <= MinHistoryForCompression { + return + } + + // Keep system prompt (usually [0]) and the very last message (user's trigger) + // We want to drop the oldest half of the *conversation* + // Assuming [0] is system, [1:] is conversation + conversation := history[1 : len(history)-1] + if len(conversation) == 0 { + return + } + + // Helper to find the mid-point of the conversation + mid := len(conversation) / 2 + + // New history structure: + // 1. System Prompt (with compression note appended) + // 2. Second half of conversation + // 3. Last message + + droppedCount := mid + keptConversation := conversation[mid:] + + newHistory := make([]providers.Message, 0) + + // Append compression note to the original system prompt instead of adding a new system message + // This avoids having two consecutive system messages which some APIs (like Zhipu) reject + compressionNote := fmt.Sprintf( + "\n\n[System Note: Emergency compression dropped %d oldest messages due to context limit]", + droppedCount, + ) + enhancedSystemPrompt := history[0] + enhancedSystemPrompt.Content = enhancedSystemPrompt.Content + compressionNote + newHistory = append(newHistory, enhancedSystemPrompt) + + newHistory = append(newHistory, keptConversation...) + newHistory = append(newHistory, history[len(history)-1]) // Last message + + // Update session + agent.Sessions.SetHistory(sessionKey, newHistory) + agent.Sessions.Save(sessionKey) + + logger.WarnCF("agent", "Forced compression executed", map[string]any{ + "session_key": sessionKey, + "dropped_msgs": droppedCount, + "new_count": len(newHistory), + }) +} + +// SummarizeSession summarizes the conversation history for a session. +func (cc *ContextCompressor) SummarizeSession(agent *AgentInstance, sessionKey string) { + ctx, cancel := context.WithTimeout(context.Background(), SummarizationTimeout) + defer cancel() + + history := agent.Sessions.GetHistory(sessionKey) + summary := agent.Sessions.GetSummary(sessionKey) + + // Keep last N messages for continuity + if len(history) <= MessagesKeptAfterSummary { + return + } + + toSummarize := history[:len(history)-MessagesKeptAfterSummary] + + // Oversized Message Guard + maxMessageTokens := agent.ContextWindow / 2 + validMessages := make([]providers.Message, 0) + omitted := false + + for _, m := range toSummarize { + if m.Role != "user" && m.Role != "assistant" { + continue + } + // Use character-based estimation (2.5 chars per token = totalChars * 2 / 5) + msgTokens := utf8.RuneCountInString(m.Content) * 2 / 5 + if msgTokens > maxMessageTokens { + omitted = true + continue + } + validMessages = append(validMessages, m) + } + + if len(validMessages) == 0 { + return + } + + // Multi-Part Summarization + var finalSummary string + if len(validMessages) > MultiPartSummarizationThreshold { + mid := len(validMessages) / 2 + part1 := validMessages[:mid] + part2 := validMessages[mid:] + + s1, _ := cc.SummarizeBatch(ctx, agent, part1, "") + s2, _ := cc.SummarizeBatch(ctx, agent, part2, "") + + mergePrompt := fmt.Sprintf( + "Merge these two conversation summaries into one cohesive summary:\n\n1: %s\n\n2: %s", + s1, + s2, + ) + resp, err := agent.Provider.Chat( + ctx, + []providers.Message{{Role: "user", Content: mergePrompt}}, + nil, + agent.Model, + map[string]any{ + "max_tokens": SummarizeMaxTokens, + "temperature": SummarizeTemperature, + }, + ) + if err == nil { + finalSummary = resp.Content + } else { + finalSummary = s1 + " " + s2 + } + } else { + finalSummary, _ = cc.SummarizeBatch(ctx, agent, validMessages, summary) + } + + if omitted && finalSummary != "" { + finalSummary += "\n[Note: Some oversized messages were omitted from this summary for efficiency.]" + } + + if finalSummary != "" { + agent.Sessions.SetSummary(sessionKey, finalSummary) + agent.Sessions.TruncateHistory(sessionKey, MessagesKeptAfterSummary) + agent.Sessions.Save(sessionKey) + } +} + +// SummarizeBatch summarizes a batch of messages. +func (cc *ContextCompressor) SummarizeBatch( + ctx context.Context, + agent *AgentInstance, + batch []providers.Message, + existingSummary string, +) (string, error) { + var sb strings.Builder + sb.WriteString("Provide a concise summary of this conversation segment, preserving core context and key points.\n") + if existingSummary != "" { + sb.WriteString("Existing context: ") + sb.WriteString(existingSummary) + sb.WriteString("\n") + } + sb.WriteString("\nCONVERSATION:\n") + for _, m := range batch { + fmt.Fprintf(&sb, "%s: %s\n", m.Role, m.Content) + } + prompt := sb.String() + + response, err := agent.Provider.Chat( + ctx, + []providers.Message{{Role: "user", Content: prompt}}, + nil, + agent.Model, + map[string]any{ + "max_tokens": SummarizeMaxTokens, + "temperature": SummarizeTemperature, + }, + ) + if err != nil { + return "", err + } + return response.Content, nil +} + +// EstimateTokens estimates the number of tokens in a message list. +// Uses a safe heuristic of 2.5 characters per token to account for CJK and other +// overheads better than the previous 3 chars/token. +func (cc *ContextCompressor) EstimateTokens(messages []providers.Message) int { + totalChars := 0 + for _, m := range messages { + totalChars += utf8.RuneCountInString(m.Content) + } + // 2.5 chars per token = totalChars * 2 / 5 + return totalChars * 2 / 5 +} diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index bf229ad74..2cfb77845 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -14,7 +14,6 @@ import ( "sync" "sync/atomic" "time" - "unicode/utf8" "github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/channels" @@ -38,6 +37,15 @@ type AgentLoop struct { summarizing sync.Map fallback *providers.FallbackChain channelManager *channels.Manager + + // Extracted components for better separation of concerns + compressor *ContextCompressor + toolExec *ToolExecutor + + // Graceful shutdown support + cancelCtx context.Context + cancelFunc context.CancelFunc + wg sync.WaitGroup } // processOptions configures how a message is processed @@ -69,13 +77,24 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers stateManager = state.NewManager(defaultAgent.Workspace) } + // Initialize extracted components + summarizing := &sync.Map{} + compressor := NewContextCompressor(msgBus, summarizing) + toolExec := NewToolExecutor(msgBus) + + // Create cancellation context for graceful shutdown + cancelCtx, cancelFunc := context.WithCancel(context.Background()) + return &AgentLoop{ - bus: msgBus, - cfg: cfg, - registry: registry, - state: stateManager, - summarizing: sync.Map{}, - fallback: fallbackChain, + bus: msgBus, + cfg: cfg, + registry: registry, + state: stateManager, + fallback: fallbackChain, + compressor: compressor, + toolExec: toolExec, + cancelCtx: cancelCtx, + cancelFunc: cancelFunc, } } @@ -202,6 +221,10 @@ func (al *AgentLoop) Run(ctx context.Context) error { func (al *AgentLoop) Stop() { al.running.Store(false) + if al.cancelFunc != nil { + al.cancelFunc() + } + al.wg.Wait() } func (al *AgentLoop) RegisterTool(tool tools.Tool) { @@ -404,7 +427,7 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt } // 1. Update tool contexts - al.updateToolContexts(agent, opts.Channel, opts.ChatID) + al.toolExec.UpdateToolContexts(agent, opts.Channel, opts.ChatID) // 2. Build messages (skip history for heartbeat) var history []providers.Message @@ -445,7 +468,7 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt // 7. Optional: summarization if opts.EnableSummary { - al.maybeSummarize(agent, opts.SessionKey, opts.Channel, opts.ChatID) + al.compressor.MaybeSummarize(agent, opts.SessionKey, opts.Channel, opts.ChatID) } // 8. Optional: send response via bus @@ -545,18 +568,14 @@ func (al *AgentLoop) runLLMIteration( } // Retry loop for context/token errors - maxRetries := 2 + maxRetries := DefaultMaxLLMRetries for retry := 0; retry <= maxRetries; retry++ { response, err = callLLM() if err == nil { break } - errMsg := strings.ToLower(err.Error()) - isContextError := strings.Contains(errMsg, "token") || - strings.Contains(errMsg, "context") || - strings.Contains(errMsg, "invalidparameter") || - strings.Contains(errMsg, "length") + isContextError := providers.IsContextWindowError(err) if isContextError && retry < maxRetries { logger.WarnCF("agent", "Context window error detected, attempting compression", map[string]any{ @@ -572,7 +591,7 @@ func (al *AgentLoop) runLLMIteration( }) } - al.forceCompression(agent, opts.SessionKey) + al.compressor.ForceCompression(agent, opts.SessionKey) newHistory := agent.Sessions.GetHistory(opts.SessionKey) newSummary := agent.Sessions.GetSummary(opts.SessionKey) messages = agent.ContextBuilder.BuildMessages( @@ -656,174 +675,17 @@ func (al *AgentLoop) runLLMIteration( // Save assistant message with tool calls to session agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg) - // Execute tool calls - for _, tc := range normalizedToolCalls { - argsJSON, _ := json.Marshal(tc.Arguments) - argsPreview := utils.Truncate(string(argsJSON), 200) - logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview), - map[string]any{ - "agent_id": agent.ID, - "tool": tc.Name, - "iteration": iteration, - }) - - // Create async callback for tools that implement AsyncTool - // NOTE: Following openclaw's design, async tools do NOT send results directly to users. - // Instead, they notify the agent via PublishInbound, and the agent decides - // whether to forward the result to the user (in processSystemMessage). - asyncCallback := func(callbackCtx context.Context, result *tools.ToolResult) { - // Log the async completion but don't send directly to user - // The agent will handle user notification via processSystemMessage - if !result.Silent && result.ForUser != "" { - logger.InfoCF("agent", "Async tool completed, agent will handle notification", - map[string]any{ - "tool": tc.Name, - "content_len": len(result.ForUser), - }) - } - } - - toolResult := agent.Tools.ExecuteWithContext( - ctx, - tc.Name, - tc.Arguments, - opts.Channel, - opts.ChatID, - asyncCallback, - ) - - // Send ForUser content to user immediately if not Silent - if !toolResult.Silent && toolResult.ForUser != "" && opts.SendResponse { - al.bus.PublishOutbound(bus.OutboundMessage{ - Channel: opts.Channel, - ChatID: opts.ChatID, - Content: toolResult.ForUser, - }) - logger.DebugCF("agent", "Sent tool result to user", - map[string]any{ - "tool": tc.Name, - "content_len": len(toolResult.ForUser), - }) - } - - // Determine content for LLM based on tool result - contentForLLM := toolResult.ForLLM - if contentForLLM == "" && toolResult.Err != nil { - contentForLLM = toolResult.Err.Error() - } - - toolResultMsg := providers.Message{ - Role: "tool", - Content: contentForLLM, - ToolCallID: tc.ID, - } - messages = append(messages, toolResultMsg) - - // Save tool result message to session - agent.Sessions.AddFullMessage(opts.SessionKey, toolResultMsg) + // Execute tool calls via the ToolExecutor component + toolResultMsgs := al.toolExec.ExecuteToolCalls(ctx, agent, normalizedToolCalls, opts) + for _, resultMsg := range toolResultMsgs { + messages = append(messages, resultMsg) + agent.Sessions.AddFullMessage(opts.SessionKey, resultMsg) } } return finalContent, iteration, nil } -// updateToolContexts updates the context for tools that need channel/chatID info. -func (al *AgentLoop) updateToolContexts(agent *AgentInstance, channel, chatID string) { - // Use ContextualTool interface instead of type assertions - if tool, ok := agent.Tools.Get("message"); ok { - if mt, ok := tool.(tools.ContextualTool); ok { - mt.SetContext(channel, chatID) - } - } - if tool, ok := agent.Tools.Get("spawn"); ok { - if st, ok := tool.(tools.ContextualTool); ok { - st.SetContext(channel, chatID) - } - } - if tool, ok := agent.Tools.Get("subagent"); ok { - if st, ok := tool.(tools.ContextualTool); ok { - st.SetContext(channel, chatID) - } - } -} - -// maybeSummarize triggers summarization if the session history exceeds thresholds. -func (al *AgentLoop) maybeSummarize(agent *AgentInstance, sessionKey, channel, chatID string) { - newHistory := agent.Sessions.GetHistory(sessionKey) - tokenEstimate := al.estimateTokens(newHistory) - threshold := agent.ContextWindow * 75 / 100 - - if len(newHistory) > 20 || tokenEstimate > threshold { - summarizeKey := agent.ID + ":" + sessionKey - if _, loading := al.summarizing.LoadOrStore(summarizeKey, true); !loading { - go func() { - defer al.summarizing.Delete(summarizeKey) - if !constants.IsInternalChannel(channel) { - al.bus.PublishOutbound(bus.OutboundMessage{ - Channel: channel, - ChatID: chatID, - Content: "Memory threshold reached. Optimizing conversation history...", - }) - } - al.summarizeSession(agent, sessionKey) - }() - } - } -} - -// forceCompression aggressively reduces context when the limit is hit. -// It drops the oldest 50% of messages (keeping system prompt and last user message). -func (al *AgentLoop) forceCompression(agent *AgentInstance, sessionKey string) { - history := agent.Sessions.GetHistory(sessionKey) - if len(history) <= 4 { - return - } - - // Keep system prompt (usually [0]) and the very last message (user's trigger) - // We want to drop the oldest half of the *conversation* - // Assuming [0] is system, [1:] is conversation - conversation := history[1 : len(history)-1] - if len(conversation) == 0 { - return - } - - // Helper to find the mid-point of the conversation - mid := len(conversation) / 2 - - // New history structure: - // 1. System Prompt (with compression note appended) - // 2. Second half of conversation - // 3. Last message - - droppedCount := mid - keptConversation := conversation[mid:] - - newHistory := make([]providers.Message, 0) - - // Append compression note to the original system prompt instead of adding a new system message - // This avoids having two consecutive system messages which some APIs (like Zhipu) reject - compressionNote := fmt.Sprintf( - "\n\n[System Note: Emergency compression dropped %d oldest messages due to context limit]", - droppedCount, - ) - enhancedSystemPrompt := history[0] - enhancedSystemPrompt.Content = enhancedSystemPrompt.Content + compressionNote - newHistory = append(newHistory, enhancedSystemPrompt) - - newHistory = append(newHistory, keptConversation...) - newHistory = append(newHistory, history[len(history)-1]) // Last message - - // Update session - agent.Sessions.SetHistory(sessionKey, newHistory) - agent.Sessions.Save(sessionKey) - - logger.WarnCF("agent", "Forced compression executed", map[string]any{ - "session_key": sessionKey, - "dropped_msgs": droppedCount, - "new_count": len(newHistory), - }) -} - // GetStartupInfo returns information about loaded tools and skills for logging. func (al *AgentLoop) GetStartupInfo() map[string]any { info := make(map[string]any) @@ -903,135 +765,6 @@ func formatToolsForLog(toolDefs []providers.ToolDefinition) string { return sb.String() } -// summarizeSession summarizes the conversation history for a session. -func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string) { - ctx, cancel := context.WithTimeout(context.Background(), 120*time.Second) - defer cancel() - - history := agent.Sessions.GetHistory(sessionKey) - summary := agent.Sessions.GetSummary(sessionKey) - - // Keep last 4 messages for continuity - if len(history) <= 4 { - return - } - - toSummarize := history[:len(history)-4] - - // Oversized Message Guard - maxMessageTokens := agent.ContextWindow / 2 - validMessages := make([]providers.Message, 0) - omitted := false - - for _, m := range toSummarize { - if m.Role != "user" && m.Role != "assistant" { - continue - } - msgTokens := len(m.Content) / 2 - if msgTokens > maxMessageTokens { - omitted = true - continue - } - validMessages = append(validMessages, m) - } - - if len(validMessages) == 0 { - return - } - - // Multi-Part Summarization - var finalSummary string - if len(validMessages) > 10 { - mid := len(validMessages) / 2 - part1 := validMessages[:mid] - part2 := validMessages[mid:] - - s1, _ := al.summarizeBatch(ctx, agent, part1, "") - s2, _ := al.summarizeBatch(ctx, agent, part2, "") - - mergePrompt := fmt.Sprintf( - "Merge these two conversation summaries into one cohesive summary:\n\n1: %s\n\n2: %s", - s1, - s2, - ) - resp, err := agent.Provider.Chat( - ctx, - []providers.Message{{Role: "user", Content: mergePrompt}}, - nil, - agent.Model, - map[string]any{ - "max_tokens": 1024, - "temperature": 0.3, - }, - ) - if err == nil { - finalSummary = resp.Content - } else { - finalSummary = s1 + " " + s2 - } - } else { - finalSummary, _ = al.summarizeBatch(ctx, agent, validMessages, summary) - } - - if omitted && finalSummary != "" { - finalSummary += "\n[Note: Some oversized messages were omitted from this summary for efficiency.]" - } - - if finalSummary != "" { - agent.Sessions.SetSummary(sessionKey, finalSummary) - agent.Sessions.TruncateHistory(sessionKey, 4) - agent.Sessions.Save(sessionKey) - } -} - -// summarizeBatch summarizes a batch of messages. -func (al *AgentLoop) summarizeBatch( - ctx context.Context, - agent *AgentInstance, - batch []providers.Message, - existingSummary string, -) (string, error) { - var sb strings.Builder - sb.WriteString("Provide a concise summary of this conversation segment, preserving core context and key points.\n") - if existingSummary != "" { - sb.WriteString("Existing context: ") - sb.WriteString(existingSummary) - sb.WriteString("\n") - } - sb.WriteString("\nCONVERSATION:\n") - for _, m := range batch { - fmt.Fprintf(&sb, "%s: %s\n", m.Role, m.Content) - } - prompt := sb.String() - - response, err := agent.Provider.Chat( - ctx, - []providers.Message{{Role: "user", Content: prompt}}, - nil, - agent.Model, - map[string]any{ - "max_tokens": 1024, - "temperature": 0.3, - }, - ) - if err != nil { - return "", err - } - return response.Content, nil -} - -// estimateTokens estimates the number of tokens in a message list. -// Uses a safe heuristic of 2.5 characters per token to account for CJK and other -// overheads better than the previous 3 chars/token. -func (al *AgentLoop) estimateTokens(messages []providers.Message) int { - totalChars := 0 - for _, m := range messages { - totalChars += utf8.RuneCountInString(m.Content) - } - // 2.5 chars per token = totalChars * 2 / 5 - return totalChars * 2 / 5 -} - func (al *AgentLoop) handleCommand(ctx context.Context, msg bus.InboundMessage) (string, bool) { content := strings.TrimSpace(msg.Content) if !strings.HasPrefix(content, "/") { diff --git a/pkg/agent/tool_executor.go b/pkg/agent/tool_executor.go new file mode 100644 index 000000000..a13a0dccd --- /dev/null +++ b/pkg/agent/tool_executor.go @@ -0,0 +1,111 @@ +package agent + +import ( + "context" + "encoding/json" + "fmt" + + "github.com/sipeed/picoclaw/pkg/bus" + "github.com/sipeed/picoclaw/pkg/logger" + "github.com/sipeed/picoclaw/pkg/providers" + "github.com/sipeed/picoclaw/pkg/tools" + "github.com/sipeed/picoclaw/pkg/utils" +) + +// ToolExecutor handles executing tool calls requested by the LLM, +// including async callback setup, result routing, and context updates. +// Extracted from AgentLoop to improve separation of concerns. +type ToolExecutor struct { + bus *bus.MessageBus +} + +// NewToolExecutor creates a new ToolExecutor. +func NewToolExecutor(msgBus *bus.MessageBus) *ToolExecutor { + return &ToolExecutor{bus: msgBus} +} + +// ExecuteToolCalls processes a list of normalized tool calls, executes each one, +// sends user-facing results to the bus, and returns the tool result messages +// to be appended to the conversation. +func (te *ToolExecutor) ExecuteToolCalls( + ctx context.Context, + agent *AgentInstance, + toolCalls []providers.ToolCall, + opts processOptions, +) []providers.Message { + var resultMessages []providers.Message + + for _, tc := range toolCalls { + argsJSON, _ := json.Marshal(tc.Arguments) + argsPreview := utils.Truncate(string(argsJSON), 200) + logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview), + map[string]any{ + "agent_id": agent.ID, + "tool": tc.Name, + }) + + // Create async callback for tools that implement AsyncTool + // NOTE: Following openclaw's design, async tools do NOT send results directly to users. + // Instead, they notify the agent via PublishInbound, and the agent decides + // whether to forward the result to the user (in processSystemMessage). + asyncCallback := func(callbackCtx context.Context, result *tools.ToolResult) { + if !result.Silent && result.ForUser != "" { + logger.InfoCF("agent", "Async tool completed, agent will handle notification", + map[string]any{ + "tool": tc.Name, + "content_len": len(result.ForUser), + }) + } + } + + toolResult := agent.Tools.ExecuteWithContext( + ctx, + tc.Name, + tc.Arguments, + opts.Channel, + opts.ChatID, + asyncCallback, + ) + + // Send ForUser content to user immediately if not Silent + if !toolResult.Silent && toolResult.ForUser != "" && opts.SendResponse { + te.bus.PublishOutbound(bus.OutboundMessage{ + Channel: opts.Channel, + ChatID: opts.ChatID, + Content: toolResult.ForUser, + }) + logger.DebugCF("agent", "Sent tool result to user", + map[string]any{ + "tool": tc.Name, + "content_len": len(toolResult.ForUser), + }) + } + + // Determine content for LLM based on tool result + contentForLLM := toolResult.ForLLM + if contentForLLM == "" && toolResult.Err != nil { + contentForLLM = toolResult.Err.Error() + } + + resultMessages = append(resultMessages, providers.Message{ + Role: "tool", + Content: contentForLLM, + ToolCallID: tc.ID, + }) + } + + return resultMessages +} + +// UpdateToolContexts updates the context for tools that need channel/chatID info. +func (te *ToolExecutor) UpdateToolContexts(agent *AgentInstance, channel, chatID string) { + // Use ContextualTool interface instead of type assertions + contextualToolNames := []string{"message", "spawn", "subagent"} + for _, name := range contextualToolNames { + if tool, ok := agent.Tools.Get(name); ok { + if ct, ok := tool.(tools.ContextualTool); ok { + ct.SetContext(channel, chatID) + } + } + } +} diff --git a/pkg/providers/error_classifier.go b/pkg/providers/error_classifier.go index a0f003006..9125f2f36 100644 --- a/pkg/providers/error_classifier.go +++ b/pkg/providers/error_classifier.go @@ -86,6 +86,18 @@ var ( rxp(`image exceeds.*mb`), } + // contextWindowPatterns detects context window / token limit exhaustion errors. + contextWindowPatterns = []errorPattern{ + rxp(`context.*(window|length).*exceed`), + rxp(`max.*token.*(limit|exceed)`), + rxp(`total tokens.*exceed`), + substr("context_length_exceeded"), + substr("maximum context length"), + substr("invalidparameter"), + rxp(`tokens?.*exceed`), + rxp(`exceeds? the (model|max).*token`), + } + // Transient HTTP status codes that map to timeout (server-side failures). transientStatusCodes = map[int]bool{ 500: true, 502: true, 503: true, @@ -225,6 +237,17 @@ func IsImageSizeError(msg string) bool { return matchesAny(msg, imageSizePatterns) } +// IsContextWindowError returns true if the error indicates a context window / token limit +// exhaustion. This is more precise than naive string matching and avoids false positives +// on errors like "invalid authentication token". +func IsContextWindowError(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return matchesAny(msg, contextWindowPatterns) +} + // matchesAny checks if msg matches any of the patterns. func matchesAny(msg string, patterns []errorPattern) bool { for _, p := range patterns { diff --git a/pkg/providers/factory.go b/pkg/providers/factory.go index cda4753ea..9a2c7a771 100644 --- a/pkg/providers/factory.go +++ b/pkg/providers/factory.go @@ -35,6 +35,268 @@ type providerSelection struct { enableWebSearch bool } +// providerDefaults holds the default API base URL and a config accessor for a standard HTTP-compatible provider. +type providerDefaults struct { + defaultBase string + getConfig func(cfg *config.Config) (apiKey, apiBase, proxy string) + // hasKey returns true if the provider has credentials configured. + // If nil, checks that getConfig returns a non-empty apiKey. + hasKey func(cfg *config.Config) bool +} + +// standardProviderRegistry maps provider names to their defaults. +// Only covers providers that follow the standard pattern: apiKey + apiBase + proxy. +// Special-case providers (CLI, OAuth, Copilot) are handled separately. +var standardProviderRegistry = map[string]providerDefaults{ + "groq": { + defaultBase: "https://api.groq.com/openai/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Groq.APIKey, cfg.Providers.Groq.APIBase, cfg.Providers.Groq.Proxy + }, + }, + "openrouter": { + defaultBase: "https://openrouter.ai/api/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.OpenRouter.APIKey, cfg.Providers.OpenRouter.APIBase, cfg.Providers.OpenRouter.Proxy + }, + }, + "zhipu": { + defaultBase: "https://open.bigmodel.cn/api/paas/v4", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Zhipu.APIKey, cfg.Providers.Zhipu.APIBase, cfg.Providers.Zhipu.Proxy + }, + }, + "gemini": { + defaultBase: "https://generativelanguage.googleapis.com/v1beta", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Gemini.APIKey, cfg.Providers.Gemini.APIBase, cfg.Providers.Gemini.Proxy + }, + }, + "vllm": { + defaultBase: "", // no default base; requires explicit config + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.VLLM.APIKey, cfg.Providers.VLLM.APIBase, cfg.Providers.VLLM.Proxy + }, + hasKey: func(cfg *config.Config) bool { + return cfg.Providers.VLLM.APIBase != "" + }, + }, + "shengsuanyun": { + defaultBase: "https://router.shengsuanyun.com/api/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.ShengSuanYun.APIKey, cfg.Providers.ShengSuanYun.APIBase, cfg.Providers.ShengSuanYun.Proxy + }, + }, + "nvidia": { + defaultBase: "https://integrate.api.nvidia.com/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Nvidia.APIKey, cfg.Providers.Nvidia.APIBase, cfg.Providers.Nvidia.Proxy + }, + }, + "deepseek": { + defaultBase: "https://api.deepseek.com/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.DeepSeek.APIKey, cfg.Providers.DeepSeek.APIBase, cfg.Providers.DeepSeek.Proxy + }, + }, + "mistral": { + defaultBase: "https://api.mistral.ai/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Mistral.APIKey, cfg.Providers.Mistral.APIBase, cfg.Providers.Mistral.Proxy + }, + }, + "ollama": { + defaultBase: "http://localhost:11434/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Ollama.APIKey, cfg.Providers.Ollama.APIBase, cfg.Providers.Ollama.Proxy + }, + }, + "moonshot": { + defaultBase: "https://api.moonshot.cn/v1", + getConfig: func(cfg *config.Config) (string, string, string) { + return cfg.Providers.Moonshot.APIKey, cfg.Providers.Moonshot.APIBase, cfg.Providers.Moonshot.Proxy + }, + }, +} + +// providerNameAliases maps alternative provider names to their canonical names in the registry. +var providerNameAliases = map[string]string{ + "glm": "zhipu", + "google": "gemini", +} + +// applyStandardProvider applies config from a standardProviderRegistry entry to sel. +// Returns true if the provider had credentials configured. +func applyStandardProvider(cfg *config.Config, sel *providerSelection, entry providerDefaults) bool { + apiKey, apiBase, proxy := entry.getConfig(cfg) + + hasCredentials := apiKey != "" + if entry.hasKey != nil { + hasCredentials = entry.hasKey(cfg) + } + + if !hasCredentials { + return false + } + + sel.apiKey = apiKey + sel.apiBase = apiBase + sel.proxy = proxy + if sel.apiBase == "" && entry.defaultBase != "" { + sel.apiBase = entry.defaultBase + } + return true +} + +// modelInferenceEntry maps a model name pattern to a provider and optional match function. +type modelInferenceEntry struct { + // matches returns true if this entry should handle the given model name (lowercase) and original model. + matches func(lowerModel, model string, cfg *config.Config) bool + // apply sets up the provider selection. Returns true on success. + apply func(cfg *config.Config, sel *providerSelection) bool +} + +// modelInferenceRegistry defines fallback model → provider inference rules. +// Order matters: first match wins. +var modelInferenceRegistry = []modelInferenceEntry{ + // Moonshot/Kimi + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "kimi") || strings.Contains(lm, "moonshot") || strings.HasPrefix(m, "moonshot/")) && + cfg.Providers.Moonshot.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["moonshot"]) + }, + }, + // OpenRouter-prefixed models (openrouter/, anthropic/, openai/, meta-llama/, deepseek/, google/) + { + matches: func(_, m string, _ *config.Config) bool { + for _, prefix := range []string{"openrouter/", "anthropic/", "openai/", "meta-llama/", "deepseek/", "google/"} { + if strings.HasPrefix(m, prefix) { + return true + } + } + return false + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["openrouter"]) + }, + }, + // Claude models → Anthropic (with OAuth support) + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "claude") || strings.HasPrefix(m, "anthropic/")) && + (cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != "") + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" { + sel.apiBase = cfg.Providers.Anthropic.APIBase + if sel.apiBase == "" { + sel.apiBase = defaultAnthropicAPIBase + } + sel.providerType = providerTypeClaudeAuth + return true + } + sel.apiKey = cfg.Providers.Anthropic.APIKey + sel.apiBase = cfg.Providers.Anthropic.APIBase + sel.proxy = cfg.Providers.Anthropic.Proxy + if sel.apiBase == "" { + sel.apiBase = defaultAnthropicAPIBase + } + return true + }, + }, + // GPT models → OpenAI (with OAuth/codex-cli support) + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "gpt") || strings.HasPrefix(m, "openai/")) && + (cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != "") + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + sel.enableWebSearch = cfg.Providers.OpenAI.WebSearch + if cfg.Providers.OpenAI.AuthMethod == "codex-cli" { + sel.providerType = providerTypeCodexCLIToken + return true + } + if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" { + sel.providerType = providerTypeCodexAuth + return true + } + sel.apiKey = cfg.Providers.OpenAI.APIKey + sel.apiBase = cfg.Providers.OpenAI.APIBase + sel.proxy = cfg.Providers.OpenAI.Proxy + if sel.apiBase == "" { + sel.apiBase = "https://api.openai.com/v1" + } + return true + }, + }, + // Gemini + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "gemini") || strings.HasPrefix(m, "google/")) && cfg.Providers.Gemini.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["gemini"]) + }, + }, + // Zhipu/GLM + { + matches: func(lm, _ string, cfg *config.Config) bool { + return (strings.Contains(lm, "glm") || strings.Contains(lm, "zhipu") || strings.Contains(lm, "zai")) && cfg.Providers.Zhipu.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["zhipu"]) + }, + }, + // Groq + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "groq") || strings.HasPrefix(m, "groq/")) && cfg.Providers.Groq.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["groq"]) + }, + }, + // Nvidia + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "nvidia") || strings.HasPrefix(m, "nvidia/")) && cfg.Providers.Nvidia.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["nvidia"]) + }, + }, + // Ollama + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "ollama") || strings.HasPrefix(m, "ollama/")) && cfg.Providers.Ollama.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["ollama"]) + }, + }, + // Mistral + { + matches: func(lm, m string, cfg *config.Config) bool { + return (strings.Contains(lm, "mistral") || strings.HasPrefix(m, "mistral/")) && cfg.Providers.Mistral.APIKey != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["mistral"]) + }, + }, + // VLLM (fallback if API base is configured) + { + matches: func(_, _ string, cfg *config.Config) bool { + return cfg.Providers.VLLM.APIBase != "" + }, + apply: func(cfg *config.Config, sel *providerSelection) bool { + return applyStandardProvider(cfg, sel, standardProviderRegistry["vllm"]) + }, + }, +} + func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { model := cfg.Agents.Defaults.Model providerName := strings.ToLower(cfg.Agents.Defaults.Provider) @@ -47,16 +309,8 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { // First, prefer explicit provider configuration. if providerName != "" { + // Handle special-case providers that have unique instantiation logic. switch providerName { - case "groq": - if cfg.Providers.Groq.APIKey != "" { - sel.apiKey = cfg.Providers.Groq.APIKey - sel.apiBase = cfg.Providers.Groq.APIBase - sel.proxy = cfg.Providers.Groq.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.groq.com/openai/v1" - } - } case "openai", "gpt": if cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != "" { sel.enableWebSearch = cfg.Providers.OpenAI.WebSearch @@ -92,58 +346,6 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { sel.apiBase = defaultAnthropicAPIBase } } - case "openrouter": - if cfg.Providers.OpenRouter.APIKey != "" { - sel.apiKey = cfg.Providers.OpenRouter.APIKey - sel.proxy = cfg.Providers.OpenRouter.Proxy - if cfg.Providers.OpenRouter.APIBase != "" { - sel.apiBase = cfg.Providers.OpenRouter.APIBase - } else { - sel.apiBase = "https://openrouter.ai/api/v1" - } - } - case "zhipu", "glm": - if cfg.Providers.Zhipu.APIKey != "" { - sel.apiKey = cfg.Providers.Zhipu.APIKey - sel.apiBase = cfg.Providers.Zhipu.APIBase - sel.proxy = cfg.Providers.Zhipu.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://open.bigmodel.cn/api/paas/v4" - } - } - case "gemini", "google": - if cfg.Providers.Gemini.APIKey != "" { - sel.apiKey = cfg.Providers.Gemini.APIKey - sel.apiBase = cfg.Providers.Gemini.APIBase - sel.proxy = cfg.Providers.Gemini.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://generativelanguage.googleapis.com/v1beta" - } - } - case "vllm": - if cfg.Providers.VLLM.APIBase != "" { - sel.apiKey = cfg.Providers.VLLM.APIKey - sel.apiBase = cfg.Providers.VLLM.APIBase - sel.proxy = cfg.Providers.VLLM.Proxy - } - case "shengsuanyun": - if cfg.Providers.ShengSuanYun.APIKey != "" { - sel.apiKey = cfg.Providers.ShengSuanYun.APIKey - sel.apiBase = cfg.Providers.ShengSuanYun.APIBase - sel.proxy = cfg.Providers.ShengSuanYun.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://router.shengsuanyun.com/api/v1" - } - } - case "nvidia": - if cfg.Providers.Nvidia.APIKey != "" { - sel.apiKey = cfg.Providers.Nvidia.APIKey - sel.apiBase = cfg.Providers.Nvidia.APIBase - sel.proxy = cfg.Providers.Nvidia.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://integrate.api.nvidia.com/v1" - } - } case "claude-cli", "claude-code", "claudecode": workspace := cfg.WorkspacePath() if workspace == "" { @@ -160,27 +362,6 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { sel.providerType = providerTypeCodexCLI sel.workspace = workspace return sel, nil - case "deepseek": - if cfg.Providers.DeepSeek.APIKey != "" { - sel.apiKey = cfg.Providers.DeepSeek.APIKey - sel.apiBase = cfg.Providers.DeepSeek.APIBase - sel.proxy = cfg.Providers.DeepSeek.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.deepseek.com/v1" - } - if model != "deepseek-chat" && model != "deepseek-reasoner" { - sel.model = "deepseek-chat" - } - } - case "mistral": - if cfg.Providers.Mistral.APIKey != "" { - sel.apiKey = cfg.Providers.Mistral.APIKey - sel.apiBase = cfg.Providers.Mistral.APIBase - sel.proxy = cfg.Providers.Mistral.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.mistral.ai/v1" - } - } case "github_copilot", "copilot": sel.providerType = providerTypeGitHubCopilot if cfg.Providers.GitHubCopilot.APIBase != "" { @@ -190,120 +371,38 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { } sel.connectMode = cfg.Providers.GitHubCopilot.ConnectMode return sel, nil + default: + // Try standard provider registry lookup. + canonicalName := providerName + if alias, ok := providerNameAliases[providerName]; ok { + canonicalName = alias + } + if entry, ok := standardProviderRegistry[canonicalName]; ok { + if applyStandardProvider(cfg, &sel, entry) { + // Handle deepseek model override. + if canonicalName == "deepseek" && model != "deepseek-chat" && model != "deepseek-reasoner" { + sel.model = "deepseek-chat" + } + } + } } } - // Fallback: infer provider from model and configured keys. + // Fallback: infer provider from model name and configured keys. if sel.apiKey == "" && sel.apiBase == "" { - switch { - case (strings.Contains(lowerModel, "kimi") || strings.Contains(lowerModel, "moonshot") || strings.HasPrefix(model, "moonshot/")) && cfg.Providers.Moonshot.APIKey != "": - sel.apiKey = cfg.Providers.Moonshot.APIKey - sel.apiBase = cfg.Providers.Moonshot.APIBase - sel.proxy = cfg.Providers.Moonshot.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.moonshot.cn/v1" + matched := false + for _, entry := range modelInferenceRegistry { + if entry.matches(lowerModel, model, cfg) { + entry.apply(cfg, &sel) + matched = true + break } - case strings.HasPrefix(model, "openrouter/") || - strings.HasPrefix(model, "anthropic/") || - strings.HasPrefix(model, "openai/") || - strings.HasPrefix(model, "meta-llama/") || - strings.HasPrefix(model, "deepseek/") || - strings.HasPrefix(model, "google/"): - sel.apiKey = cfg.Providers.OpenRouter.APIKey - sel.proxy = cfg.Providers.OpenRouter.Proxy - if cfg.Providers.OpenRouter.APIBase != "" { - sel.apiBase = cfg.Providers.OpenRouter.APIBase - } else { - sel.apiBase = "https://openrouter.ai/api/v1" - } - case (strings.Contains(lowerModel, "claude") || strings.HasPrefix(model, "anthropic/")) && - (cfg.Providers.Anthropic.APIKey != "" || cfg.Providers.Anthropic.AuthMethod != ""): - if cfg.Providers.Anthropic.AuthMethod == "oauth" || cfg.Providers.Anthropic.AuthMethod == "token" { - sel.apiBase = cfg.Providers.Anthropic.APIBase - if sel.apiBase == "" { - sel.apiBase = defaultAnthropicAPIBase - } - sel.providerType = providerTypeClaudeAuth - return sel, nil - } - sel.apiKey = cfg.Providers.Anthropic.APIKey - sel.apiBase = cfg.Providers.Anthropic.APIBase - sel.proxy = cfg.Providers.Anthropic.Proxy - if sel.apiBase == "" { - sel.apiBase = defaultAnthropicAPIBase - } - case (strings.Contains(lowerModel, "gpt") || strings.HasPrefix(model, "openai/")) && - (cfg.Providers.OpenAI.APIKey != "" || cfg.Providers.OpenAI.AuthMethod != ""): - sel.enableWebSearch = cfg.Providers.OpenAI.WebSearch - if cfg.Providers.OpenAI.AuthMethod == "codex-cli" { - sel.providerType = providerTypeCodexCLIToken - return sel, nil - } - if cfg.Providers.OpenAI.AuthMethod == "oauth" || cfg.Providers.OpenAI.AuthMethod == "token" { - sel.providerType = providerTypeCodexAuth - return sel, nil - } - sel.apiKey = cfg.Providers.OpenAI.APIKey - sel.apiBase = cfg.Providers.OpenAI.APIBase - sel.proxy = cfg.Providers.OpenAI.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.openai.com/v1" - } - case (strings.Contains(lowerModel, "gemini") || strings.HasPrefix(model, "google/")) && cfg.Providers.Gemini.APIKey != "": - sel.apiKey = cfg.Providers.Gemini.APIKey - sel.apiBase = cfg.Providers.Gemini.APIBase - sel.proxy = cfg.Providers.Gemini.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://generativelanguage.googleapis.com/v1beta" - } - case (strings.Contains(lowerModel, "glm") || strings.Contains(lowerModel, "zhipu") || strings.Contains(lowerModel, "zai")) && cfg.Providers.Zhipu.APIKey != "": - sel.apiKey = cfg.Providers.Zhipu.APIKey - sel.apiBase = cfg.Providers.Zhipu.APIBase - sel.proxy = cfg.Providers.Zhipu.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://open.bigmodel.cn/api/paas/v4" - } - case (strings.Contains(lowerModel, "groq") || strings.HasPrefix(model, "groq/")) && cfg.Providers.Groq.APIKey != "": - sel.apiKey = cfg.Providers.Groq.APIKey - sel.apiBase = cfg.Providers.Groq.APIBase - sel.proxy = cfg.Providers.Groq.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.groq.com/openai/v1" - } - case (strings.Contains(lowerModel, "nvidia") || strings.HasPrefix(model, "nvidia/")) && cfg.Providers.Nvidia.APIKey != "": - sel.apiKey = cfg.Providers.Nvidia.APIKey - sel.apiBase = cfg.Providers.Nvidia.APIBase - sel.proxy = cfg.Providers.Nvidia.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://integrate.api.nvidia.com/v1" - } - case (strings.Contains(lowerModel, "ollama") || strings.HasPrefix(model, "ollama/")) && cfg.Providers.Ollama.APIKey != "": - sel.apiKey = cfg.Providers.Ollama.APIKey - sel.apiBase = cfg.Providers.Ollama.APIBase - sel.proxy = cfg.Providers.Ollama.Proxy - if sel.apiBase == "" { - sel.apiBase = "http://localhost:11434/v1" - } - case (strings.Contains(lowerModel, "mistral") || strings.HasPrefix(model, "mistral/")) && cfg.Providers.Mistral.APIKey != "": - sel.apiKey = cfg.Providers.Mistral.APIKey - sel.apiBase = cfg.Providers.Mistral.APIBase - sel.proxy = cfg.Providers.Mistral.Proxy - if sel.apiBase == "" { - sel.apiBase = "https://api.mistral.ai/v1" - } - case cfg.Providers.VLLM.APIBase != "": - sel.apiKey = cfg.Providers.VLLM.APIKey - sel.apiBase = cfg.Providers.VLLM.APIBase - sel.proxy = cfg.Providers.VLLM.Proxy - default: + } + + // Ultimate fallback: try OpenRouter if configured. + if !matched { if cfg.Providers.OpenRouter.APIKey != "" { - sel.apiKey = cfg.Providers.OpenRouter.APIKey - sel.proxy = cfg.Providers.OpenRouter.Proxy - if cfg.Providers.OpenRouter.APIBase != "" { - sel.apiBase = cfg.Providers.OpenRouter.APIBase - } else { - sel.apiBase = "https://openrouter.ai/api/v1" - } + applyStandardProvider(cfg, &sel, standardProviderRegistry["openrouter"]) } else { return providerSelection{}, fmt.Errorf("no API key configured for model: %s", model) } @@ -321,3 +420,4 @@ func resolveProviderSelection(cfg *config.Config) (providerSelection, error) { return sel, nil } + diff --git a/pkg/providers/types.go b/pkg/providers/types.go index f711e7803..9a4ee3533 100644 --- a/pkg/providers/types.go +++ b/pkg/providers/types.go @@ -34,13 +34,14 @@ type LLMProvider interface { type FailoverReason string const ( - FailoverAuth FailoverReason = "auth" - FailoverRateLimit FailoverReason = "rate_limit" - FailoverBilling FailoverReason = "billing" - FailoverTimeout FailoverReason = "timeout" - FailoverFormat FailoverReason = "format" - FailoverOverloaded FailoverReason = "overloaded" - FailoverUnknown FailoverReason = "unknown" + FailoverAuth FailoverReason = "auth" + FailoverRateLimit FailoverReason = "rate_limit" + FailoverBilling FailoverReason = "billing" + FailoverTimeout FailoverReason = "timeout" + FailoverFormat FailoverReason = "format" + FailoverOverloaded FailoverReason = "overloaded" + FailoverContextWindow FailoverReason = "context_window" + FailoverUnknown FailoverReason = "unknown" ) // FailoverError wraps an LLM provider error with classification metadata.