| 1 | package agent |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "fmt" |
| 7 | "strings" |
| 8 | "sync" |
| 9 | "testing" |
| 10 | |
| 11 | "reasonix/internal/event" |
| 12 | "reasonix/internal/provider" |
| 13 | "reasonix/internal/tool" |
| 14 | ) |
| 15 | |
| 16 | type overflowLoopTool struct{ output string } |
| 17 | |
| 18 | func (overflowLoopTool) Name() string { return "grow_context" } |
| 19 | func (overflowLoopTool) Description() string { return "Return a large deterministic result." } |
| 20 | func (overflowLoopTool) Schema() json.RawMessage { return json.RawMessage(`{"type":"object"}`) } |
| 21 | func (overflowLoopTool) ReadOnly() bool { return true } |
| 22 | func (t overflowLoopTool) Execute(context.Context, json.RawMessage) (string, error) { |
| 23 | return t.output, nil |
| 24 | } |
| 25 | |
| 26 | // repeatedOverflowProvider forces two independent context-limit recoveries in |
| 27 | // one Run. Summary requests are answered separately and never advance the tool |
| 28 | // loop, so the test exercises the production recovery/compaction wiring. |
| 29 | type repeatedOverflowProvider struct { |
| 30 | mu sync.Mutex |
| 31 | toolCalls int |
| 32 | overflows int |
| 33 | summaries int |
| 34 | rejectedAt map[int]bool |
| 35 | requestsAt map[int]int |
| 36 | maxToolCalls int |
| 37 | } |
| 38 | |
| 39 | func (p *repeatedOverflowProvider) Name() string { return "repeated-overflow" } |
| 40 | func (p *repeatedOverflowProvider) ContextBudgetPolicy() provider.ContextBudgetPolicy { |
| 41 | return provider.ContextBudgetPolicy{ |
| 42 | WindowMode: provider.ContextWindowShared, |
| 43 | AutoOutputTokens: 1024, |
| 44 | MaxOutputTokens: 1024, |
| 45 | LimitMode: provider.OutputLimitAlways, |
| 46 | } |
| 47 | } |
| 48 | |
| 49 | func (p *repeatedOverflowProvider) Stream(_ context.Context, req provider.Request) (<-chan provider.Chunk, error) { |
| 50 | p.mu.Lock() |
| 51 | defer p.mu.Unlock() |
| 52 | |
| 53 | if len(req.Messages) > 0 && strings.Contains(req.Messages[len(req.Messages)-1].Content, "Compact the preceding conversation prefix") { |
| 54 | p.summaries++ |
| 55 | return chunks( |
| 56 | provider.Chunk{Type: provider.ChunkText, Text: "- goal: finish the tool loop\n- pending: continue"}, |
| 57 | provider.Chunk{Type: provider.ChunkDone}, |
| 58 | ), nil |
| 59 | } |
| 60 | p.requestsAt[p.toolCalls]++ |
| 61 | |
| 62 | if (p.toolCalls == 3 || p.toolCalls == 6) && !p.rejectedAt[p.toolCalls] { |
| 63 | p.rejectedAt[p.toolCalls] = true |
| 64 | p.overflows++ |
| 65 | return nil, &provider.ContextLimitError{ |
| 66 | APIError: &provider.APIError{Provider: p.Name(), Status: 400, Body: "context limit exceeded"}, |
| 67 | WindowTokens: 24_000, |
| 68 | PromptTokens: 23_000, |
| 69 | CompletionTokens: 2_000, |
| 70 | RequestedTokens: 25_000, |
| 71 | } |
| 72 | } |
| 73 | |
| 74 | if p.toolCalls >= p.maxToolCalls { |
| 75 | return chunks( |
| 76 | provider.Chunk{Type: provider.ChunkText, Text: "Done."}, |
| 77 | provider.Chunk{Type: provider.ChunkDone}, |
| 78 | ), nil |
| 79 | } |
| 80 | |
| 81 | p.toolCalls++ |
| 82 | return chunks( |
| 83 | provider.Chunk{Type: provider.ChunkToolCall, ToolCall: &provider.ToolCall{ |
| 84 | ID: fmt.Sprintf("grow-%d", p.toolCalls), Name: "grow_context", Arguments: `{}`, |
| 85 | }}, |
| 86 | provider.Chunk{Type: provider.ChunkDone}, |
| 87 | ), nil |
| 88 | } |
| 89 | |
| 90 | func chunks(items ...provider.Chunk) <-chan provider.Chunk { |
| 91 | ch := make(chan provider.Chunk, len(items)) |
| 92 | for _, item := range items { |
| 93 | ch <- item |
| 94 | } |
| 95 | close(ch) |
| 96 | return ch |
| 97 | } |
| 98 | |
| 99 | func TestToolLoopRetriesAfterEachCompleteToolResult(t *testing.T) { |
| 100 | prov := &repeatedOverflowProvider{rejectedAt: make(map[int]bool), requestsAt: make(map[int]int), maxToolCalls: 9} |
| 101 | reg := tool.NewRegistry() |
| 102 | reg.Add(overflowLoopTool{output: strings.Repeat("large deterministic tool output. ", 700)}) |
| 103 | |
| 104 | a := New(prov, reg, NewSession("system"), Options{ |
| 105 | ContextWindow: 100_000, |
| 106 | CompactRatio: defaultCompactRatio, |
| 107 | MaxOutputTokens: 1024, |
| 108 | }, event.Discard) |
| 109 | |
| 110 | err := a.Run(context.Background(), "keep using the tool until the provider says the task is done") |
| 111 | if err != nil { |
| 112 | t.Fatalf("complete repeated tool results should advance projection recovery: %v", err) |
| 113 | } |
| 114 | |
| 115 | prov.mu.Lock() |
| 116 | defer prov.mu.Unlock() |
| 117 | if prov.overflows != 2 { |
| 118 | t.Fatalf("provider overflows = %d, want 2", prov.overflows) |
| 119 | } |
| 120 | if prov.requestsAt[3] != 2 || prov.requestsAt[6] != 2 { |
| 121 | t.Fatalf("requests at overflow points = %v, want one retry after each new tool-result span", prov.requestsAt) |
| 122 | } |
| 123 | } |
| 124 |