fix(provider): harden text tool-call extraction parsing

This commit is contained in:
Alix-007 2026-03-29 16:13:35 +08:00
parent e70928cc6f
commit fd3a0fa782
3 changed files with 185 additions and 35 deletions

View file

@ -164,10 +164,32 @@ func (p *ClaudeCliProvider) stripToolCallsJSON(text string) string {
// findMatchingBrace finds the index after the closing brace matching the opening brace at pos. // findMatchingBrace finds the index after the closing brace matching the opening brace at pos.
func findMatchingBrace(text string, pos int) int { func findMatchingBrace(text string, pos int) int {
depth := 0 depth := 0
inString := false
escaped := false
for i := pos; i < len(text); i++ { for i := pos; i < len(text); i++ {
if text[i] == '{' { ch := text[i]
if inString {
if escaped {
escaped = false
continue
}
if ch == '\\' {
escaped = true
continue
}
if ch == '"' {
inString = false
}
continue
}
switch ch {
case '"':
inString = true
case '{':
depth++ depth++
} else if text[i] == '}' { case '}':
depth-- depth--
if depth == 0 { if depth == 0 {
return i + 1 return i + 1

View file

@ -5,41 +5,28 @@ import (
"strings" "strings"
) )
// extractToolCallsFromText parses tool call JSON from response text. type textToolCall struct {
// Both ClaudeCliProvider and CodexCliProvider use this to extract
// tool calls that the model outputs in its response text.
func extractToolCallsFromText(text string) []ToolCall {
start := strings.Index(text, `{"tool_calls"`)
if start == -1 {
return nil
}
end := findMatchingBrace(text, start)
if end == start {
return nil
}
jsonStr := text[start:end]
var wrapper struct {
ToolCalls []struct {
ID string `json:"id"` ID string `json:"id"`
Type string `json:"type"` Type string `json:"type"`
Function struct { Function struct {
Name string `json:"name"` Name string `json:"name"`
Arguments string `json:"arguments"` Arguments string `json:"arguments"`
} `json:"function"` } `json:"function"`
} `json:"tool_calls"`
} }
if err := json.Unmarshal([]byte(jsonStr), &wrapper); err != nil { // 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.
func extractToolCallsFromText(text string) []ToolCall {
start, end, rawToolCalls, ok := locateToolCallObject(text)
if !ok || end <= start {
return nil return nil
} }
var result []ToolCall var result []ToolCall
for _, tc := range wrapper.ToolCalls { for _, tc := range rawToolCalls {
var args map[string]any var args map[string]any
json.Unmarshal([]byte(tc.Function.Arguments), &args) _ = json.Unmarshal([]byte(tc.Function.Arguments), &args)
result = append(result, ToolCall{ result = append(result, ToolCall{
ID: tc.ID, ID: tc.ID,
@ -58,15 +45,51 @@ func extractToolCallsFromText(text string) []ToolCall {
// stripToolCallsFromText removes tool call JSON from response text. // stripToolCallsFromText removes tool call JSON from response text.
func stripToolCallsFromText(text string) string { func stripToolCallsFromText(text string) string {
start := strings.Index(text, `{"tool_calls"`) start, end, _, ok := locateToolCallObject(text)
if start == -1 { if !ok || end <= start {
return text
}
end := findMatchingBrace(text, start)
if end == start {
return text return text
} }
return strings.TrimSpace(text[:start] + text[end:]) return strings.TrimSpace(text[:start] + text[end:])
} }
func locateToolCallObject(text string) (start int, end int, toolCalls []textToolCall, ok bool) {
for i := 0; i < len(text); i++ {
if text[i] != '{' {
continue
}
candidateEnd := findMatchingBrace(text, i)
if candidateEnd <= i {
continue
}
candidateToolCalls, matched := parseToolCallsObject(text[i:candidateEnd])
if !matched {
continue
}
return i, candidateEnd, candidateToolCalls, true
}
return 0, 0, nil, false
}
func parseToolCallsObject(jsonStr string) ([]textToolCall, bool) {
var obj map[string]json.RawMessage
if err := json.Unmarshal([]byte(jsonStr), &obj); err != nil {
return nil, false
}
rawToolCalls, exists := obj["tool_calls"]
if !exists {
return nil, false
}
var toolCalls []textToolCall
if err := json.Unmarshal(rawToolCalls, &toolCalls); err != nil {
return nil, false
}
return toolCalls, true
}

View file

@ -0,0 +1,105 @@
package providers
import "testing"
func TestExtractToolCallsFromText_IndentedObjectWithExtraText(t *testing.T) {
text := `Planning tool execution...
{
"message": "calling tools",
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {
"name": "read_file",
"arguments": "{\"path\":\"/tmp/test.txt\"}"
}
}
],
"status": "ok"
}
Done.`
got := extractToolCallsFromText(text)
if len(got) != 1 {
t.Fatalf("extractToolCallsFromText() len = %d, want 1", len(got))
}
if got[0].Name != "read_file" {
t.Fatalf("tool name = %q, want %q", got[0].Name, "read_file")
}
if got[0].Arguments["path"] != "/tmp/test.txt" {
t.Fatalf("tool args[path] = %v, want /tmp/test.txt", got[0].Arguments["path"])
}
}
func TestExtractToolCallsFromText_ToolCallsNotFirstField(t *testing.T) {
text := `{"kind":"assistant_result","metadata":{"provider":"x"},"tool_calls":[{"id":"call_7","type":"function","function":{"name":"write_file","arguments":"{\"path\":\"a.txt\",\"content\":\"hello\"}"}}]}`
got := extractToolCallsFromText(text)
if len(got) != 1 {
t.Fatalf("extractToolCallsFromText() len = %d, want 1", len(got))
}
if got[0].ID != "call_7" {
t.Fatalf("tool id = %q, want %q", got[0].ID, "call_7")
}
if got[0].Name != "write_file" {
t.Fatalf("tool name = %q, want %q", got[0].Name, "write_file")
}
}
func TestExtractToolCallsFromText_SkipsInvalidCandidateAndFindsValidObject(t *testing.T) {
text := `prefix {"tool_calls":invalid} middle {"note":"valid-json-without-tools"} tail {
"note": "valid-json-with-tools",
"tool_calls": [
{
"id": "call_2",
"type": "function",
"function": {
"name": "get_weather",
"arguments": "{\"city\":\"Tokyo\"}"
}
}
]
}`
got := extractToolCallsFromText(text)
if len(got) != 1 {
t.Fatalf("extractToolCallsFromText() len = %d, want 1", len(got))
}
if got[0].Name != "get_weather" {
t.Fatalf("tool name = %q, want %q", got[0].Name, "get_weather")
}
}
func TestStripToolCallsFromText_DoesNotStripInvalidToolCallsObject(t *testing.T) {
text := `before {"tool_calls":"not-an-array","other":1} after`
got := stripToolCallsFromText(text)
if got != text {
t.Fatalf("stripToolCallsFromText() = %q, want unchanged", got)
}
}
func TestStripToolCallsFromText_StripsValidIndentedObject(t *testing.T) {
text := `before
{
"message": "tool call follows",
"tool_calls": [
{
"id": "call_3",
"type": "function",
"function": {
"name": "list_files",
"arguments": "{}"
}
}
]
}
after`
got := stripToolCallsFromText(text)
want := "before\n\nafter"
if got != want {
t.Fatalf("stripToolCallsFromText() = %q, want %q", got, want)
}
}