fix: add tool call extraction to Codex CLI provider
- Extract shared tool call parsing into tool_call_extract.go (extractToolCallsFromText, stripToolCallsFromText, findMatchingBrace) - Both ClaudeCliProvider and CodexCliProvider now share the same tool call extraction logic for PicoClaw-specific tools - Fix cache token accounting: include cached_input_tokens in total - Add 2 new tests for tool call extraction from JSONL events - Update existing tests for corrected token calculations
This commit is contained in:
parent
d2d9de5614
commit
8a47a6fb67
4 changed files with 163 additions and 68 deletions
|
|
@ -171,68 +171,14 @@ func (p *ClaudeCliProvider) parseClaudeCliResponse(output string) (*LLMResponse,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// extractToolCalls parses tool call JSON from the response text.
|
// extractToolCalls delegates to the shared extractToolCallsFromText function.
|
||||||
func (p *ClaudeCliProvider) extractToolCalls(text string) []ToolCall {
|
func (p *ClaudeCliProvider) extractToolCalls(text string) []ToolCall {
|
||||||
start := strings.Index(text, `{"tool_calls"`)
|
return extractToolCallsFromText(text)
|
||||||
if start == -1 {
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
end := findMatchingBrace(text, start)
|
// stripToolCallsJSON delegates to the shared stripToolCallsFromText function.
|
||||||
if end == start {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
jsonStr := text[start:end]
|
|
||||||
|
|
||||||
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"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := json.Unmarshal([]byte(jsonStr), &wrapper); err != nil {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var result []ToolCall
|
|
||||||
for _, tc := range wrapper.ToolCalls {
|
|
||||||
var args map[string]interface{}
|
|
||||||
json.Unmarshal([]byte(tc.Function.Arguments), &args)
|
|
||||||
|
|
||||||
result = append(result, ToolCall{
|
|
||||||
ID: tc.ID,
|
|
||||||
Type: tc.Type,
|
|
||||||
Name: tc.Function.Name,
|
|
||||||
Arguments: args,
|
|
||||||
Function: &FunctionCall{
|
|
||||||
Name: tc.Function.Name,
|
|
||||||
Arguments: tc.Function.Arguments,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
return result
|
|
||||||
}
|
|
||||||
|
|
||||||
// stripToolCallsJSON removes tool call JSON from response text.
|
|
||||||
func (p *ClaudeCliProvider) stripToolCallsJSON(text string) string {
|
func (p *ClaudeCliProvider) stripToolCallsJSON(text string) string {
|
||||||
start := strings.Index(text, `{"tool_calls"`)
|
return stripToolCallsFromText(text)
|
||||||
if start == -1 {
|
|
||||||
return text
|
|
||||||
}
|
|
||||||
|
|
||||||
end := findMatchingBrace(text, start)
|
|
||||||
if end == start {
|
|
||||||
return text
|
|
||||||
}
|
|
||||||
|
|
||||||
return strings.TrimSpace(text[:start] + text[end:])
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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.
|
||||||
|
|
|
||||||
|
|
@ -199,10 +199,11 @@ func (p *CodexCliProvider) parseJSONLEvents(output string) (*LLMResponse, error)
|
||||||
}
|
}
|
||||||
case "turn.completed":
|
case "turn.completed":
|
||||||
if event.Usage != nil {
|
if event.Usage != nil {
|
||||||
|
promptTokens := event.Usage.InputTokens + event.Usage.CachedInputTokens
|
||||||
usage = &UsageInfo{
|
usage = &UsageInfo{
|
||||||
PromptTokens: event.Usage.InputTokens,
|
PromptTokens: promptTokens,
|
||||||
CompletionTokens: event.Usage.OutputTokens,
|
CompletionTokens: event.Usage.OutputTokens,
|
||||||
TotalTokens: event.Usage.InputTokens + event.Usage.OutputTokens,
|
TotalTokens: promptTokens + event.Usage.OutputTokens,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
case "error":
|
case "error":
|
||||||
|
|
@ -220,10 +221,19 @@ func (p *CodexCliProvider) parseJSONLEvents(output string) (*LLMResponse, error)
|
||||||
|
|
||||||
content := strings.Join(contentParts, "\n")
|
content := strings.Join(contentParts, "\n")
|
||||||
|
|
||||||
|
// Extract tool calls from response text (same pattern as ClaudeCliProvider)
|
||||||
|
toolCalls := extractToolCallsFromText(content)
|
||||||
|
|
||||||
|
finishReason := "stop"
|
||||||
|
if len(toolCalls) > 0 {
|
||||||
|
finishReason = "tool_calls"
|
||||||
|
content = stripToolCallsFromText(content)
|
||||||
|
}
|
||||||
|
|
||||||
return &LLMResponse{
|
return &LLMResponse{
|
||||||
Content: strings.TrimSpace(content),
|
Content: strings.TrimSpace(content),
|
||||||
ToolCalls: nil,
|
ToolCalls: toolCalls,
|
||||||
FinishReason: "stop",
|
FinishReason: finishReason,
|
||||||
Usage: usage,
|
Usage: usage,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
"os/exec"
|
"os/exec"
|
||||||
|
|
@ -32,20 +33,86 @@ func TestParseJSONLEvents_AgentMessage(t *testing.T) {
|
||||||
if resp.Usage == nil {
|
if resp.Usage == nil {
|
||||||
t.Fatal("Usage should not be nil")
|
t.Fatal("Usage should not be nil")
|
||||||
}
|
}
|
||||||
if resp.Usage.PromptTokens != 100 {
|
if resp.Usage.PromptTokens != 150 {
|
||||||
t.Errorf("PromptTokens = %d, want 100", resp.Usage.PromptTokens)
|
t.Errorf("PromptTokens = %d, want 150", resp.Usage.PromptTokens)
|
||||||
}
|
}
|
||||||
if resp.Usage.CompletionTokens != 20 {
|
if resp.Usage.CompletionTokens != 20 {
|
||||||
t.Errorf("CompletionTokens = %d, want 20", resp.Usage.CompletionTokens)
|
t.Errorf("CompletionTokens = %d, want 20", resp.Usage.CompletionTokens)
|
||||||
}
|
}
|
||||||
if resp.Usage.TotalTokens != 120 {
|
if resp.Usage.TotalTokens != 170 {
|
||||||
t.Errorf("TotalTokens = %d, want 120", resp.Usage.TotalTokens)
|
t.Errorf("TotalTokens = %d, want 170", resp.Usage.TotalTokens)
|
||||||
}
|
}
|
||||||
if len(resp.ToolCalls) != 0 {
|
if len(resp.ToolCalls) != 0 {
|
||||||
t.Errorf("ToolCalls should be empty, got %d", len(resp.ToolCalls))
|
t.Errorf("ToolCalls should be empty, got %d", len(resp.ToolCalls))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestParseJSONLEvents_ToolCallExtraction(t *testing.T) {
|
||||||
|
p := &CodexCliProvider{}
|
||||||
|
toolCallText := `Let me read that file.
|
||||||
|
{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"read_file","arguments":"{\"path\":\"/tmp/test.txt\"}"}}]}`
|
||||||
|
// Build valid JSONL by marshaling the event
|
||||||
|
item := codexEvent{
|
||||||
|
Type: "item.completed",
|
||||||
|
Item: &codexEventItem{ID: "item_1", Type: "agent_message", Text: toolCallText},
|
||||||
|
}
|
||||||
|
itemJSON, _ := json.Marshal(item)
|
||||||
|
usageEvt := `{"type":"turn.completed","usage":{"input_tokens":50,"cached_input_tokens":0,"output_tokens":20}}`
|
||||||
|
events := `{"type":"turn.started"}` + "\n" + string(itemJSON) + "\n" + usageEvt
|
||||||
|
|
||||||
|
resp, err := p.parseJSONLEvents(events)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parseJSONLEvents() error: %v", err)
|
||||||
|
}
|
||||||
|
if resp.FinishReason != "tool_calls" {
|
||||||
|
t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "tool_calls")
|
||||||
|
}
|
||||||
|
if len(resp.ToolCalls) != 1 {
|
||||||
|
t.Fatalf("ToolCalls count = %d, want 1", len(resp.ToolCalls))
|
||||||
|
}
|
||||||
|
if resp.ToolCalls[0].Name != "read_file" {
|
||||||
|
t.Errorf("ToolCalls[0].Name = %q, want %q", resp.ToolCalls[0].Name, "read_file")
|
||||||
|
}
|
||||||
|
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)
|
||||||
|
}
|
||||||
|
// Content should have the tool call JSON stripped
|
||||||
|
if strings.Contains(resp.Content, "tool_calls") {
|
||||||
|
t.Errorf("Content should not contain tool_calls JSON, got: %q", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseJSONLEvents_MultipleToolCalls(t *testing.T) {
|
||||||
|
p := &CodexCliProvider{}
|
||||||
|
toolCallText := `{"tool_calls":[{"id":"call_1","type":"function","function":{"name":"read_file","arguments":"{\"path\":\"a.txt\"}"}},{"id":"call_2","type":"function","function":{"name":"write_file","arguments":"{\"path\":\"b.txt\",\"content\":\"hello\"}"}}]}`
|
||||||
|
item := codexEvent{
|
||||||
|
Type: "item.completed",
|
||||||
|
Item: &codexEventItem{ID: "item_1", Type: "agent_message", Text: toolCallText},
|
||||||
|
}
|
||||||
|
itemJSON, _ := json.Marshal(item)
|
||||||
|
events := `{"type":"turn.started"}` + "\n" + string(itemJSON) + "\n" + `{"type":"turn.completed"}`
|
||||||
|
|
||||||
|
resp, err := p.parseJSONLEvents(events)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("parseJSONLEvents() error: %v", err)
|
||||||
|
}
|
||||||
|
if len(resp.ToolCalls) != 2 {
|
||||||
|
t.Fatalf("ToolCalls count = %d, want 2", len(resp.ToolCalls))
|
||||||
|
}
|
||||||
|
if resp.ToolCalls[0].Name != "read_file" {
|
||||||
|
t.Errorf("ToolCalls[0].Name = %q, want %q", resp.ToolCalls[0].Name, "read_file")
|
||||||
|
}
|
||||||
|
if resp.ToolCalls[1].Name != "write_file" {
|
||||||
|
t.Errorf("ToolCalls[1].Name = %q, want %q", resp.ToolCalls[1].Name, "write_file")
|
||||||
|
}
|
||||||
|
if resp.FinishReason != "tool_calls" {
|
||||||
|
t.Errorf("FinishReason = %q, want %q", resp.FinishReason, "tool_calls")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestParseJSONLEvents_MultipleMessages(t *testing.T) {
|
func TestParseJSONLEvents_MultipleMessages(t *testing.T) {
|
||||||
p := &CodexCliProvider{}
|
p := &CodexCliProvider{}
|
||||||
events := `{"type":"turn.started"}
|
events := `{"type":"turn.started"}
|
||||||
|
|
@ -372,8 +439,8 @@ func TestCodexCliProvider_MockCLI_Success(t *testing.T) {
|
||||||
if resp.Usage == nil {
|
if resp.Usage == nil {
|
||||||
t.Fatal("Usage should not be nil")
|
t.Fatal("Usage should not be nil")
|
||||||
}
|
}
|
||||||
if resp.Usage.PromptTokens != 50 {
|
if resp.Usage.PromptTokens != 60 {
|
||||||
t.Errorf("PromptTokens = %d, want 50", resp.Usage.PromptTokens)
|
t.Errorf("PromptTokens = %d, want 60", resp.Usage.PromptTokens)
|
||||||
}
|
}
|
||||||
if resp.Usage.CompletionTokens != 15 {
|
if resp.Usage.CompletionTokens != 15 {
|
||||||
t.Errorf("CompletionTokens = %d, want 15", resp.Usage.CompletionTokens)
|
t.Errorf("CompletionTokens = %d, want 15", resp.Usage.CompletionTokens)
|
||||||
|
|
|
||||||
72
pkg/providers/tool_call_extract.go
Normal file
72
pkg/providers/tool_call_extract.go
Normal file
|
|
@ -0,0 +1,72 @@
|
||||||
|
package providers
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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 := 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"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
Function struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Arguments string `json:"arguments"`
|
||||||
|
} `json:"function"`
|
||||||
|
} `json:"tool_calls"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal([]byte(jsonStr), &wrapper); err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var result []ToolCall
|
||||||
|
for _, tc := range wrapper.ToolCalls {
|
||||||
|
var args map[string]interface{}
|
||||||
|
json.Unmarshal([]byte(tc.Function.Arguments), &args)
|
||||||
|
|
||||||
|
result = append(result, ToolCall{
|
||||||
|
ID: tc.ID,
|
||||||
|
Type: tc.Type,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Arguments: args,
|
||||||
|
Function: &FunctionCall{
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Arguments: tc.Function.Arguments,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
}
|
||||||
|
|
||||||
|
// stripToolCallsFromText removes tool call JSON from response text.
|
||||||
|
func stripToolCallsFromText(text string) string {
|
||||||
|
start := strings.Index(text, `{"tool_calls"`)
|
||||||
|
if start == -1 {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
|
||||||
|
end := findMatchingBrace(text, start)
|
||||||
|
if end == start {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
|
||||||
|
return strings.TrimSpace(text[:start] + text[end:])
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue