fix: harden tool call extraction scanner to handle braces in strings
This commit is contained in:
parent
394d1d1197
commit
885ee1c8ea
3 changed files with 149 additions and 16 deletions
|
|
@ -181,22 +181,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 {
|
||||
|
|
|
|||
|
|
@ -70,3 +70,42 @@ func stripToolCallsFromText(text string) string {
|
|||
|
||||
return strings.TrimSpace(text[:start] + text[end:])
|
||||
}
|
||||
|
||||
// findMatchingBrace finds the index after the closing brace matching the opening brace at pos.
|
||||
// It accounts for braces inside strings and escaped characters.
|
||||
func findMatchingBrace(text string, pos int) int {
|
||||
depth := 0
|
||||
inString := false
|
||||
escaped := false
|
||||
|
||||
for i := pos; i < len(text); i++ {
|
||||
char := text[i]
|
||||
|
||||
if escaped {
|
||||
escaped = false
|
||||
continue
|
||||
}
|
||||
|
||||
if char == '\\' {
|
||||
escaped = true
|
||||
continue
|
||||
}
|
||||
|
||||
if char == '"' {
|
||||
inString = !inString
|
||||
continue
|
||||
}
|
||||
|
||||
if !inString {
|
||||
if char == '{' {
|
||||
depth++
|
||||
} else if char == '}' {
|
||||
depth--
|
||||
if depth == 0 {
|
||||
return i + 1
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return pos
|
||||
}
|
||||
|
|
|
|||
110
pkg/providers/tool_call_extract_test.go
Normal file
110
pkg/providers/tool_call_extract_test.go
Normal file
|
|
@ -0,0 +1,110 @@
|
|||
package providers
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestExtractToolCallsFromText(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
text string
|
||||
want []ToolCall
|
||||
}{
|
||||
{
|
||||
name: "Basic tool call",
|
||||
text: `Here is the tool call: {"tool_calls":[{"id":"call_1","type":"function","function":{"name":"search","arguments":"{\"query\":\"hello\"}"}}]} and some more text.`,
|
||||
want: []ToolCall{
|
||||
{
|
||||
ID: "call_1",
|
||||
Type: "function",
|
||||
Name: "search",
|
||||
Arguments: map[string]interface{}{
|
||||
"query": "hello",
|
||||
},
|
||||
Function: &FunctionCall{
|
||||
Name: "search",
|
||||
Arguments: `{"query":"hello"}`,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "Brace in string",
|
||||
text: `Tool call with brace in string: {"tool_calls":[{"id":"call_2","type":"function","function":{"name":"msg","arguments":"{\"text\":\"Hello { world }\"}"}}]} post-text.`,
|
||||
want: []ToolCall{
|
||||
{
|
||||
ID: "call_2",
|
||||
Type: "function",
|
||||
Name: "msg",
|
||||
Arguments: map[string]interface{}{
|
||||
"text": "Hello { world }",
|
||||
},
|
||||
Function: &FunctionCall{
|
||||
Name: "msg",
|
||||
Arguments: `{"text":"Hello { world }"}`,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "Escaped quote and brace in arguments",
|
||||
text: `Complex: {"tool_calls":[{"id":"call_3","type":"function","function":{"name":"exec","arguments":"{\"cmd\":\"echo \\\"}\\\"\"}"}}]}`,
|
||||
want: []ToolCall{
|
||||
{
|
||||
ID: "call_3",
|
||||
Type: "function",
|
||||
Name: "exec",
|
||||
Arguments: map[string]interface{}{
|
||||
"cmd": `echo "}"`,
|
||||
},
|
||||
Function: &FunctionCall{
|
||||
Name: "exec",
|
||||
Arguments: `{"cmd":"echo \"}\""}`,
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "No tool calls",
|
||||
text: "Just some normal text here.",
|
||||
want: nil,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := extractToolCallsFromText(tt.text)
|
||||
if !reflect.DeepEqual(got, tt.want) {
|
||||
t.Errorf("extractToolCallsFromText() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripToolCallsFromText(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
text string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "Basic strip",
|
||||
text: "Prefix text. {\"tool_calls\":[]} Suffix text.",
|
||||
want: "Prefix text. Suffix text.",
|
||||
},
|
||||
{
|
||||
name: "No tool calls to strip",
|
||||
text: "Normal text.",
|
||||
want: "Normal text.",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := stripToolCallsFromText(tt.text); got != tt.want {
|
||||
t.Errorf("stripToolCallsFromText() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue