refactor: decompose AgentLoop and improve code quality
- Extract ContextCompressor (compression, summarization, token estimation) from AgentLoop - Extract ToolExecutor (tool dispatch, result routing, context updates) from AgentLoop - Create named constants replacing magic numbers (retry limits, thresholds, timeouts) - Replace fragile string-matching error detection with structured IsContextWindowError() - Add contextWindowPatterns and FailoverContextWindow to error classifier - Refactor provider factory to table-driven registry approach - Add graceful shutdown with context cancellation and WaitGroup - Net reduction of ~143 lines from loop.go via component extraction
This commit is contained in:
parent
7cbfa89a96
commit
05f80c5f3d
7 changed files with 752 additions and 504 deletions
43
pkg/agent/constants.go
Normal file
43
pkg/agent/constants.go
Normal file
|
|
@ -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
|
||||
)
|
||||
237
pkg/agent/context_compressor.go
Normal file
237
pkg/agent/context_compressor.go
Normal file
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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, "/") {
|
||||
|
|
|
|||
111
pkg/agent/tool_executor.go
Normal file
111
pkg/agent/tool_executor.go
Normal file
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue