返回 DeepSeek-Reasonix
pi_recovery_test.go
根目录 / internal / agent / pi_recovery_test.go
1 package agent
2
3 import (
4 "context"
5 "errors"
6 "testing"
7
8 "reasonix/internal/agent/testutil"
9 "reasonix/internal/event"
10 "reasonix/internal/provider"
11 )
12
13 func TestCompatibleMissingReasoningDoesNotRegenerate(t *testing.T) {
14 mock := testutil.NewMock("compatible", testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "one", Name: "echo", Arguments: `{"text":"hi"}`}}}, testutil.Turn{Text: "done"})
15 sink := &recordSink{}
16 a := New(toolCallReasoningRequiredProvider{mock}, echoRegistry(), NewSession(""), Options{MissingReasoningWarnStateDir: t.TempDir()}, sink)
17 if err := a.Run(withNoClosedLoop(context.Background()), "go"); err != nil {
18 t.Fatal(err)
19 }
20 if mock.CallCount() != 2 || len(sink.kinds(event.ToolResult)) != 1 || len(sink.kinds(event.Retrying)) != 0 {
21 t.Fatalf("calls=%d tools=%d retries=%d", mock.CallCount(), len(sink.kinds(event.ToolResult)), len(sink.kinds(event.Retrying)))
22 }
23 }
24
25 func TestTruncatedArgumentsNeverExecute(t *testing.T) {
26 mock := testutil.NewMock("m", testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "cut", Name: "echo", Arguments: `{"text":"partial"}`}}, Usage: &provider.Usage{FinishReason: "length"}}, testutil.Turn{Text: "recovered"})
27 sink := &recordSink{}
28 a := New(mock, echoRegistry(), NewSession(""), Options{}, sink)
29 if err := a.Run(withNoClosedLoop(context.Background()), "go"); err != nil {
30 t.Fatal(err)
31 }
32 results := sink.kinds(event.ToolResult)
33 if len(results) != 1 || results[0].Tool.Output == "echoed: partial" {
34 t.Fatalf("results=%+v", results)
35 }
36 for _, m := range a.Session().Snapshot() {
37 if m.ToolCallID == "cut" && m.ToolRunState != provider.ToolRunNotStarted {
38 t.Fatalf("state=%s", m.ToolRunState)
39 }
40 }
41 }
42
43 func TestPartialStreamNeverEntersContinuousWaiting(t *testing.T) {
44 turns := make([]testutil.Turn, maxSamplingAttempts)
45 for i := range turns {
46 turns[i] = testutil.Turn{Text: "partial", ChunkError: provider.StreamInterrupt(errors.New("closed"), provider.StreamInterruptPrematureEOF)}
47 }
48 mock := testutil.NewMock("m", turns...)
49 sink := &recordSink{}
50 a := New(mock, echoRegistry(), NewSession(""), Options{}, sink)
51 if err := a.Run(withNoClosedLoop(context.Background()), "go"); err == nil {
52 t.Fatal("partial stream accepted")
53 }
54 if mock.CallCount() != 1 {
55 t.Fatalf("calls=%d", mock.CallCount())
56 }
57 for _, e := range sink.kinds(event.Retrying) {
58 if e.Recovery != nil && e.Recovery.Waiting {
59 t.Fatal("partial generation waited indefinitely")
60 }
61 }
62 }
63
64 func TestKnownSpendStillStopsRecoveryWithUnknownUsage(t *testing.T) {
65 var b runBudget
66 b.observe(&provider.Usage{PromptTokens: 1000000, Unknown: true}, &provider.Pricing{Input: 2, Currency: "USD"})
67 if axis, _ := b.exceeded(TaskBudget{Cost: 1}); axis != "cost" || b.totals().Priced {
68 t.Fatalf("axis=%s totals=%+v", axis, b.totals())
69 }
70 }
71
72 func TestPiReferenceRetryScenarios(t *testing.T) {
73 for _, tc := range []struct {
74 name string
75 turns []testutil.Turn
76 calls int
77 success bool
78 }{
79 {"temporary_failure_is_terminal", []testutil.Turn{{StreamError: &provider.APIError{Status: 503}}, {Text: "done"}}, 1, false},
80 {"quota", []testutil.Turn{{StreamError: &provider.APIError{Status: 429, Body: "insufficient_quota"}}}, 1, false},
81 } {
82 t.Run(tc.name, func(t *testing.T) {
83 p := testutil.NewMock("reference", tc.turns...)
84 a := New(p, echoRegistry(), NewSession(""), Options{}, event.Discard)
85 err := a.Run(withNoClosedLoop(context.Background()), "go")
86 if (err == nil) != tc.success || p.CallCount() != tc.calls {
87 t.Fatalf("calls=%d err=%v", p.CallCount(), err)
88 }
89 })
90 }
91 }
92
93 func TestMixedFailuresCannotRenewRecoveryBudget(t *testing.T) {
94 missing := testutil.Turn{ToolCalls: []provider.ToolCall{{ID: "missing", Name: "echo", Arguments: `{"text":"unsafe"}`}}}
95 p := testutil.NewMock("strict", testutil.Turn{StreamError: &provider.APIError{Status: 503}}, testutil.Turn{StreamError: &provider.APIError{Status: 503}}, missing, missing)
96 sink := &recordSink{}
97 a := New(strictToolCallReasoningProvider{p}, echoRegistry(), NewSession(""), Options{}, sink)
98 if err := a.Run(withNoClosedLoop(context.Background()), "go"); err == nil {
99 t.Fatal("invalid reasoning accepted")
100 }
101 if p.CallCount() != 1 || len(sink.kinds(event.ToolResult)) != 0 {
102 t.Fatalf("calls=%d tools=%d", p.CallCount(), len(sink.kinds(event.ToolResult)))
103 }
104 }
105
106 type canceledCompletionProvider struct{ cancel context.CancelFunc }
107
108 func (*canceledCompletionProvider) Name() string { return "late" }
109 func (p *canceledCompletionProvider) Stream(context.Context, provider.Request) (<-chan provider.Chunk, error) {
110 ch := make(chan provider.Chunk, 2)
111 ch <- provider.Chunk{Type: provider.ChunkToolCall, ToolCall: &provider.ToolCall{ID: "late", Name: "echo", Arguments: `{"text":"late"}`}}
112 ch <- provider.Chunk{Type: provider.ChunkDone}
113 close(ch)
114 p.cancel()
115 return ch, nil
116 }
117 func TestCanceledCompletionCannotStartTools(t *testing.T) {
118 ctx, cancel := context.WithCancel(context.Background())
119 defer cancel()
120 sink := &recordSink{}
121 a := New(&canceledCompletionProvider{cancel}, echoRegistry(), NewSession(""), Options{}, sink)
122 if err := a.Run(withNoClosedLoop(ctx), "go"); !errors.Is(err, context.Canceled) {
123 t.Fatalf("err=%v", err)
124 }
125 if len(sink.kinds(event.ToolResult)) != 0 {
126 t.Fatal("late completion executed a tool")
127 }
128 }
129
129 lines GO