| 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 |