返回 DeepSeek-Reasonix
termination_policy_test.go
根目录 / internal / control / termination_policy_test.go
1 package control
2
3 import (
4 "bytes"
5 "context"
6 "encoding/json"
7 "io"
8 "net/http"
9 "net/http/httptest"
10 "path/filepath"
11 "reflect"
12 "sync"
13 "testing"
14 "time"
15
16 "reasonix/internal/provider"
17 "reasonix/internal/provider/anthropic"
18 "reasonix/internal/provider/openai"
19 "reasonix/internal/provider/responses"
20 "reasonix/internal/session"
21 )
22
23 func TestTerminationPolicyPreservesSemanticsAndProviderBytes(t *testing.T) {
24 system := provider.Message{ID: "system", Role: provider.RoleSystem, Content: "stable system"}
25 user := provider.Message{ID: "user", Role: provider.RoleUser, Content: "update files"}
26 call := provider.Message{ID: "assistant", Role: provider.RoleAssistant, Content: "calling", ReasoningContent: "original reasoning", ReasoningID: "reason-1", ReasoningStatus: "completed", ReasoningSignature: "original-proof", ToolCalls: []provider.ToolCall{{ID: "done", Name: "lookup", Arguments: `{"q":"value"}`, ThoughtSignature: "thought-proof"}}}
27 result := provider.Message{ID: "result", Role: provider.RoleTool, ToolCallID: "done", Name: "lookup", Content: "result", ToolRunState: provider.ToolRunCompleted}
28 partial := provider.Message{ID: "partial", Role: provider.RoleAssistant, Content: "unfinished", ReasoningContent: "partial reasoning", ReasoningSignature: "partial-proof"}
29 summary := provider.Message{ID: "summary", Role: provider.RoleUser, Content: "<compaction-summary>\nearlier work\n</compaction-summary>"}
30 local := provider.Message{ID: "local", Role: provider.RoleTool, Content: "old partial", LocalOnly: true, ToolCallID: provider.LocalOnlyToolID, Name: provider.LocalOnlyToolName, InterruptedTurn: &provider.InterruptedTurnRecovery{Pending: true, DroppedPartialText: true}}
31 partialBatch := call
32 partialBatch.ToolCalls = append(append([]provider.ToolCall{}, call.ToolCalls...), provider.ToolCall{ID: "unknown", Name: "lookup", Arguments: `{}`})
33 fixtures := []struct {
34 name string
35 input, expected []provider.Message
36 fallback provider.Message
37 replay bool
38 droppedReasoning bool
39 }{
40 {name: "partial-reasoning", input: []provider.Message{system, user, partial}, expected: []provider.Message{system, user}, replay: true, droppedReasoning: true},
41 {name: "paired-tool-round", input: []provider.Message{system, user, call, result, partial}, expected: []provider.Message{system, user, call, result}, replay: true, droppedReasoning: true},
42 {name: "partial-tool-batch", input: []provider.Message{system, user, partialBatch, result}, expected: []provider.Message{system, user}, replay: true, droppedReasoning: true},
43 {name: "unreplayable-pair", input: []provider.Message{system, user, call, result}, expected: []provider.Message{system, user}, replay: false, droppedReasoning: true},
44 {name: "compaction-and-partial", input: []provider.Message{system, user, summary, call, result, partial}, expected: []provider.Message{system, user, summary, call, result}, replay: true, droppedReasoning: true},
45 {name: "existing-local-only", input: []provider.Message{system, user, local}, expected: []provider.Message{system, user}, replay: true},
46 {name: "pre-executor-fallback", input: []provider.Message{system}, expected: []provider.Message{system, user}, fallback: user, replay: true},
47 }
48 for _, fixture := range fixtures {
49 t.Run(fixture.name, func(t *testing.T) {
50 before, err := json.Marshal(fixture.input)
51 if err != nil {
52 t.Fatal(err)
53 }
54 fallbackBefore, _ := json.Marshal(fixture.fallback)
55 got := planCancelledMessages(fixture.input, 1, fixture.fallback, time.Time{}, func(provider.Message) bool { return fixture.replay }, nil)
56 after, _ := json.Marshal(fixture.input)
57 fallbackAfter, _ := json.Marshal(fixture.fallback)
58 if !bytes.Equal(before, after) || !bytes.Equal(fallbackBefore, fallbackAfter) {
59 t.Fatal("planner mutated its inputs")
60 }
61 expected := provider.ModelMessages(fixture.expected)
62 if !reflect.DeepEqual(provider.ModelMessages(got), expected) {
63 t.Fatalf("model projection differs from explicit legacy-policy fixture:\ngot=%+v\nwant=%+v", provider.ModelMessages(got), expected)
64 }
65 if len(got) == 0 || got[len(got)-1].InterruptedTurn == nil {
66 t.Fatalf("missing recovery handoff: %+v", got)
67 }
68 recovery := got[len(got)-1].InterruptedTurn
69 if !recovery.Pending || recovery.DroppedPartialReasoning != fixture.droppedReasoning {
70 t.Fatalf("recovery=%+v", recovery)
71 }
72 if fixture.name == "paired-tool-round" && len(recovery.CompletedTools) != 1 {
73 t.Fatalf("lost completed tool fact: %+v", recovery)
74 }
75 if fixture.name == "partial-tool-batch" && (len(recovery.CompletedTools) != 1 || len(recovery.UnknownTools) != 1) {
76 t.Fatalf("partial batch facts: %+v", recovery)
77 }
78 reopenedMessages := persistTerminationPolicyModel(t, got)
79 for _, adapter := range []string{"openai", "anthropic", "responses"} {
80 t.Run(adapter, func(t *testing.T) {
81 var mu sync.Mutex
82 var bodies [][]byte
83 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
84 body, _ := io.ReadAll(r.Body)
85 mu.Lock()
86 bodies = append(bodies, body)
87 mu.Unlock()
88 w.Header().Set("Content-Type", "application/json")
89 w.WriteHeader(http.StatusBadRequest)
90 _, _ = io.WriteString(w, `{"error":{"message":"captured"}}`)
91 }))
92 defer server.Close()
93 var p provider.Provider
94 switch adapter {
95 case "openai":
96 p, err = openai.New(provider.Config{Name: "openai", BaseURL: server.URL, Model: "cache-model", APIKey: "test"})
97 case "anthropic":
98 p, err = anthropic.New(provider.Config{Name: "anthropic", BaseURL: server.URL, Model: "cache-model", APIKey: "test"})
99 case "responses":
100 p = responses.New(responses.Config{Name: "responses", BaseURL: server.URL, Model: "cache-model", APIKey: "test", Mode: "stateless"})
101 }
102 if err != nil {
103 t.Fatal(err)
104 }
105 for _, messages := range [][]provider.Message{fixture.expected, got, reopenedMessages} {
106 stream, err := p.Stream(t.Context(), provider.Request{Messages: messages, Tools: []provider.ToolSchema{{Name: "lookup", Description: "look up", Parameters: json.RawMessage(`{"type":"object","properties":{"q":{"type":"string"}}}`)}}})
107 if err == nil {
108 for range stream {
109 }
110 }
111 }
112 mu.Lock()
113 defer mu.Unlock()
114 if len(bodies) != 3 {
115 t.Fatalf("captured %d requests, want 3", len(bodies))
116 }
117 if !bytes.Equal(bodies[0], bodies[1]) {
118 t.Fatalf("provider bytes changed:\nexpected:%s\nactual:%s", bodies[0], bodies[1])
119 }
120 if !bytes.Equal(bodies[0], bodies[2]) {
121 t.Fatalf("provider bytes changed after reopen:\nexpected:%s\nactual:%s", bodies[0], bodies[2])
122 }
123 })
124 }
125 })
126 }
127 }
128
129 func persistTerminationPolicyModel(t *testing.T, messages []provider.Message) []provider.Message {
130 t.Helper()
131 dir := filepath.Join(t.TempDir(), "termination-policy")
132 store, err := session.CreateStore(dir, "termination-policy")
133 if err != nil {
134 t.Fatal(err)
135 }
136 payload, err := json.Marshal(map[string]any{"messages": messages})
137 if err != nil {
138 t.Fatal(err)
139 }
140 if _, err := store.Append(t.Context(), session.Batch{OperationID: "model-cleanup", Events: []session.Event{{Kind: "model/context-replace", Payload: payload}}}); err != nil {
141 t.Fatal(err)
142 }
143 if _, err := store.Flush(t.Context()); err != nil {
144 t.Fatal(err)
145 }
146 if err := store.Close(t.Context()); err != nil {
147 t.Fatal(err)
148 }
149 reopened, err := session.Open(dir, "termination-policy")
150 if err != nil {
151 t.Fatal(err)
152 }
153 defer reopened.Close(context.Background())
154 return reopened.DeriveMessages()
155 }
156
157 func TestTerminationPolicyFallbackOwnsImageSlice(t *testing.T) {
158 fallback := provider.Message{ID: "user", Role: provider.RoleUser, Content: "inspect image", Images: []string{"original"}}
159 got := planCancelledMessages(nil, 0, fallback, time.Time{}, func(provider.Message) bool { return true }, nil)
160 if len(got) < 1 || len(got[0].Images) != 1 {
161 t.Fatalf("fallback missing: %+v", got)
162 }
163 got[0].Images[0] = "changed"
164 if fallback.Images[0] != "original" {
165 t.Fatal("fallback image slice aliases input")
166 }
167 }
168
168 lines GO