From da7a4e498ba9e135ed4802bc72ba9b630f96f1be Mon Sep 17 00:00:00 2001 From: Eric Jacksch Date: Thu, 19 Mar 2026 14:37:09 -0400 Subject: [PATCH] fix(providers): robust CLI tool call extraction and mixed response handling Tool call detection previously relied on a literal strings.Index for '{"tool_calls"' which failed whenever the LLM returned pretty-printed JSON (newline after '{') or wrapped the output in markdown code fences. Arguments typed as a JSON object instead of an encoded string also caused a silent parse failure and leaked the raw JSON block to the user. Changes: - Strip markdown code fences (` + "```json" + ` / ` + "```" + `) before parsing - Locate JSON candidate via first '{' / last '}' instead of literal match - Unmarshal directly and check for top-level "tool_calls" key - Accept arguments as either a JSON-encoded string or a plain JSON object - Remove dead findMatchingBrace function and its tests - Publish response.Content to the user immediately when a response contains both text and tool calls (previously the text was silently discarded into session history) - Fix pre-existing test bug: TestCreateProvider_GeminiCliDefaultWorkspace now clears Agents.Defaults.Workspace before testing the '.' fallback Co-Authored-By: Claude Sonnet 4.6 --- pkg/agent/loop.go | 12 + pkg/providers/claude_cli_provider.go | 16 -- pkg/providers/claude_cli_provider_test.go | 24 -- pkg/providers/gemini_cli_provider_test.go | 1 + pkg/providers/tool_call_extract.go | 150 +++++++++--- pkg/providers/tool_call_extract_test.go | 265 ++++++++++++++++++++++ 6 files changed, 397 insertions(+), 71 deletions(-) create mode 100644 pkg/providers/tool_call_extract_test.go diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index cbd177708..979c56daf 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -1231,6 +1231,18 @@ func (al *AgentLoop) runLLMIteration( "iteration": iteration, }) + // If the LLM returned both text content and tool calls, publish the + // text to the user immediately so it is visible before tool execution. + if response.Content != "" && opts.Channel != "" { + pubCtx, pubCancel := context.WithTimeout(ctx, 5*time.Second) + _ = al.bus.PublishOutbound(pubCtx, bus.OutboundMessage{ + Channel: opts.Channel, + ChatID: opts.ChatID, + Content: response.Content, + }) + pubCancel() + } + // Build assistant message with tool calls assistantMsg := providers.Message{ Role: "assistant", diff --git a/pkg/providers/claude_cli_provider.go b/pkg/providers/claude_cli_provider.go index 40b581490..44de4613e 100644 --- a/pkg/providers/claude_cli_provider.go +++ b/pkg/providers/claude_cli_provider.go @@ -161,22 +161,6 @@ func (p *ClaudeCliProvider) stripToolCallsJSON(text string) string { return stripToolCallsFromText(text) } -// findMatchingBrace finds the index after the closing brace matching the opening brace at pos. -func findMatchingBrace(text string, pos int) int { - depth := 0 - for i := pos; i < len(text); i++ { - if text[i] == '{' { - depth++ - } else if text[i] == '}' { - depth-- - if depth == 0 { - return i + 1 - } - } - } - return pos -} - // claudeCliJSONResponse represents the JSON output from the claude CLI. // Matches the real claude CLI v2.x output format. type claudeCliJSONResponse struct { diff --git a/pkg/providers/claude_cli_provider_test.go b/pkg/providers/claude_cli_provider_test.go index d4d648f5a..efd8f2125 100644 --- a/pkg/providers/claude_cli_provider_test.go +++ b/pkg/providers/claude_cli_provider_test.go @@ -960,27 +960,3 @@ func TestStripToolCallsJSON_OnlyToolCalls(t *testing.T) { t.Errorf("stripToolCallsJSON() = %q, want empty", got) } } - -// --- findMatchingBrace tests --- - -func TestFindMatchingBrace(t *testing.T) { - tests := []struct { - text string - pos int - want int - }{ - {`{"a":1}`, 0, 7}, - {`{"a":{"b":2}}`, 0, 13}, - {`text {"a":1} more`, 5, 12}, - {`{unclosed`, 0, 0}, // no match returns pos - {`{}`, 0, 2}, // empty object - {`{{{}}}`, 0, 6}, // deeply nested - {`{"a":"b{c}d"}`, 0, 13}, // braces in strings (simplified matcher) - } - for _, tt := range tests { - got := findMatchingBrace(tt.text, tt.pos) - if got != tt.want { - t.Errorf("findMatchingBrace(%q, %d) = %d, want %d", tt.text, tt.pos, got, tt.want) - } - } -} diff --git a/pkg/providers/gemini_cli_provider_test.go b/pkg/providers/gemini_cli_provider_test.go index d513d8380..53e86f10d 100644 --- a/pkg/providers/gemini_cli_provider_test.go +++ b/pkg/providers/gemini_cli_provider_test.go @@ -261,6 +261,7 @@ func TestCreateProvider_GeminiCliWithModel(t *testing.T) { func TestCreateProvider_GeminiCliDefaultWorkspace(t *testing.T) { cfg := config.DefaultConfig() + cfg.Agents.Defaults.Workspace = "" // no workspace configured cfg.ModelList = []config.ModelConfig{ {ModelName: "gemini-cli", Model: "gemini-cli/gemini-cli"}, } diff --git a/pkg/providers/tool_call_extract.go b/pkg/providers/tool_call_extract.go index 7ddea0e99..3930e91d8 100644 --- a/pkg/providers/tool_call_extract.go +++ b/pkg/providers/tool_call_extract.go @@ -8,47 +8,60 @@ import ( // extractToolCallsFromText parses tool call JSON from response text. // Both ClaudeCliProvider and CodexCliProvider use this to extract // tool calls that the model outputs in its response text. +// +// The algorithm is robust to pretty-printed JSON and markdown code fences: +// 1. Strip any surrounding code fences (```json ... ``` or ``` ... ```) +// 2. Find the first '{' and last '}' to extract the JSON candidate +// 3. Unmarshal the candidate and check for a non-empty "tool_calls" array +// 4. Parse each tool call, handling arguments as either a JSON string or object func extractToolCallsFromText(text string) []ToolCall { - start := strings.Index(text, `{"tool_calls"`) - if start == -1 { + text = stripCodeFences(text) + + firstBrace := strings.Index(text, "{") + if firstBrace == -1 { + return nil + } + lastBrace := strings.LastIndex(text, "}") + if lastBrace == -1 || lastBrace <= firstBrace { return nil } - end := findMatchingBrace(text, start) - if end == start { - return nil - } - - jsonStr := text[start:end] + candidate := text[firstBrace : lastBrace+1] var wrapper struct { - ToolCalls []struct { - ID string `json:"id"` - Type string `json:"type"` - Function struct { - Name string `json:"name"` - Arguments string `json:"arguments"` - } `json:"function"` - } `json:"tool_calls"` + ToolCalls []json.RawMessage `json:"tool_calls"` } - - if err := json.Unmarshal([]byte(jsonStr), &wrapper); err != nil { + if err := json.Unmarshal([]byte(candidate), &wrapper); err != nil { + return nil + } + if len(wrapper.ToolCalls) == 0 { return nil } var result []ToolCall - for _, tc := range wrapper.ToolCalls { - var args map[string]any - json.Unmarshal([]byte(tc.Function.Arguments), &args) + for _, raw := range wrapper.ToolCalls { + var tc struct { + ID string `json:"id"` + Type string `json:"type"` + Function struct { + Name string `json:"name"` + Arguments json.RawMessage `json:"arguments"` + } `json:"function"` + } + if err := json.Unmarshal(raw, &tc); err != nil { + continue + } + + args, rawArgsStr := parseToolCallArguments(tc.Function.Arguments) result = append(result, ToolCall{ - ID: tc.ID, - Type: tc.Type, - Name: tc.Function.Name, + ID: tc.ID, + Type: tc.Type, + Name: tc.Function.Name, Arguments: args, Function: &FunctionCall{ Name: tc.Function.Name, - Arguments: tc.Function.Arguments, + Arguments: rawArgsStr, }, }) } @@ -56,17 +69,92 @@ func extractToolCallsFromText(text string) []ToolCall { return result } -// stripToolCallsFromText removes tool call JSON from response text. +// stripToolCallsFromText removes tool call JSON from response text, +// preserving any text before the first '{' or after the last '}'. func stripToolCallsFromText(text string) string { - start := strings.Index(text, `{"tool_calls"`) - if start == -1 { + stripped := stripCodeFences(text) + + firstBrace := strings.Index(stripped, "{") + if firstBrace == -1 { + return text + } + lastBrace := strings.LastIndex(stripped, "}") + if lastBrace == -1 || lastBrace <= firstBrace { return text } - end := findMatchingBrace(text, start) - if end == start { + // Validate that the candidate is actually a tool_calls JSON blob + candidate := stripped[firstBrace : lastBrace+1] + var wrapper struct { + ToolCalls []json.RawMessage `json:"tool_calls"` + } + if err := json.Unmarshal([]byte(candidate), &wrapper); err != nil || len(wrapper.ToolCalls) == 0 { return text } - return strings.TrimSpace(text[:start] + text[end:]) + // Use positions in the original (pre-fence-stripped) text to preserve + // any content outside code fences. If the text was fence-wrapped we strip + // the whole thing; otherwise cut out just the JSON range. + originalFirst := strings.Index(text, "{") + originalLast := strings.LastIndex(text, "}") + if originalFirst == -1 || originalLast == -1 || originalLast <= originalFirst { + return text + } + + return strings.TrimSpace(text[:originalFirst] + text[originalLast+1:]) } + +// parseToolCallArguments handles arguments that may be either a JSON string +// (which itself encodes a JSON object) or a JSON object directly. +// Returns the parsed map and a canonical string representation. +func parseToolCallArguments(raw json.RawMessage) (map[string]any, string) { + if len(raw) == 0 { + return nil, "" + } + + trimmed := strings.TrimSpace(string(raw)) + + if strings.HasPrefix(trimmed, `"`) { + // Arguments is a JSON-encoded string; decode the string then parse it as JSON. + var encoded string + if err := json.Unmarshal(raw, &encoded); err != nil { + return nil, trimmed + } + var args map[string]any + if err := json.Unmarshal([]byte(encoded), &args); err != nil { + return nil, encoded + } + return args, encoded + } + + if strings.HasPrefix(trimmed, "{") { + // Arguments is already a JSON object. + var args map[string]any + if err := json.Unmarshal(raw, &args); err != nil { + return nil, trimmed + } + return args, trimmed + } + + return nil, trimmed +} + +// stripCodeFences removes markdown code fences (```json...``` or ```...```) +// from text, returning the inner content trimmed of surrounding whitespace. +// If no code fence is present the original text is returned unchanged. +func stripCodeFences(text string) string { + s := strings.TrimSpace(text) + + for _, fence := range []string{"```json", "```"} { + if strings.HasPrefix(s, fence) { + rest := s[len(fence):] + // The fence opener may be immediately followed by content or a newline. + if idx := strings.Index(rest, "```"); idx != -1 { + return strings.TrimSpace(rest[:idx]) + } + } + } + + return text +} + diff --git a/pkg/providers/tool_call_extract_test.go b/pkg/providers/tool_call_extract_test.go new file mode 100644 index 000000000..f42f5ad93 --- /dev/null +++ b/pkg/providers/tool_call_extract_test.go @@ -0,0 +1,265 @@ +package providers + +import ( + "testing" +) + +// --- extractToolCallsFromText tests --- + +func TestExtractToolCallsFromText_NoJSON(t *testing.T) { + got := extractToolCallsFromText("Just plain text with no JSON.") + if len(got) != 0 { + t.Errorf("expected 0 tool calls, got %d", len(got)) + } +} + +func TestExtractToolCallsFromText_CompactJSON(t *testing.T) { + text := `{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"get_weather","arguments":"{\"location\":\"NYC\"}"}}]}` + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call, got %d", len(got)) + } + if got[0].Name != "get_weather" { + t.Errorf("Name = %q, want %q", got[0].Name, "get_weather") + } + if got[0].Arguments["location"] != "NYC" { + t.Errorf("Arguments[location] = %v, want NYC", got[0].Arguments["location"]) + } +} + +func TestExtractToolCallsFromText_PrettyPrintedJSON(t *testing.T) { + // Bug case: LLM returns pretty-printed JSON with newlines after { + text := `{ + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": "{\"location\": \"NYC\"}" + } + } + ] +}` + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call from pretty-printed JSON, got %d", len(got)) + } + if got[0].Name != "get_weather" { + t.Errorf("Name = %q, want %q", got[0].Name, "get_weather") + } + if got[0].Arguments["location"] != "NYC" { + t.Errorf("Arguments[location] = %v, want NYC", got[0].Arguments["location"]) + } +} + +func TestExtractToolCallsFromText_JSONInCodeFence(t *testing.T) { + // LLM wraps output in markdown code fences + text := "```json\n" + `{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"list_files","arguments":"{\"path\":\"/tmp\"}"}}]}` + "\n```" + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call from code-fenced JSON, got %d", len(got)) + } + if got[0].Name != "list_files" { + t.Errorf("Name = %q, want %q", got[0].Name, "list_files") + } + if got[0].Arguments["path"] != "/tmp" { + t.Errorf("Arguments[path] = %v, want /tmp", got[0].Arguments["path"]) + } +} + +func TestExtractToolCallsFromText_JSONInPlainCodeFence(t *testing.T) { + // Code fence without language tag + text := "```\n" + `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"do_thing","arguments":"{}"}}]}` + "\n```" + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call, got %d", len(got)) + } + if got[0].Name != "do_thing" { + t.Errorf("Name = %q, want %q", got[0].Name, "do_thing") + } +} + +func TestExtractToolCallsFromText_ArgumentsAsObject(t *testing.T) { + // Arguments field is a JSON object, not a JSON string + text := `{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"create_file","arguments":{"path":"/tmp/out.txt","content":"hello world"}}}]}` + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call, got %d", len(got)) + } + if got[0].Name != "create_file" { + t.Errorf("Name = %q, want %q", got[0].Name, "create_file") + } + if got[0].Arguments["path"] != "/tmp/out.txt" { + t.Errorf("Arguments[path] = %v, want /tmp/out.txt", got[0].Arguments["path"]) + } + if got[0].Arguments["content"] != "hello world" { + t.Errorf("Arguments[content] = %v, want 'hello world'", got[0].Arguments["content"]) + } +} + +func TestExtractToolCallsFromText_MixedTextAroundJSON(t *testing.T) { + // Text both before and after the JSON block + text := "Let me search for that.\n" + `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"search","arguments":"{\"query\":\"golang\"}"}}]}` + "\nI'll get back to you." + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call with surrounding text, got %d", len(got)) + } + if got[0].Name != "search" { + t.Errorf("Name = %q, want %q", got[0].Name, "search") + } + if got[0].Arguments["query"] != "golang" { + t.Errorf("Arguments[query] = %v, want golang", got[0].Arguments["query"]) + } +} + +func TestExtractToolCallsFromText_InvalidJSON(t *testing.T) { + got := extractToolCallsFromText(`{"tool_calls":invalid}`) + if len(got) != 0 { + t.Errorf("expected 0 tool calls for invalid JSON, got %d", len(got)) + } +} + +func TestExtractToolCallsFromText_EmptyToolCalls(t *testing.T) { + got := extractToolCallsFromText(`{"tool_calls":[]}`) + if len(got) != 0 { + t.Errorf("expected 0 tool calls for empty array, got %d", len(got)) + } +} + +func TestExtractToolCallsFromText_NoToolCallsKey(t *testing.T) { + got := extractToolCallsFromText(`{"other_key":"value"}`) + if len(got) != 0 { + t.Errorf("expected 0 tool calls for unrelated JSON, got %d", len(got)) + } +} + +func TestExtractToolCallsFromText_MultipleToolCalls(t *testing.T) { + text := `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"read_file","arguments":"{\"path\":\"/a\"}"}},{"id":"c2","type":"function","function":{"name":"write_file","arguments":"{\"path\":\"/b\",\"content\":\"x\"}"}}]}` + got := extractToolCallsFromText(text) + if len(got) != 2 { + t.Fatalf("expected 2 tool calls, got %d", len(got)) + } + if got[0].Name != "read_file" { + t.Errorf("[0].Name = %q, want read_file", got[0].Name) + } + if got[1].Name != "write_file" { + t.Errorf("[1].Name = %q, want write_file", got[1].Name) + } +} + +func TestExtractToolCallsFromText_FunctionFieldPreserved(t *testing.T) { + text := `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"fn","arguments":"{\"k\":\"v\"}"}}]}` + got := extractToolCallsFromText(text) + if len(got) != 1 { + t.Fatalf("expected 1 tool call, got %d", len(got)) + } + if got[0].Function == nil { + t.Fatal("Function field should not be nil") + } + if got[0].Function.Name != "fn" { + t.Errorf("Function.Name = %q, want fn", got[0].Function.Name) + } + if got[0].Function.Arguments == "" { + t.Error("Function.Arguments should contain raw JSON string") + } +} + +// --- stripToolCallsFromText tests --- + +func TestStripToolCallsFromText_NoToolCalls(t *testing.T) { + text := "Just regular text." + got := stripToolCallsFromText(text) + if got != text { + t.Errorf("stripToolCallsFromText() = %q, want %q", got, text) + } +} + +func TestStripToolCallsFromText_OnlyJSON(t *testing.T) { + text := `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"fn","arguments":"{}"}}]}` + got := stripToolCallsFromText(text) + if got != "" { + t.Errorf("stripToolCallsFromText() = %q, want empty", got) + } +} + +func TestStripToolCallsFromText_TextBeforeJSON(t *testing.T) { + text := "Let me check the weather.\n" + `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"fn","arguments":"{}"}}]}` + got := stripToolCallsFromText(text) + if got != "Let me check the weather." { + t.Errorf("stripToolCallsFromText() = %q, want %q", got, "Let me check the weather.") + } +} + +func TestStripToolCallsFromText_TextAfterJSON(t *testing.T) { + text := `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"fn","arguments":"{}"}}]}` + "\nDone." + got := stripToolCallsFromText(text) + if got != "Done." { + t.Errorf("stripToolCallsFromText() = %q, want %q", got, "Done.") + } +} + +func TestStripToolCallsFromText_MixedTextAroundJSON(t *testing.T) { + text := "Before.\n" + `{"tool_calls":[{"id":"c1","type":"function","function":{"name":"fn","arguments":"{}"}}]}` + "\nAfter." + got := stripToolCallsFromText(text) + // TrimSpace only trims the outer ends; internal whitespace from surrounding text is kept. + if got != "Before.\n\nAfter." { + t.Errorf("stripToolCallsFromText() = %q, want %q", got, "Before.\n\nAfter.") + } +} + +func TestStripToolCallsFromText_PrettyPrintedJSON(t *testing.T) { + text := "Thinking...\n" + "{\n \"tool_calls\": [{\"id\":\"c1\",\"type\":\"function\",\"function\":{\"name\":\"fn\",\"arguments\":\"{}\"}}]\n}" + got := stripToolCallsFromText(text) + if got == text { + t.Error("stripToolCallsFromText() should remove pretty-printed JSON block") + } + if got != "Thinking..." { + t.Errorf("stripToolCallsFromText() = %q, want %q", got, "Thinking...") + } +} + +func TestStripToolCallsFromText_UnrelatedJSON(t *testing.T) { + // A JSON object that is not a tool_calls wrapper should be left alone + text := `{"key":"value"}` + got := stripToolCallsFromText(text) + if got != text { + t.Errorf("stripToolCallsFromText() = %q, want %q (unrelated JSON should be preserved)", got, text) + } +} + +// --- stripCodeFences tests --- + +func TestStripCodeFences_JSONFence(t *testing.T) { + inner := `{"tool_calls":[]}` + text := "```json\n" + inner + "\n```" + got := stripCodeFences(text) + if got != inner { + t.Errorf("stripCodeFences() = %q, want %q", got, inner) + } +} + +func TestStripCodeFences_PlainFence(t *testing.T) { + inner := `{"tool_calls":[]}` + text := "```\n" + inner + "\n```" + got := stripCodeFences(text) + if got != inner { + t.Errorf("stripCodeFences() = %q, want %q", got, inner) + } +} + +func TestStripCodeFences_NoFence(t *testing.T) { + text := `{"tool_calls":[]}` + got := stripCodeFences(text) + if got != text { + t.Errorf("stripCodeFences() = %q, want %q (no fence should return original)", got, text) + } +} + +func TestStripCodeFences_PlainText(t *testing.T) { + text := "Just some text." + got := stripCodeFences(text) + if got != text { + t.Errorf("stripCodeFences() = %q, want %q", got, text) + } +}