返回 DeepSeek-Reasonix
context_recovery_tool_loop_test.go
根目录 / internal / agent / context_recovery_tool_loop_test.go
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
124 lines GO