From 1772a1ccbd208927085a49a84976821e614ed935 Mon Sep 17 00:00:00 2001 From: dj-oyu <68707227+dj-oyu@users.noreply.github.com> Date: Fri, 13 Mar 2026 13:14:57 +0900 Subject: [PATCH] fix: resolve post-merge type mismatches and duplicate declarations Fix FunctionCall.Arguments (map[string]any) and Parameters (json.RawMessage) type usage in new upstream providers (antigravity, codex). Remove duplicate FallbackCandidate/StreamEvent from types.go. Add cloneToolArgs, fix session Close signature, remove duplicate functions from loop.go, and fix test type assertions. Co-Authored-By: Claude Opus 4.6 --- cmd/picoclaw/internal/gateway/command.go | 6 +- cmd/picoclaw/internal/helpers_ext_test.go | 51 +-- pkg/agent/instance.go | 9 + pkg/agent/loop.go | 380 --------------------- pkg/agent/loop_ext_test.go | 88 +++-- pkg/memory/jsonl_test.go | 2 +- pkg/memory/migration_test.go | 2 +- pkg/providers/anthropic/provider.go | 11 +- pkg/providers/anthropic/provider_test.go | 8 +- pkg/providers/antigravity_provider.go | 14 +- pkg/providers/antigravity_provider_test.go | 2 +- pkg/providers/claude_cli_provider_test.go | 12 +- pkg/providers/codex_cli_provider_test.go | 11 +- pkg/providers/codex_provider.go | 10 +- pkg/providers/codex_provider_test.go | 17 +- pkg/providers/factory_provider.go | 87 +++++ pkg/providers/tool_call_extract.go | 13 + pkg/providers/toolcall_utils.go | 14 +- pkg/providers/types.go | 11 +- pkg/session/jsonl_backend_test.go | 6 +- pkg/session/legacy_adapter.go | 6 +- pkg/session/legacy_adapter_test.go | 2 +- pkg/session/session_store.go | 4 +- pkg/tools/filesystem.go | 13 + pkg/tools/registry.go | 22 +- pkg/tools/registry_ext_test.go | 2 +- pkg/tools/registry_test.go | 4 +- pkg/tools/spawn_test.go | 4 +- pkg/tools/subagent_tool_test.go | 18 +- 29 files changed, 302 insertions(+), 527 deletions(-) diff --git a/cmd/picoclaw/internal/gateway/command.go b/cmd/picoclaw/internal/gateway/command.go index bfa69f072..5e42a677f 100644 --- a/cmd/picoclaw/internal/gateway/command.go +++ b/cmd/picoclaw/internal/gateway/command.go @@ -12,6 +12,8 @@ import ( func NewGatewayCommand() *cobra.Command { var debug bool var noTruncate bool + var orchestration bool + var enableStats bool cmd := &cobra.Command{ Use: "gateway", @@ -31,12 +33,14 @@ func NewGatewayCommand() *cobra.Command { return nil }, RunE: func(_ *cobra.Command, _ []string) error { - return gatewayCmd(debug) + return gatewayCmd(debug, orchestration, enableStats) }, } cmd.Flags().BoolVarP(&debug, "debug", "d", false, "Enable debug logging") cmd.Flags().BoolVarP(&noTruncate, "no-truncate", "T", false, "Disable string truncation in debug logs") + cmd.Flags().BoolVar(&orchestration, "orchestration", false, "Enable subagent orchestration") + cmd.Flags().BoolVar(&enableStats, "stats", false, "Enable stats collection") return cmd } diff --git a/cmd/picoclaw/internal/helpers_ext_test.go b/cmd/picoclaw/internal/helpers_ext_test.go index 7e9ca5097..fafa63be8 100644 --- a/cmd/picoclaw/internal/helpers_ext_test.go +++ b/cmd/picoclaw/internal/helpers_ext_test.go @@ -1,63 +1,66 @@ package internal import ( - "github.com/stretchr/testify/assert" "runtime" "testing" + + "github.com/stretchr/testify/assert" + + "github.com/sipeed/picoclaw/pkg/config" ) func TestFormatVersion_NoGitCommit(t *testing.T) { - oldVersion, oldGit := version, gitCommit - t.Cleanup(func() { version, gitCommit = oldVersion, oldGit }) + oldVersion, oldGit := config.Version, config.GitCommit + t.Cleanup(func() { config.Version, config.GitCommit = oldVersion, oldGit }) - version = "1.2.3" - gitCommit = "" + config.Version = "1.2.3" + config.GitCommit = "" assert.Equal(t, "1.2.3", FormatVersion()) } func TestFormatVersion_WithGitCommit(t *testing.T) { - oldVersion, oldGit := version, gitCommit - t.Cleanup(func() { version, gitCommit = oldVersion, oldGit }) + oldVersion, oldGit := config.Version, config.GitCommit + t.Cleanup(func() { config.Version, config.GitCommit = oldVersion, oldGit }) - version = "1.2.3" - gitCommit = "abc123" + config.Version = "1.2.3" + config.GitCommit = "abc123" assert.Equal(t, "1.2.3 (git: abc123)", FormatVersion()) } func TestFormatBuildInfo_UsesBuildTimeAndGoVersion_WhenSet(t *testing.T) { - oldBuildTime, oldGoVersion := buildTime, goVersion - t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion }) + oldBuildTime, oldGoVersion := config.BuildTime, config.GoVersion + t.Cleanup(func() { config.BuildTime, config.GoVersion = oldBuildTime, oldGoVersion }) - buildTime = "2026-02-20T00:00:00Z" - goVersion = "go1.23.0" + config.BuildTime = "2026-02-20T00:00:00Z" + config.GoVersion = "go1.23.0" build, goVer := FormatBuildInfo() - assert.Equal(t, buildTime, build) - assert.Equal(t, goVersion, goVer) + assert.Equal(t, config.BuildTime, build) + assert.Equal(t, config.GoVersion, goVer) } func TestFormatBuildInfo_EmptyBuildTime_ReturnsEmptyBuild(t *testing.T) { - oldBuildTime, oldGoVersion := buildTime, goVersion - t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion }) + oldBuildTime, oldGoVersion := config.BuildTime, config.GoVersion + t.Cleanup(func() { config.BuildTime, config.GoVersion = oldBuildTime, oldGoVersion }) - buildTime = "" - goVersion = "go1.23.0" + config.BuildTime = "" + config.GoVersion = "go1.23.0" build, goVer := FormatBuildInfo() assert.Empty(t, build) - assert.Equal(t, goVersion, goVer) + assert.Equal(t, config.GoVersion, goVer) } func TestFormatBuildInfo_EmptyGoVersion_FallsBackToRuntimeVersion(t *testing.T) { - oldBuildTime, oldGoVersion := buildTime, goVersion - t.Cleanup(func() { buildTime, goVersion = oldBuildTime, oldGoVersion }) + oldBuildTime, oldGoVersion := config.BuildTime, config.GoVersion + t.Cleanup(func() { config.BuildTime, config.GoVersion = oldBuildTime, oldGoVersion }) - buildTime = "x" - goVersion = "" + config.BuildTime = "x" + config.GoVersion = "" build, goVer := FormatBuildInfo() diff --git a/pkg/agent/instance.go b/pkg/agent/instance.go index 684396427..809c89d49 100644 --- a/pkg/agent/instance.go +++ b/pkg/agent/instance.go @@ -65,6 +65,15 @@ type AgentInstance struct { worktreeMu sync.RWMutex } +// Close releases resources held by the agent instance. +// If the provider implements StatefulProvider, its Close method is called. +func (ai *AgentInstance) Close() error { + if sp, ok := ai.Provider.(providers.StatefulProvider); ok { + sp.Close() + } + return nil +} + // NewAgentInstance creates an agent instance from config. func NewAgentInstance( agentCfg *config.AgentConfig, diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index f6abfb472..9c61a134b 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -17,7 +17,6 @@ import ( "sync" "sync/atomic" "time" - "unicode/utf8" "github.com/sipeed/picoclaw/pkg/bus" "github.com/sipeed/picoclaw/pkg/channels" @@ -1998,52 +1997,6 @@ func (al *AgentLoop) targetReasoningChannelID(channelName string) (chatID string return "" } -func (al *AgentLoop) handleReasoning( - ctx context.Context, - reasoningContent, channelName, channelID string, -) { - if reasoningContent == "" || channelName == "" || channelID == "" { - return - } - - // Check context cancellation before attempting to publish, - // since PublishOutbound's select may race between send and ctx.Done(). - if ctx.Err() != nil { - return - } - - // Use a short timeout so the goroutine does not block indefinitely when - // the outbound bus is full. Reasoning output is best-effort; dropping it - // is acceptable to avoid goroutine accumulation. - pubCtx, pubCancel := context.WithTimeout(ctx, 5*time.Second) - defer pubCancel() - - if err := al.bus.PublishOutbound(pubCtx, bus.OutboundMessage{ - Channel: channelName, - ChatID: channelID, - Content: reasoningContent, - }); err != nil { - // Treat context.DeadlineExceeded / context.Canceled as expected - // (bus full under load, or parent canceled). Check the error - // itself rather than ctx.Err(), because pubCtx may time out - // (5 s) while the parent ctx is still active. - // Also treat ErrBusClosed as expected — it occurs during normal - // shutdown when the bus is closed before all goroutines finish. - if errors.Is(err, context.DeadlineExceeded) || errors.Is(err, context.Canceled) || - errors.Is(err, bus.ErrBusClosed) { - logger.DebugCF("agent", "Reasoning publish skipped (timeout/cancel)", map[string]any{ - "channel": channelName, - "error": err.Error(), - }) - } else { - logger.WarnCF("agent", "Failed to publish reasoning (best-effort)", map[string]any{ - "channel": channelName, - "error": err.Error(), - }) - } - } -} - // runLLMIteration executes the LLM call loop with tool handling using hooks. func (al *AgentLoop) runLLMIteration( ctx context.Context, @@ -2620,237 +2573,6 @@ func (al *AgentLoop) selectCandidates( return agent.LightCandidates, agent.Router.LightModel() } -// 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 * agent.SummarizeTokenPercent / 100 - - if len(newHistory) > agent.SummarizeMessageThreshold || tokenEstimate > threshold { - summarizeKey := agent.ID + ":" + sessionKey - if _, loading := al.summarizing.LoadOrStore(summarizeKey, true); !loading { - go func() { - defer al.summarizing.Delete(summarizeKey) - logger.Debug("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, 1+len(keptConversation)+1) - - // 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) - - agent := al.registry.GetDefaultAgent() - if agent == nil { - return info - } - - // Tools info - toolsList := agent.Tools.List() - info["tools"] = map[string]any{ - "count": len(toolsList), - "names": toolsList, - } - - // Skills info - info["skills"] = agent.ContextBuilder.GetSkillsInfo() - - // Agents info - info["agents"] = map[string]any{ - "count": len(al.registry.ListAgentIDs()), - "ids": al.registry.ListAgentIDs(), - } - - return info -} - -// formatMessagesForLog formats messages for logging -func formatMessagesForLog(messages []providers.Message) string { - if len(messages) == 0 { - return "[]" - } - - var sb strings.Builder - sb.WriteString("[\n") - for i, msg := range messages { - fmt.Fprintf(&sb, " [%d] Role: %s\n", i, msg.Role) - if len(msg.ToolCalls) > 0 { - sb.WriteString(" ToolCalls:\n") - for _, tc := range msg.ToolCalls { - fmt.Fprintf(&sb, " - ID: %s, Type: %s, Name: %s\n", tc.ID, tc.Type, tc.Name) - if tc.Function != nil { - fmt.Fprintf( - &sb, - " Arguments: %s\n", - utils.Truncate(tc.Function.Arguments, 200), - ) - } - } - } - if msg.Content != "" { - content := utils.Truncate(msg.Content, 200) - fmt.Fprintf(&sb, " Content: %s\n", content) - } - if msg.ToolCallID != "" { - fmt.Fprintf(&sb, " ToolCallID: %s\n", msg.ToolCallID) - } - sb.WriteString("\n") - } - sb.WriteString("]") - return sb.String() -} - -// formatToolsForLog formats tool definitions for logging -func formatToolsForLog(toolDefs []providers.ToolDefinition) string { - if len(toolDefs) == 0 { - return "[]" - } - - var sb strings.Builder - sb.WriteString("[\n") - for i, tool := range toolDefs { - fmt.Fprintf(&sb, " [%d] Type: %s, Name: %s\n", i, tool.Type, tool.Function.Name) - fmt.Fprintf(&sb, " Description: %s\n", tool.Function.Description) - if len(tool.Function.Parameters) > 0 { - fmt.Fprintf( - &sb, - " Parameters: %s\n", - utils.Truncate(fmt.Sprintf("%v", tool.Function.Parameters), 200), - ) - } - } - sb.WriteString("]") - 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 - } - - // Oversized Message Guard - maxMessageTokens := agent.ContextWindow / 2 - validMessages := make([]providers.Message, 0) - omitted := false - - for _, m := range history[:len(history)-4] { - msgTokens := len(m.Content) / 2 - if msgTokens > maxMessageTokens { - omitted = true - continue - } - validMessages = append(validMessages, m) - } - - if omitted { - logger.WarnCF("agent", "Oversized messages omitted during summarization", - map[string]any{"session_key": sessionKey}) - } - - const ( - maxSummarizationMessages = 10 - llmMaxRetries = 3 - ) - - // Multi-Part Summarization - var finalSummary string - if len(validMessages) > maxSummarizationMessages { - mid := len(validMessages) / 2 - - mid = al.findNearestUserMessage(validMessages, mid) - - 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 := al.retryLLMCall(ctx, agent, mergePrompt, llmMaxRetries) - if err == nil && resp.Content != "" { - finalSummary = resp.Content - } else { - finalSummary = s1 + " " + s2 - } - } else { - finalSummary, _ = al.summarizeBatch(ctx, agent, validMessages, summary) - } - - if finalSummary != "" { - agent.Sessions.SetSummary(sessionKey, finalSummary) - agent.Sessions.TruncateHistory(sessionKey, 4) - agent.Sessions.Save(sessionKey) - } -} - // findNearestUserMessage finds the nearest user message to the given index. // It searches backward first, then forward if no user message is found. func (al *AgentLoop) findNearestUserMessage(messages []providers.Message, mid int) int { @@ -2913,82 +2635,6 @@ func (al *AgentLoop) retryLLMCall( return resp, err } -// summarizeBatch summarizes a batch of messages. -func (al *AgentLoop) summarizeBatch( - ctx context.Context, - agent *AgentInstance, - batch []providers.Message, - existingSummary string, -) (string, error) { - const ( - llmMaxRetries = 3 - fallbackMinContentLength = 200 - fallbackMaxContentPercent = 10 - ) - - 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 := al.retryLLMCall(ctx, agent, prompt, llmMaxRetries) - if err == nil && response.Content != "" { - return strings.TrimSpace(response.Content), nil - } - - var fallback strings.Builder - fallback.WriteString("Conversation summary: ") - for i, m := range batch { - if i > 0 { - fallback.WriteString(" | ") - } - content := strings.TrimSpace(m.Content) - runes := []rune(content) - if len(runes) == 0 { - fallback.WriteString(fmt.Sprintf("%s: ", m.Role)) - continue - } - - keepLength := len(runes) * fallbackMaxContentPercent / 100 - if keepLength < fallbackMinContentLength { - keepLength = fallbackMinContentLength - } - - if keepLength > len(runes) { - keepLength = len(runes) - } - - content = string(runes[:keepLength]) - if keepLength < len(runes) { - content += "..." - } - fallback.WriteString(fmt.Sprintf("%s: %s", m.Role, content)) - } - return fallback.String(), 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 -} - // updateToolContexts updates the context for tools that need channel/chatID info. func (al *AgentLoop) updateToolContexts(agent *AgentInstance, channel, chatID string) { @@ -3013,35 +2659,9 @@ func (al *AgentLoop) updateToolContexts(agent *AgentInstance, channel, chatID st } } -// extractPeer extracts the routing peer from the inbound message's structured Peer field. -func extractPeer(msg bus.InboundMessage) *routing.RoutePeer { - if msg.Peer.Kind == "" { - return nil - } - peerID := msg.Peer.ID - if peerID == "" { - if msg.Peer.Kind == "direct" { - peerID = msg.SenderID - } else { - peerID = msg.ChatID - } - } - return &routing.RoutePeer{Kind: msg.Peer.Kind, ID: peerID} -} - func inboundMetadata(msg bus.InboundMessage, key string) string { if msg.Metadata == nil { return "" } return msg.Metadata[key] } - -// extractParentPeer extracts the parent peer (reply-to) from inbound message metadata. -func extractParentPeer(msg bus.InboundMessage) *routing.RoutePeer { - parentKind := inboundMetadata(msg, metadataKeyParentPeerKind) - parentID := inboundMetadata(msg, metadataKeyParentPeerID) - if parentKind == "" || parentID == "" { - return nil - } - return &routing.RoutePeer{Kind: parentKind, ID: parentID} -} diff --git a/pkg/agent/loop_ext_test.go b/pkg/agent/loop_ext_test.go index a26a3fa24..a16682711 100644 --- a/pkg/agent/loop_ext_test.go +++ b/pkg/agent/loop_ext_test.go @@ -407,7 +407,7 @@ func TestBuildPlanReminder(t *testing.T) { } func TestPlanCommand_ShowNoPlan(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -466,7 +466,7 @@ func TestSplitChatAndThread(t *testing.T) { } func TestHeartbeatCommandThreadHerePersistsConfig(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -522,7 +522,7 @@ func TestHeartbeatCommandThreadHerePersistsConfig(t *testing.T) { } func TestHeartbeatCommandThreadOff(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -550,7 +550,7 @@ func TestHeartbeatCommandThreadOff(t *testing.T) { } func TestPlanCommand_StartNewPlan(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -588,7 +588,7 @@ func TestPlanCommand_StartNewPlan(t *testing.T) { } func TestPlanCommand_StartBlockedByExisting(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -606,7 +606,7 @@ func TestPlanCommand_StartBlockedByExisting(t *testing.T) { } func TestPlanCommand_Clear(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -626,7 +626,7 @@ func TestPlanCommand_Clear(t *testing.T) { } func TestPlanCommand_ClearNoPlan(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -638,7 +638,7 @@ func TestPlanCommand_ClearNoPlan(t *testing.T) { } func TestPlanCommand_Start(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -664,7 +664,7 @@ func TestPlanCommand_Start(t *testing.T) { } func TestPlanCommand_StartFromReview(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -690,7 +690,7 @@ func TestPlanCommand_StartFromReview(t *testing.T) { } func TestPlanCommand_StartNoPhases(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -714,7 +714,7 @@ func TestPlanCommand_StartNoPhases(t *testing.T) { } func TestPlanCommand_StartAlreadyExecuting(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -740,7 +740,7 @@ func TestPlanCommand_StartAlreadyExecuting(t *testing.T) { } func TestPlanCommand_Done(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -782,7 +782,7 @@ Test context } func TestPlanCommand_DoneInvalidStep(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -796,7 +796,7 @@ func TestPlanCommand_DoneInvalidStep(t *testing.T) { } func TestPlanCommand_Add(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -842,7 +842,7 @@ Test context } func TestPlanCommand_Next(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -892,7 +892,7 @@ Test } func TestPlanCommand_ShowActivePlan(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -1864,8 +1864,8 @@ func TestPlanNudge_ForegroundExecution(t *testing.T) { t.Fatalf("processMessage failed: %v", err) } - if provider.callCount < 2 { - t.Errorf("expected at least 2 provider calls (nudge should trigger continuation), got %d", provider.callCount) + if provider.calls < 2 { + t.Errorf("expected at least 2 provider calls (nudge should trigger continuation), got %d", provider.calls) } } @@ -1928,8 +1928,8 @@ func TestPlanNudge_NoNudgeWhenAllStepsComplete(t *testing.T) { t.Fatalf("processMessage failed: %v", err) } - if provider.callCount != 1 { - t.Errorf("expected exactly 1 provider call (no nudge needed), got %d", provider.callCount) + if provider.calls != 1 { + t.Errorf("expected exactly 1 provider call (no nudge needed), got %d", provider.calls) } } @@ -2019,13 +2019,34 @@ func TestPlanNudge_ProgressMessage(t *testing.T) { } type nudgeCaptureMockProvider struct { - callCount int + calls int onFirstCall func() onSecondCall func([]providers.Message) } +func (m *nudgeCaptureMockProvider) Chat( + _ context.Context, + messages []providers.Message, + _ []providers.ToolDefinition, + _ string, + _ map[string]any, +) (*providers.LLMResponse, error) { + m.calls++ + if m.calls == 1 && m.onFirstCall != nil { + m.onFirstCall() + } + if m.calls == 2 && m.onSecondCall != nil { + m.onSecondCall(messages) + } + return &providers.LLMResponse{Content: "ok"}, nil +} + +func (m *nudgeCaptureMockProvider) GetDefaultModel() string { + return "nudge-mock" +} + func TestConsumeStream_NormalCompletion(t *testing.T) { ch := make(chan protocoltypes.StreamEvent, 8) @@ -2308,6 +2329,27 @@ type modelCapturingMockProvider struct { response string } +func (m *modelCapturingMockProvider) Chat( + _ context.Context, + _ []providers.Message, + _ []providers.ToolDefinition, + model string, + _ map[string]any, +) (*providers.LLMResponse, error) { + m.mu.Lock() + m.models = append(m.models, model) + m.mu.Unlock() + resp := m.response + if resp == "" { + resp = "ok" + } + return &providers.LLMResponse{Content: resp}, nil +} + +func (m *modelCapturingMockProvider) GetDefaultModel() string { + return "model-capturing-mock" +} + func TestAgentLoop_PlanModel_UsedDuringInterviewing(t *testing.T) { tmpDir, err := os.MkdirTemp("", "agent-test-planmodel-*") if err != nil { @@ -2564,7 +2606,7 @@ func TestAgentLoop_PlanModel_ResolvesProviderForSingleCandidate(t *testing.T) { } func TestPlanCommand_StartClear(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() @@ -2630,7 +2672,7 @@ func TestPlanCommand_StartClear(t *testing.T) { } func TestPlanCommand_StartWithoutClear_PreservesHistory(t *testing.T) { - al, cleanup := newTestAgentLoop(t) + al, _, _, _, cleanup := newTestAgentLoop(t) defer cleanup() diff --git a/pkg/memory/jsonl_test.go b/pkg/memory/jsonl_test.go index 356ff14ff..31af5f320 100644 --- a/pkg/memory/jsonl_test.go +++ b/pkg/memory/jsonl_test.go @@ -96,7 +96,7 @@ func TestAddFullMessage_WithToolCalls(t *testing.T) { Type: "function", Function: &providers.FunctionCall{ Name: "web_search", - Arguments: `{"q":"golang jsonl"}`, + Arguments: map[string]any{"q": "golang jsonl"}, }, }, }, diff --git a/pkg/memory/migration_test.go b/pkg/memory/migration_test.go index 4466c96f9..cc6d05da2 100644 --- a/pkg/memory/migration_test.go +++ b/pkg/memory/migration_test.go @@ -86,7 +86,7 @@ func TestMigrateFromJSON_WithToolCalls(t *testing.T) { Type: "function", Function: &providers.FunctionCall{ Name: "web_search", - Arguments: `{"q":"test"}`, + Arguments: map[string]any{"q": "test"}, }, }, }, diff --git a/pkg/providers/anthropic/provider.go b/pkg/providers/anthropic/provider.go index 242ded175..17d70314a 100644 --- a/pkg/providers/anthropic/provider.go +++ b/pkg/providers/anthropic/provider.go @@ -181,10 +181,8 @@ func buildParams( } for _, tc := range msg.ToolCalls { args := tc.Arguments - if args == nil && tc.Function != nil && tc.Function.Arguments != "" { - if err := json.Unmarshal([]byte(tc.Function.Arguments), &args); err != nil { - args = map[string]any{} - } + if args == nil && tc.Function != nil && len(tc.Function.Arguments) > 0 { + args = tc.Function.Arguments } if args == nil { args = map[string]any{} @@ -308,16 +306,17 @@ func levelToBudget(level string) int { func translateTools(tools []ToolDefinition) []anthropic.ToolUnionParam { result := make([]anthropic.ToolUnionParam, 0, len(tools)) for _, t := range tools { + params := t.Function.ParametersMap() tool := anthropic.ToolParam{ Name: t.Function.Name, InputSchema: anthropic.ToolInputSchemaParam{ - Properties: t.Function.Parameters["properties"], + Properties: params["properties"], }, } if desc := t.Function.Description; desc != "" { tool.Description = anthropic.String(desc) } - if req, ok := t.Function.Parameters["required"].([]any); ok { + if req, ok := params["required"].([]any); ok { required := make([]string, 0, len(req)) for _, r := range req { if s, ok := r.(string); ok { diff --git a/pkg/providers/anthropic/provider_test.go b/pkg/providers/anthropic/provider_test.go index b1aed17b5..c37b7d52d 100644 --- a/pkg/providers/anthropic/provider_test.go +++ b/pkg/providers/anthropic/provider_test.go @@ -84,13 +84,7 @@ func TestBuildParams_WithTools(t *testing.T) { Function: ToolFunctionDefinition{ Name: "get_weather", Description: "Get weather for a city", - Parameters: map[string]any{ - "type": "object", - "properties": map[string]any{ - "city": map[string]any{"type": "string"}, - }, - "required": []any{"city"}, - }, + Parameters: json.RawMessage(`{"type":"object","properties":{"city":{"type":"string"}},"required":["city"]}`), }, }, } diff --git a/pkg/providers/antigravity_provider.go b/pkg/providers/antigravity_provider.go index 8a1890212..d49873570 100644 --- a/pkg/providers/antigravity_provider.go +++ b/pkg/providers/antigravity_provider.go @@ -297,7 +297,7 @@ func (p *AntigravityProvider) buildRequest( if t.Type != "function" { continue } - params := sanitizeSchemaForGemini(t.Function.Parameters) + params := sanitizeSchemaForGemini(t.Function.ParametersMap()) funcDecls = append(funcDecls, antigravityFuncDecl{ Name: t.Function.Name, Description: t.Function.Description, @@ -344,11 +344,8 @@ func normalizeStoredToolCall(tc ToolCall) (string, map[string]any, string) { args = map[string]any{} } - if len(args) == 0 && tc.Function != nil && tc.Function.Arguments != "" { - var parsed map[string]any - if err := json.Unmarshal([]byte(tc.Function.Arguments), &parsed); err == nil && parsed != nil { - args = parsed - } + if len(args) == 0 && tc.Function != nil && len(tc.Function.Arguments) > 0 { + args = tc.Function.Arguments } return name, args, thoughtSignature @@ -436,14 +433,13 @@ func (p *AntigravityProvider) parseSSEResponse(body string) (*LLMResponse, error contentParts = append(contentParts, part.Text) } if part.FunctionCall != nil { - argumentsJSON, _ := json.Marshal(part.FunctionCall.Args) - toolCalls = append(toolCalls, ToolCall{ + toolCalls = append(toolCalls, ToolCall{ ID: fmt.Sprintf("call_%s_%d", part.FunctionCall.Name, time.Now().UnixNano()), Name: part.FunctionCall.Name, Arguments: part.FunctionCall.Args, Function: &FunctionCall{ Name: part.FunctionCall.Name, - Arguments: string(argumentsJSON), + Arguments: cloneToolArgs(part.FunctionCall.Args), ThoughtSignature: extractPartThoughtSignature( part.ThoughtSignature, part.ThoughtSignatureSnake, diff --git a/pkg/providers/antigravity_provider_test.go b/pkg/providers/antigravity_provider_test.go index 238765321..a7eebbedc 100644 --- a/pkg/providers/antigravity_provider_test.go +++ b/pkg/providers/antigravity_provider_test.go @@ -12,7 +12,7 @@ func TestBuildRequestUsesFunctionFieldsWhenToolCallNameMissing(t *testing.T) { ID: "call_read_file_123", Function: &FunctionCall{ Name: "read_file", - Arguments: `{"path":"README.md"}`, + Arguments: map[string]any{"path": "README.md"}, }, }}, }, diff --git a/pkg/providers/claude_cli_provider_test.go b/pkg/providers/claude_cli_provider_test.go index d4d648f5a..b1045c904 100644 --- a/pkg/providers/claude_cli_provider_test.go +++ b/pkg/providers/claude_cli_provider_test.go @@ -2,6 +2,7 @@ package providers import ( "context" + "encoding/json" "fmt" "os" "path/filepath" @@ -619,12 +620,7 @@ func TestBuildSystemPrompt_WithTools(t *testing.T) { Function: ToolFunctionDefinition{ Name: "get_weather", Description: "Get weather for a location", - Parameters: map[string]any{ - "type": "object", - "properties": map[string]any{ - "location": map[string]any{"type": "string"}, - }, - }, + Parameters: json.RawMessage(`{"type":"object","properties":{"location":{"type":"string"}}}`), }, }, } @@ -918,8 +914,8 @@ func TestExtractToolCalls_ToolCallArgumentsParsing(t *testing.T) { t.Errorf("Arguments[name] = %v, want test", got[0].Arguments["name"]) } // Verify raw arguments string is preserved in FunctionCall - if got[0].Function.Arguments == "" { - t.Error("Function.Arguments should contain raw JSON string") + if len(got[0].Function.Arguments) == 0 { + t.Error("Function.Arguments should contain parsed arguments") } } diff --git a/pkg/providers/codex_cli_provider_test.go b/pkg/providers/codex_cli_provider_test.go index 0f66e25f4..f088d9dbd 100644 --- a/pkg/providers/codex_cli_provider_test.go +++ b/pkg/providers/codex_cli_provider_test.go @@ -76,8 +76,8 @@ func TestParseJSONLEvents_ToolCallExtraction(t *testing.T) { if resp.ToolCalls[0].ID != "call_1" { t.Errorf("ToolCalls[0].ID = %q, want %q", resp.ToolCalls[0].ID, "call_1") } - if resp.ToolCalls[0].Function.Arguments != `{"path":"/tmp/test.txt"}` { - t.Errorf("ToolCalls[0].Function.Arguments = %q", resp.ToolCalls[0].Function.Arguments) + if resp.ToolCalls[0].Function.Arguments["path"] != "/tmp/test.txt" { + t.Errorf("ToolCalls[0].Function.Arguments = %v", resp.ToolCalls[0].Function.Arguments) } // Content should have the tool call JSON stripped if strings.Contains(resp.Content, "tool_calls") { @@ -292,12 +292,7 @@ func TestBuildPrompt_WithTools(t *testing.T) { Function: ToolFunctionDefinition{ Name: "get_weather", Description: "Get current weather", - Parameters: map[string]any{ - "type": "object", - "properties": map[string]any{ - "city": map[string]any{"type": "string"}, - }, - }, + Parameters: json.RawMessage(`{"type":"object","properties":{"city":{"type":"string"}}}`), }, }, } diff --git a/pkg/providers/codex_provider.go b/pkg/providers/codex_provider.go index cf5c2d876..1270813d6 100644 --- a/pkg/providers/codex_provider.go +++ b/pkg/providers/codex_provider.go @@ -325,8 +325,12 @@ func resolveCodexToolCall(tc ToolCall) (name string, arguments string, ok bool) return name, string(argsJSON), true } - if tc.Function != nil && tc.Function.Arguments != "" { - return name, tc.Function.Arguments, true + if tc.Function != nil && len(tc.Function.Arguments) > 0 { + argsJSON, err := json.Marshal(tc.Function.Arguments) + if err != nil { + return "", "", false + } + return name, string(argsJSON), true } return name, "{}", true @@ -347,7 +351,7 @@ func translateToolsForCodex(tools []ToolDefinition, enableWebSearch bool) []resp } ft := responses.FunctionToolParam{ Name: t.Function.Name, - Parameters: t.Function.Parameters, + Parameters: t.Function.ParametersMap(), Strict: openai.Opt(false), } if t.Function.Description != "" { diff --git a/pkg/providers/codex_provider_test.go b/pkg/providers/codex_provider_test.go index dd5ad2637..45546b011 100644 --- a/pkg/providers/codex_provider_test.go +++ b/pkg/providers/codex_provider_test.go @@ -79,7 +79,7 @@ func TestBuildCodexParams_ToolCallFunctionFallback(t *testing.T) { Type: "function", Function: &FunctionCall{ Name: "read_file", - Arguments: `{"path":"README.md"}`, + Arguments: map[string]any{"path": "README.md"}, }, }, }, @@ -114,12 +114,7 @@ func TestBuildCodexParams_WithTools(t *testing.T) { Function: ToolFunctionDefinition{ Name: "get_weather", Description: "Get weather", - Parameters: map[string]any{ - "type": "object", - "properties": map[string]any{ - "city": map[string]any{"type": "string"}, - }, - }, + Parameters: json.RawMessage(`{"type":"object","properties":{"city":{"type":"string"}}}`), }, }, } @@ -166,9 +161,7 @@ func TestBuildCodexParams_WebSearchFunctionReplacedWithBuiltin(t *testing.T) { Function: ToolFunctionDefinition{ Name: "web_search", Description: "local web search", - Parameters: map[string]any{ - "type": "object", - }, + Parameters: json.RawMessage(`{"type":"object"}`), }, }, { @@ -176,9 +169,7 @@ func TestBuildCodexParams_WebSearchFunctionReplacedWithBuiltin(t *testing.T) { Function: ToolFunctionDefinition{ Name: "read_file", Description: "read file", - Parameters: map[string]any{ - "type": "object", - }, + Parameters: json.RawMessage(`{"type":"object"}`), }, }, } diff --git a/pkg/providers/factory_provider.go b/pkg/providers/factory_provider.go index 9749e7a15..217d8f1de 100644 --- a/pkg/providers/factory_provider.go +++ b/pkg/providers/factory_provider.go @@ -174,6 +174,93 @@ func CreateProviderFromConfig(cfg *config.ModelConfig) (LLMProvider, string, err } } +// CreateProviderByName creates a provider from the legacy ProvidersConfig by +// provider name (case-insensitive). It builds a ModelConfig from the named +// provider and delegates to CreateProviderFromConfig. +func CreateProviderByName(cfg *config.Config, name string) (LLMProvider, error) { + name = strings.ToLower(name) + p := cfg.Providers + + var pc config.ProviderConfig + protocol := name + + switch name { + case "openai", "gpt": + pc = config.ProviderConfig{ + APIKey: p.OpenAI.APIKey, + APIBase: p.OpenAI.APIBase, + Proxy: p.OpenAI.Proxy, + RequestTimeout: p.OpenAI.RequestTimeout, + AuthMethod: p.OpenAI.AuthMethod, + } + protocol = "openai" + case "anthropic", "claude": + pc = p.Anthropic + protocol = "anthropic" + case "litellm": + pc = p.LiteLLM + case "openrouter": + pc = p.OpenRouter + case "groq": + pc = p.Groq + case "zhipu": + pc = p.Zhipu + case "vllm": + pc = p.VLLM + case "gemini": + pc = p.Gemini + case "nvidia": + pc = p.Nvidia + case "ollama": + pc = p.Ollama + case "moonshot": + pc = p.Moonshot + case "shengsuanyun": + pc = p.ShengSuanYun + case "deepseek": + pc = p.DeepSeek + case "cerebras": + pc = p.Cerebras + case "vivgrid": + pc = p.Vivgrid + case "volcengine": + pc = p.VolcEngine + case "github-copilot", "copilot": + pc = p.GitHubCopilot + protocol = "github-copilot" + case "antigravity": + pc = p.Antigravity + case "qwen": + pc = p.Qwen + case "mistral": + pc = p.Mistral + case "avian": + pc = p.Avian + case "minimax": + pc = p.Minimax + case "longcat": + pc = p.LongCat + default: + return nil, fmt.Errorf("unknown provider %q", name) + } + + mc := &config.ModelConfig{ + ModelName: name, + Model: protocol + "/default", + APIKey: pc.APIKey, + APIBase: pc.APIBase, + Proxy: pc.Proxy, + RequestTimeout: pc.RequestTimeout, + AuthMethod: pc.AuthMethod, + } + + provider, _, err := CreateProviderFromConfig(mc) + if err != nil { + return nil, err + } + return provider, nil +} + // getDefaultAPIBase returns the default API base URL for a given protocol. func getDefaultAPIBase(protocol string) string { switch protocol { diff --git a/pkg/providers/tool_call_extract.go b/pkg/providers/tool_call_extract.go index cb04ff9cf..6fc81c042 100644 --- a/pkg/providers/tool_call_extract.go +++ b/pkg/providers/tool_call_extract.go @@ -72,3 +72,16 @@ func stripToolCallsFromText(text string) string { return strings.TrimSpace(text[:start] + text[end:]) } + +// cloneToolArgs returns a shallow copy of the given map so that +// ToolCall.Arguments and FunctionCall.Arguments do not alias. +func cloneToolArgs(m map[string]any) map[string]any { + if m == nil { + return nil + } + out := make(map[string]any, len(m)) + for k, v := range m { + out[k] = v + } + return out +} diff --git a/pkg/providers/toolcall_utils.go b/pkg/providers/toolcall_utils.go index a33e1eb5c..9daf18c5a 100644 --- a/pkg/providers/toolcall_utils.go +++ b/pkg/providers/toolcall_utils.go @@ -60,19 +60,15 @@ func NormalizeToolCall(tc ToolCall) ToolCall { } // Parse Arguments from Function.Arguments if not already set - if len(normalized.Arguments) == 0 && normalized.Function != nil && normalized.Function.Arguments != "" { - var parsed map[string]any - if err := json.Unmarshal([]byte(normalized.Function.Arguments), &parsed); err == nil && parsed != nil { - normalized.Arguments = parsed - } + if len(normalized.Arguments) == 0 && normalized.Function != nil && len(normalized.Function.Arguments) > 0 { + normalized.Arguments = cloneToolArgs(normalized.Function.Arguments) } // Ensure Function is populated with consistent values - argsJSON, _ := json.Marshal(normalized.Arguments) if normalized.Function == nil { normalized.Function = &FunctionCall{ Name: normalized.Name, - Arguments: string(argsJSON), + Arguments: cloneToolArgs(normalized.Arguments), } } else { if normalized.Function.Name == "" { @@ -81,8 +77,8 @@ func NormalizeToolCall(tc ToolCall) ToolCall { if normalized.Name == "" { normalized.Name = normalized.Function.Name } - if normalized.Function.Arguments == "" { - normalized.Function.Arguments = string(argsJSON) + if len(normalized.Function.Arguments) == 0 { + normalized.Function.Arguments = cloneToolArgs(normalized.Arguments) } } diff --git a/pkg/providers/types.go b/pkg/providers/types.go index 5e4b12dee..cbf6258a1 100644 --- a/pkg/providers/types.go +++ b/pkg/providers/types.go @@ -97,16 +97,7 @@ type ModelConfig struct { type StreamingProvider interface { LLMProvider CanStream() bool - ChatStream(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]any) (<-chan StreamEvent, error) -} - -// FallbackCandidate represents a model that can be tried if the primary model fails. -type FallbackCandidate struct { - ModelName string - Model string - Protocol string - Provider LLMProvider - Options map[string]any + ChatStream(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]any) (<-chan protocoltypes.StreamEvent, error) } // UnmarshalArguments is a helper to parse FunctionCall.Arguments from json.RawMessage. diff --git a/pkg/session/jsonl_backend_test.go b/pkg/session/jsonl_backend_test.go index 40fa019cb..f5fce98ed 100644 --- a/pkg/session/jsonl_backend_test.go +++ b/pkg/session/jsonl_backend_test.go @@ -11,8 +11,8 @@ import ( // Compile-time interface satisfaction checks. var ( - _ session.SessionStore = (*session.SessionManager)(nil) - _ session.SessionStore = (*session.JSONLBackend)(nil) + _ session.LegacyStore = (*session.SessionManager)(nil) + _ session.LegacyStore = (*session.JSONLBackend)(nil) ) func newBackend(t *testing.T) *session.JSONLBackend { @@ -50,7 +50,7 @@ func TestJSONLBackend_AddFullMessage(t *testing.T) { Role: "assistant", Content: "done", ToolCalls: []providers.ToolCall{ - {ID: "tc1", Function: &providers.FunctionCall{Name: "read_file", Arguments: `{"path":"x"}`}}, + {ID: "tc1", Function: &providers.FunctionCall{Name: "read_file", Arguments: map[string]any{"path": "x"}}}, }, } b.AddFullMessage("s1", msg) diff --git a/pkg/session/legacy_adapter.go b/pkg/session/legacy_adapter.go index 29e7b4719..d921dc8ba 100644 --- a/pkg/session/legacy_adapter.go +++ b/pkg/session/legacy_adapter.go @@ -598,11 +598,11 @@ func (la *LegacyAdapter) AdvanceStored(key string, delta int) { // Close stops the background flush loop and persists all dirty sessions. -func (la *LegacyAdapter) Close() { +func (la *LegacyAdapter) Close() error { select { case <-la.done: - return // already closed + return nil // already closed default: } @@ -611,7 +611,7 @@ func (la *LegacyAdapter) Close() { la.FlushDirty() - la.store.Close() + return la.store.Close() } func (la *LegacyAdapter) flushLoop() { diff --git a/pkg/session/legacy_adapter_test.go b/pkg/session/legacy_adapter_test.go index b313140cb..13f738225 100644 --- a/pkg/session/legacy_adapter_test.go +++ b/pkg/session/legacy_adapter_test.go @@ -33,7 +33,7 @@ type sessionBackend interface { //nolint:interfacebloat // test helper mirrors S Save(key string) error - Close() + Close() error } func backends(t *testing.T) map[string]sessionBackend { diff --git a/pkg/session/session_store.go b/pkg/session/session_store.go index 1d1a2f967..3b22ac85c 100644 --- a/pkg/session/session_store.go +++ b/pkg/session/session_store.go @@ -2,7 +2,7 @@ package session import "github.com/sipeed/picoclaw/pkg/providers" -// SessionStore defines the persistence operations used by the agent loop. +// LegacyStore defines the persistence operations used by the agent loop. // Both SessionManager (legacy JSON backend) and JSONLBackend satisfy this // interface, allowing the storage layer to be swapped without touching the // agent loop code. @@ -10,7 +10,7 @@ import "github.com/sipeed/picoclaw/pkg/providers" // Write methods (Add*, Set*, Truncate*) are fire-and-forget: they do not // return errors. Implementations should log failures internally. This // matches the original SessionManager contract that the agent loop relies on. -type SessionStore interface { +type LegacyStore interface { // AddMessage appends a simple role/content message to the session. AddMessage(sessionKey, role, content string) // AddFullMessage appends a complete message including tool calls. diff --git a/pkg/tools/filesystem.go b/pkg/tools/filesystem.go index a44e372c3..7a165e123 100644 --- a/pkg/tools/filesystem.go +++ b/pkg/tools/filesystem.go @@ -489,6 +489,19 @@ func formatDirEntries(entries []os.DirEntry) *ToolResult { return NewToolResult(result.String()) } +// buildFs returns the appropriate fileSystem implementation based on the +// restrict flag and optional allow-path patterns. +func buildFs(workspace string, restrict bool, patterns []*regexp.Regexp) fileSystem { + if !restrict { + return &hostFs{} + } + sb := &sandboxFs{workspace: workspace} + if len(patterns) > 0 { + return &whitelistFs{sandbox: sb, patterns: patterns} + } + return sb +} + // fileSystem abstracts reading, writing, and listing files, allowing both // unrestricted (host filesystem) and sandbox (os.Root) implementations to share the same polymorphic interface. type fileSystem interface { diff --git a/pkg/tools/registry.go b/pkg/tools/registry.go index 0635f47d7..5b52e07a1 100644 --- a/pkg/tools/registry.go +++ b/pkg/tools/registry.go @@ -2,8 +2,10 @@ package tools import ( "context" + "encoding/json" "fmt" "sort" + "strings" "sync" "sync/atomic" "time" @@ -283,18 +285,36 @@ func (r *ToolRegistry) ToProviderDefs() []providers.ToolDefinition { desc, _ := fn["description"].(string) params, _ := fn["parameters"].(map[string]any) + var paramsRaw json.RawMessage + if params != nil { + paramsRaw, _ = json.Marshal(params) + } + definitions = append(definitions, providers.ToolDefinition{ Type: "function", Function: providers.ToolFunctionDefinition{ Name: name, Description: desc, - Parameters: params, + Parameters: paramsRaw, }, }) } return definitions } +// NormalizeToolName lowercases the name and strips underscores/hyphens so that +// "read_file", "ReadFile", and "read-file" all map to "readfile". +func NormalizeToolName(name string) string { + var b strings.Builder + b.Grow(len(name)) + for _, r := range strings.ToLower(name) { + if r != '_' && r != '-' { + b.WriteRune(r) + } + } + return b.String() +} + // List returns a list of all registered tool names. func (r *ToolRegistry) List() []string { r.mu.RLock() diff --git a/pkg/tools/registry_ext_test.go b/pkg/tools/registry_ext_test.go index cc42d2e4c..7a2857ff7 100644 --- a/pkg/tools/registry_ext_test.go +++ b/pkg/tools/registry_ext_test.go @@ -21,7 +21,7 @@ func (m *mockCtxTool) SetContext(channel, chatID string) { } func (m *mockAsyncRegistryTool) SetCallback(cb AsyncCallback) { - m.cb = cb + m.lastCB = cb } func TestNormalizeToolName(t *testing.T) { diff --git a/pkg/tools/registry_test.go b/pkg/tools/registry_test.go index 92d7d5abd..b1b822243 100644 --- a/pkg/tools/registry_test.go +++ b/pkg/tools/registry_test.go @@ -2,6 +2,7 @@ package tools import ( "context" + "encoding/json" "strings" "sync" "testing" @@ -240,12 +241,13 @@ func TestToolRegistry_ToProviderDefs(t *testing.T) { t.Fatalf("expected 1 provider def, got %d", len(defs)) } + paramsRaw, _ := json.Marshal(params) want := providers.ToolDefinition{ Type: "function", Function: providers.ToolFunctionDefinition{ Name: "beta", Description: "tool B", - Parameters: params, + Parameters: paramsRaw, }, } got := defs[0] diff --git a/pkg/tools/spawn_test.go b/pkg/tools/spawn_test.go index 43223b8db..b5652784a 100644 --- a/pkg/tools/spawn_test.go +++ b/pkg/tools/spawn_test.go @@ -8,7 +8,7 @@ import ( func TestSpawnTool_Execute_EmptyTask(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSpawnTool(manager) ctx := context.Background() @@ -42,7 +42,7 @@ func TestSpawnTool_Execute_EmptyTask(t *testing.T) { func TestSpawnTool_Execute_ValidTask(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSpawnTool(manager) ctx := context.Background() diff --git a/pkg/tools/subagent_tool_test.go b/pkg/tools/subagent_tool_test.go index 4b6f130a5..4e91a2af8 100644 --- a/pkg/tools/subagent_tool_test.go +++ b/pkg/tools/subagent_tool_test.go @@ -46,7 +46,7 @@ func (m *MockLLMProvider) GetContextWindow() int { func TestSubagentManager_SetLLMOptions_AppliesToRunToolLoop(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) manager.SetLLMOptions(2048, 0.6) tool := NewSubagentTool(manager) @@ -72,7 +72,7 @@ func TestSubagentManager_SetLLMOptions_AppliesToRunToolLoop(t *testing.T) { // TestSubagentTool_Name verifies tool name func TestSubagentTool_Name(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) if tool.Name() != "subagent" { @@ -83,7 +83,7 @@ func TestSubagentTool_Name(t *testing.T) { // TestSubagentTool_Description verifies tool description func TestSubagentTool_Description(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) desc := tool.Description() @@ -98,7 +98,7 @@ func TestSubagentTool_Description(t *testing.T) { // TestSubagentTool_Parameters verifies tool parameters schema func TestSubagentTool_Parameters(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) params := tool.Parameters() @@ -148,7 +148,7 @@ func TestSubagentTool_Parameters(t *testing.T) { // TestSubagentTool_Execute_Success tests successful execution func TestSubagentTool_Execute_Success(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) ctx := WithToolContext(context.Background(), "telegram", "chat-123") @@ -202,7 +202,7 @@ func TestSubagentTool_Execute_Success(t *testing.T) { // TestSubagentTool_Execute_NoLabel tests execution without label func TestSubagentTool_Execute_NoLabel(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) ctx := context.Background() @@ -225,7 +225,7 @@ func TestSubagentTool_Execute_NoLabel(t *testing.T) { // TestSubagentTool_Execute_MissingTask tests error handling for missing task func TestSubagentTool_Execute_MissingTask(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) ctx := context.Background() @@ -275,7 +275,7 @@ func TestSubagentTool_Execute_NilManager(t *testing.T) { // TestSubagentTool_Execute_ContextPassing verifies context is properly used func TestSubagentTool_Execute_ContextPassing(t *testing.T) { provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) channel := "test-channel" @@ -300,7 +300,7 @@ func TestSubagentTool_Execute_ContextPassing(t *testing.T) { func TestSubagentTool_ForUserTruncation(t *testing.T) { // Create a mock provider that returns very long content provider := &MockLLMProvider{} - manager := NewSubagentManager(provider, "test-model", "/tmp/test") + manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil, nil, WebSearchToolOptions{}) tool := NewSubagentTool(manager) ctx := context.Background()