rony-llm-agent/pkg/agent/loop_test.go
Victor Vargas 0f835a0802 feat(agent): wire AGENTS.md discovery into the agent loop's system prompt
Export persona.DiscoverAgentsMD and add agent.Config.AgentsMD so project
and global AGENTS.md rules actually reach the model. Previously
buildInitialMessages always called AssembleSystemPrompt with an empty
string, so no AGENTS.md content was ever injected despite the discovery
logic already existing.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-08 23:27:17 -07:00

749 lines
22 KiB
Go

package agent
import (
"context"
"encoding/json"
"errors"
"fmt"
"iter"
"strings"
"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)
}
}
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)
}
}
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))
}
}
// 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")
}
}
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)
}
}
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)
}
}
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")
}
}