2026-07-01 06:53:26 +00:00
|
|
|
package agent
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
"context"
|
|
|
|
|
"encoding/json"
|
|
|
|
|
"errors"
|
|
|
|
|
"fmt"
|
|
|
|
|
"iter"
|
2026-07-09 06:27:17 +00:00
|
|
|
"strings"
|
2026-07-01 06:53:26 +00:00
|
|
|
"testing"
|
|
|
|
|
"time"
|
|
|
|
|
|
|
|
|
|
llm "github.com/VictorVargas/rony-llm-agent/pkg/llm"
|
|
|
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/persona"
|
|
|
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/tools"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
type mockLLM struct {
|
|
|
|
|
generateFunc func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error)
|
|
|
|
|
streamFunc func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error]
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (m *mockLLM) Generate(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
return m.generateFunc(ctx, req)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (m *mockLLM) Stream(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
|
|
|
return m.streamFunc(ctx, req)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (m *mockLLM) Name() string { return "mock" }
|
|
|
|
|
func (m *mockLLM) Capabilities() llm.ProviderCapabilities { return llm.ProviderCapabilities{} }
|
|
|
|
|
|
|
|
|
|
type mockSandbox struct {
|
|
|
|
|
validateFunc func(tool tools.Tool, call llm.ToolCall) error
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func (m *mockSandbox) ValidateToolCall(tool tools.Tool, call llm.ToolCall) error {
|
|
|
|
|
if m.validateFunc != nil {
|
|
|
|
|
return m.validateFunc(tool, call)
|
|
|
|
|
}
|
|
|
|
|
return nil
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestNew(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
if loop == nil {
|
|
|
|
|
t.Fatal("expected non-nil loop")
|
|
|
|
|
}
|
|
|
|
|
if loop.cfg.MaxIters != DefaultMaxIterations {
|
|
|
|
|
t.Errorf("expected default max iters %d, got %d", DefaultMaxIterations, loop.cfg.MaxIters)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestNew_CustomMaxIters(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
MaxIters: 10,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
if loop.cfg.MaxIters != 10 {
|
|
|
|
|
t.Errorf("expected max iters 10, got %d", loop.cfg.MaxIters)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_NoToolCalls(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
Content: "I understand.",
|
|
|
|
|
Usage: llm.TokenUsage{
|
|
|
|
|
InputTokens: 10,
|
|
|
|
|
OutputTokens: 5,
|
|
|
|
|
TotalTokens: 15,
|
|
|
|
|
},
|
|
|
|
|
}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
resp, err := loop.Run(context.Background(), "Hello")
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if resp.Content != "I understand." {
|
|
|
|
|
t.Errorf("expected 'I understand.', got %q", resp.Content)
|
|
|
|
|
}
|
|
|
|
|
if resp.Iterations != 1 {
|
|
|
|
|
t.Errorf("expected 1 iteration, got %d", resp.Iterations)
|
|
|
|
|
}
|
|
|
|
|
if resp.TokenUsage.InputTokens != 10 {
|
|
|
|
|
t.Errorf("expected 10 input tokens, got %d", resp.TokenUsage.InputTokens)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-09 06:27:17 +00:00
|
|
|
func TestRun_IncludesAgentsMD(t *testing.T) {
|
|
|
|
|
var capturedSystemPrompt string
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
capturedSystemPrompt = req.Messages[0].Content
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
AgentsMD: "Never edit go.mod directly.",
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
if _, err := loop.Run(context.Background(), "Hello"); err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if !strings.Contains(capturedSystemPrompt, "Never edit go.mod directly.") {
|
|
|
|
|
t.Errorf("expected system prompt to include AGENTS.md content, got %q", capturedSystemPrompt)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-01 06:53:26 +00:00
|
|
|
func TestRun_ToolCalls(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "greet",
|
|
|
|
|
Description: "Greet someone",
|
|
|
|
|
InputSchema: json.RawMessage(`{"type":"object","properties":{"name":{"type":"string"}}}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "Hello!"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
ToolCalls: []llm.ToolCall{
|
|
|
|
|
{ID: "call-1", Name: "greet", Arguments: json.RawMessage(`{"name":"World"}`)},
|
|
|
|
|
},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
Content: "Done!",
|
|
|
|
|
}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
resp, err := loop.Run(context.Background(), "Say hi")
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if resp.Content != "Done!" {
|
|
|
|
|
t.Errorf("expected 'Done!', got %q", resp.Content)
|
|
|
|
|
}
|
|
|
|
|
if resp.Iterations != 2 {
|
|
|
|
|
t.Errorf("expected 2 iterations, got %d", resp.Iterations)
|
|
|
|
|
}
|
|
|
|
|
if len(resp.ToolCalls) != 1 {
|
|
|
|
|
t.Errorf("expected 1 tool call, got %d", len(resp.ToolCalls))
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
Fix tool call tracking and streaming assembly for all providers
- agent/loop.go: Record assistant message with ToolCalls before tool results,
and set ToolCallID on tool-result messages so follow-up requests have complete
context (prevents models from losing track of already-attempted tools).
- llm/providers/llamacpp/client.go: Buffer fragmented tool call deltas during
streaming, assemble them into complete calls when finish_reason arrives.
Add ToolCalls, ToolCallID, Name fields to request building.
- llm/providers/openai/client.go: Send ToolCalls, ToolCallID, Name when
building chat requests so messages are wire-format correct.
- llm/types.go: Add ToolCalls field to Message struct for serialization
back into conversation history.
- agent/integration_test.go: Move integration test skip from TestMain to a
per-test skipUnlessIntegration() so it doesn't hide other package tests.
- sandbox & tools: Add edge-case tests (relative traversal, array paths,
non-path strings, zero-value guards, sentinel errors).
2026-07-08 23:11:57 +00:00
|
|
|
// TestRun_ToolCalls_RecordsAssistantTurnAndToolCallID is the regression test
|
|
|
|
|
// for a bug where the follow-up request sent to the model, after executing
|
|
|
|
|
// a tool, never included the assistant message that requested the call
|
|
|
|
|
// (with its ToolCalls) nor set ToolCallID on the tool-result message. Some
|
|
|
|
|
// chat templates get confused by a "tool" message with nothing to attribute
|
|
|
|
|
// it to and the model loses track of what it already tried, which produced
|
|
|
|
|
// exactly the symptom reported in production: the model re-greeting and
|
|
|
|
|
// re-attempting the same search over and over instead of ever converging.
|
|
|
|
|
func TestRun_ToolCalls_RecordsAssistantTurnAndToolCallID(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "greet",
|
|
|
|
|
Description: "Greet someone",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "Hello!"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
var requests []llm.CompletionRequest
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
requests = append(requests, req)
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
Content: "Voy a saludar.",
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "call-1", Name: "greet", Arguments: json.RawMessage(`{"name":"World"}`)}},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
return llm.CompletionResponse{Content: "Done!"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
if _, err := loop.Run(context.Background(), "Say hi"); err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(requests) != 2 {
|
|
|
|
|
t.Fatalf("expected 2 requests to the model, got %d", len(requests))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// The second request (the follow-up after the tool ran) must contain
|
|
|
|
|
// the assistant's own tool_calls turn, immediately followed by a tool
|
|
|
|
|
// message whose ToolCallID matches it.
|
|
|
|
|
second := requests[1].Messages
|
|
|
|
|
var assistantIdx, toolIdx = -1, -1
|
|
|
|
|
for i, m := range second {
|
|
|
|
|
if m.Role == llm.RoleAssistant && len(m.ToolCalls) > 0 {
|
|
|
|
|
assistantIdx = i
|
|
|
|
|
}
|
|
|
|
|
if m.Role == llm.RoleTool {
|
|
|
|
|
toolIdx = i
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
if assistantIdx == -1 {
|
|
|
|
|
t.Fatalf("expected an assistant message carrying ToolCalls in the follow-up request, got %+v", second)
|
|
|
|
|
}
|
|
|
|
|
if second[assistantIdx].Content != "Voy a saludar." {
|
|
|
|
|
t.Errorf("expected the assistant message to keep its original content, got %q", second[assistantIdx].Content)
|
|
|
|
|
}
|
|
|
|
|
if second[assistantIdx].ToolCalls[0].ID != "call-1" || second[assistantIdx].ToolCalls[0].Name != "greet" {
|
|
|
|
|
t.Errorf("expected the recorded tool call to match what was requested, got %+v", second[assistantIdx].ToolCalls[0])
|
|
|
|
|
}
|
|
|
|
|
if toolIdx == -1 {
|
|
|
|
|
t.Fatalf("expected a tool-result message in the follow-up request, got %+v", second)
|
|
|
|
|
}
|
|
|
|
|
if second[toolIdx].ToolCallID != "call-1" {
|
|
|
|
|
t.Errorf("expected the tool message's ToolCallID to be %q, got %q", "call-1", second[toolIdx].ToolCallID)
|
|
|
|
|
}
|
|
|
|
|
if toolIdx <= assistantIdx {
|
|
|
|
|
t.Errorf("expected the tool-result message to come after the assistant's tool_calls message")
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// TestRun_Stream_ToolCalls_RecordsAssistantTurnAndToolCallID is the
|
|
|
|
|
// streaming counterpart of the test above.
|
|
|
|
|
func TestRun_Stream_ToolCalls_RecordsAssistantTurnAndToolCallID(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "greet",
|
|
|
|
|
Description: "Greet",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "greeted"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
var requests []llm.CompletionRequest
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
streamFunc: func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
|
|
|
requests = append(requests, req)
|
|
|
|
|
callCount++
|
|
|
|
|
return func(yield func(llm.StreamChunk, error) bool) {
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
yield(llm.StreamChunk{Delta: "Voy a saludar."}, nil)
|
|
|
|
|
yield(llm.StreamChunk{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "call-9", Name: "greet", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
FinishReason: "tool_calls",
|
|
|
|
|
}, nil)
|
|
|
|
|
} else {
|
|
|
|
|
yield(llm.StreamChunk{Delta: "done"}, nil)
|
|
|
|
|
yield(llm.StreamChunk{FinishReason: "stop"}, nil)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
for _, err := range loop.RunStream(context.Background(), "test") {
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected stream error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(requests) != 2 {
|
|
|
|
|
t.Fatalf("expected 2 requests to the model, got %d", len(requests))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
second := requests[1].Messages
|
|
|
|
|
var assistantIdx, toolIdx = -1, -1
|
|
|
|
|
for i, m := range second {
|
|
|
|
|
if m.Role == llm.RoleAssistant && len(m.ToolCalls) > 0 {
|
|
|
|
|
assistantIdx = i
|
|
|
|
|
}
|
|
|
|
|
if m.Role == llm.RoleTool {
|
|
|
|
|
toolIdx = i
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
if assistantIdx == -1 {
|
|
|
|
|
t.Fatalf("expected an assistant message carrying ToolCalls in the follow-up request, got %+v", second)
|
|
|
|
|
}
|
|
|
|
|
if second[assistantIdx].Content != "Voy a saludar." {
|
|
|
|
|
t.Errorf("expected the assistant message to carry the content streamed before the tool call, got %q", second[assistantIdx].Content)
|
|
|
|
|
}
|
|
|
|
|
if second[assistantIdx].ToolCalls[0].ID != "call-9" {
|
|
|
|
|
t.Errorf("expected the recorded tool call ID to be %q, got %q", "call-9", second[assistantIdx].ToolCalls[0].ID)
|
|
|
|
|
}
|
|
|
|
|
if toolIdx == -1 || second[toolIdx].ToolCallID != "call-9" {
|
|
|
|
|
t.Fatalf("expected a tool-result message with ToolCallID %q, got %+v", "call-9", second)
|
|
|
|
|
}
|
|
|
|
|
if toolIdx <= assistantIdx {
|
|
|
|
|
t.Errorf("expected the tool-result message to come after the assistant's tool_calls message")
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-01 06:53:26 +00:00
|
|
|
func TestRun_MaxIterations(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "x", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "x",
|
|
|
|
|
Description: "x",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "ok"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
MaxIters: 3,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
_, err := loop.Run(context.Background(), "test")
|
|
|
|
|
if err == nil {
|
|
|
|
|
t.Fatal("expected error, got nil")
|
|
|
|
|
}
|
|
|
|
|
if err.Error() != "max iterations (3) reached" {
|
|
|
|
|
t.Errorf("expected 'max iterations (3) reached', got %q", err.Error())
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_ToolNotFound(t *testing.T) {
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "nonexistent", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
resp, err := loop.Run(context.Background(), "test")
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if resp.Content != "done" {
|
|
|
|
|
t.Errorf("expected 'done', got %q", resp.Content)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_ApprovalDenied(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "dangerous",
|
|
|
|
|
Description: "Do something dangerous",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "executed"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Ask,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "dangerous", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
Approver: func(tool tools.Tool, call llm.ToolCall) bool {
|
|
|
|
|
return false // deny all
|
|
|
|
|
},
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
resp, err := loop.Run(context.Background(), "test")
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if len(resp.ToolCalls) != 1 {
|
|
|
|
|
t.Errorf("expected 1 tool call, got %d", len(resp.ToolCalls))
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_SandboxViolation(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "restricted",
|
|
|
|
|
Description: "Restricted tool",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "executed"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "restricted", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
Sandbox: &mockSandbox{
|
|
|
|
|
validateFunc: func(tool tools.Tool, call llm.ToolCall) error {
|
|
|
|
|
return fmt.Errorf("path traversal detected")
|
|
|
|
|
},
|
|
|
|
|
},
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
resp, err := loop.Run(context.Background(), "test")
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if len(resp.ToolCalls) != 0 {
|
|
|
|
|
t.Errorf("expected 0 tool calls (sandbox rejected), got %d", len(resp.ToolCalls))
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_OnIterationHook(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "greet",
|
|
|
|
|
Description: "Greet",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "hello"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
return llm.CompletionResponse{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "greet", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil
|
|
|
|
|
}
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
var iterations []Iteration
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
OnIteration: func(iter Iteration) {
|
|
|
|
|
iterations = append(iterations, iter)
|
|
|
|
|
},
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
_, err := loop.Run(context.Background(), "test")
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
if len(iterations) != 1 {
|
|
|
|
|
t.Errorf("expected 1 iteration hook call, got %d", len(iterations))
|
|
|
|
|
}
|
|
|
|
|
if iterations[0].ToolsUsed != 1 {
|
|
|
|
|
t.Errorf("expected 1 tool used, got %d", iterations[0].ToolsUsed)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_Stream_NoToolCalls(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
streamFunc: func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
|
|
|
return func(yield func(llm.StreamChunk, error) bool) {
|
|
|
|
|
yield(llm.StreamChunk{Delta: "Hello"}, nil)
|
|
|
|
|
yield(llm.StreamChunk{Delta: " world"}, nil)
|
|
|
|
|
yield(llm.StreamChunk{FinishReason: "stop"}, nil)
|
|
|
|
|
}
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
var chunks []llm.StreamChunk
|
|
|
|
|
stream := loop.RunStream(context.Background(), "test")
|
|
|
|
|
for chunk, err := range stream {
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
chunks = append(chunks, chunk)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(chunks) != 2 {
|
|
|
|
|
t.Errorf("expected 2 chunks, got %d", len(chunks))
|
|
|
|
|
}
|
|
|
|
|
if chunks[0].Delta != "Hello" {
|
|
|
|
|
t.Errorf("expected 'Hello', got %q", chunks[0].Delta)
|
|
|
|
|
}
|
|
|
|
|
if chunks[1].Delta != " world" {
|
|
|
|
|
t.Errorf("expected ' world', got %q", chunks[1].Delta)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-06 07:05:30 +00:00
|
|
|
func TestRun_Stream_ForwardsTrailingUsageOnlyChunk(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
streamFunc: func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
|
|
|
return func(yield func(llm.StreamChunk, error) bool) {
|
|
|
|
|
yield(llm.StreamChunk{Delta: "Hello"}, nil)
|
|
|
|
|
// No Delta/ReasoningDelta, as providers report usage in a
|
|
|
|
|
// separate trailing event; it must still be forwarded.
|
|
|
|
|
yield(llm.StreamChunk{Usage: llm.TokenUsage{InputTokens: 10, OutputTokens: 3, TotalTokens: 13}}, nil)
|
|
|
|
|
}
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
var chunks []llm.StreamChunk
|
|
|
|
|
stream := loop.RunStream(context.Background(), "test")
|
|
|
|
|
for chunk, err := range stream {
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
chunks = append(chunks, chunk)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(chunks) != 2 {
|
|
|
|
|
t.Fatalf("expected 2 chunks (content + usage-only), got %d", len(chunks))
|
|
|
|
|
}
|
|
|
|
|
usage := chunks[len(chunks)-1].Usage
|
|
|
|
|
if usage.InputTokens != 10 || usage.OutputTokens != 3 || usage.TotalTokens != 13 {
|
|
|
|
|
t.Errorf("expected the trailing usage-only chunk to be forwarded, got %+v", usage)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
2026-07-01 06:53:26 +00:00
|
|
|
func TestRun_Stream_WithToolCalls(t *testing.T) {
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "greet",
|
|
|
|
|
Description: "Greet",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "greeted"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
callCount := 0
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
streamFunc: func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
|
|
|
return func(yield func(llm.StreamChunk, error) bool) {
|
|
|
|
|
callCount++
|
|
|
|
|
if callCount == 1 {
|
|
|
|
|
yield(llm.StreamChunk{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "greet", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil)
|
|
|
|
|
} else {
|
|
|
|
|
yield(llm.StreamChunk{Delta: "done"}, nil)
|
|
|
|
|
yield(llm.StreamChunk{FinishReason: "stop"}, nil)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
var chunks []llm.StreamChunk
|
|
|
|
|
stream := loop.RunStream(context.Background(), "test")
|
|
|
|
|
for chunk, err := range stream {
|
|
|
|
|
if err != nil {
|
|
|
|
|
t.Fatalf("unexpected error: %v", err)
|
|
|
|
|
}
|
|
|
|
|
chunks = append(chunks, chunk)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if len(chunks) != 1 {
|
|
|
|
|
t.Errorf("expected 1 chunk (only the 'done' chunk), got %d", len(chunks))
|
|
|
|
|
}
|
|
|
|
|
if chunks[0].Delta != "done" {
|
|
|
|
|
t.Errorf("expected 'done', got %q", chunks[0].Delta)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_Stream_MaxIterations(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
streamFunc: func(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
|
|
|
return func(yield func(llm.StreamChunk, error) bool) {
|
|
|
|
|
yield(llm.StreamChunk{
|
|
|
|
|
ToolCalls: []llm.ToolCall{{ID: "1", Name: "x", Arguments: json.RawMessage("{}")}},
|
|
|
|
|
}, nil)
|
|
|
|
|
}
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
registry := tools.NewRegistry()
|
|
|
|
|
registry.Register(tools.Tool{
|
|
|
|
|
Name: "x",
|
|
|
|
|
Description: "x",
|
|
|
|
|
InputSchema: json.RawMessage(`{}`),
|
|
|
|
|
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
|
|
|
|
|
return tools.ToolResult{Content: "ok"}, nil
|
|
|
|
|
},
|
|
|
|
|
Permission: tools.Allow,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: registry,
|
|
|
|
|
MaxIters: 1,
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
var chunks []llm.StreamChunk
|
|
|
|
|
stream := loop.RunStream(context.Background(), "test")
|
|
|
|
|
for chunk, err := range stream {
|
|
|
|
|
if err != nil {
|
|
|
|
|
// expect max iterations error
|
|
|
|
|
if !errors.Is(err, context.DeadlineExceeded) && err.Error() != "max iterations (1) reached" {
|
|
|
|
|
t.Fatalf("expected max iterations error, got: %v", err)
|
|
|
|
|
}
|
|
|
|
|
continue
|
|
|
|
|
}
|
|
|
|
|
chunks = append(chunks, chunk)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func TestRun_Timeout(t *testing.T) {
|
|
|
|
|
mockClient := &mockLLM{
|
|
|
|
|
generateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
|
|
|
select {
|
|
|
|
|
case <-ctx.Done():
|
|
|
|
|
return llm.CompletionResponse{}, ctx.Err()
|
|
|
|
|
default:
|
|
|
|
|
return llm.CompletionResponse{Content: "done"}, nil
|
|
|
|
|
}
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond)
|
|
|
|
|
defer cancel()
|
|
|
|
|
time.Sleep(20 * time.Millisecond) // ensure context is cancelled before calling
|
|
|
|
|
|
|
|
|
|
loop := New(Config{
|
|
|
|
|
LLM: mockClient,
|
|
|
|
|
Persona: persona.DefaultPersona(),
|
|
|
|
|
Tools: tools.NewRegistry(),
|
|
|
|
|
})
|
|
|
|
|
|
|
|
|
|
_, err := loop.Run(ctx, "test")
|
|
|
|
|
if err == nil {
|
|
|
|
|
t.Fatal("expected timeout error, got nil")
|
|
|
|
|
}
|
|
|
|
|
}
|