| 1 | package control |
| 2 | |
| 3 | import ( |
| 4 | "bytes" |
| 5 | "context" |
| 6 | "encoding/json" |
| 7 | "errors" |
| 8 | "net/http" |
| 9 | "net/url" |
| 10 | "strings" |
| 11 | "testing" |
| 12 | "time" |
| 13 | |
| 14 | "golang.org/x/net/http2" |
| 15 | "reasonix/internal/agent" |
| 16 | "reasonix/internal/event" |
| 17 | "reasonix/internal/provider" |
| 18 | "reasonix/internal/provider/openai" |
| 19 | "reasonix/internal/session" |
| 20 | "reasonix/internal/tool" |
| 21 | ) |
| 22 | |
| 23 | type diagnosticFailureTransport struct { |
| 24 | calls int |
| 25 | beforeFailure func() |
| 26 | } |
| 27 | |
| 28 | func (r *diagnosticFailureTransport) RoundTrip(*http.Request) (*http.Response, error) { |
| 29 | r.calls++ |
| 30 | if r.beforeFailure != nil { |
| 31 | r.beforeFailure() |
| 32 | } |
| 33 | return nil, http2.ConnectionError(http2.ErrCodeProtocol) |
| 34 | } |
| 35 | |
| 36 | func TestProviderFailureSurvivesCleanupAndColdExport(t *testing.T) { |
| 37 | for _, async := range []bool{false, true} { |
| 38 | t.Run(map[bool]string{false: "synchronous", true: "desktop-send"}[async], func(t *testing.T) { |
| 39 | for _, broken := range []bool{false, true} { |
| 40 | t.Run(map[bool]string{false: "valid-evidence", true: "unencodable-evidence"}[broken], func(t *testing.T) { |
| 41 | testProviderFailureColdExport(t, async, broken) |
| 42 | }) |
| 43 | } |
| 44 | }) |
| 45 | } |
| 46 | } |
| 47 | |
| 48 | func testProviderFailureColdExport(t *testing.T, async, broken bool) { |
| 49 | service, err := session.NewService("local", session.NewFilesystemPersistence(t.TempDir())) |
| 50 | if err != nil { |
| 51 | t.Fatal(err) |
| 52 | } |
| 53 | runtime, err := service.Create(t.Context(), session.CreateOptions{SessionID: "provider-diagnostic"}) |
| 54 | if err != nil { |
| 55 | t.Fatal(err) |
| 56 | } |
| 57 | transport := &diagnosticFailureTransport{} |
| 58 | p, err := openai.New(provider.Config{Name: "test", Protocol: "openai", Model: "test-model", APIKey: "private-key", BaseURL: "https://provider.test/v1", HTTPClient: &http.Client{Transport: transport}}) |
| 59 | if err != nil { |
| 60 | t.Fatal(err) |
| 61 | } |
| 62 | exec := agent.New(p, tool.NewRegistry(), agent.NewSession("stable system"), agent.Options{}, event.Discard) |
| 63 | done := make(chan event.Event, 1) |
| 64 | sink := event.FuncSink(func(e event.Event) { |
| 65 | if e.Kind == event.TurnDone { |
| 66 | done <- e |
| 67 | } |
| 68 | }) |
| 69 | c := newOwnedTestController(t, Options{Runner: exec, Executor: exec, Sink: sink, SessionService: service, SessionRuntime: runtime, ExclusiveSession: true}) |
| 70 | if broken { |
| 71 | transport.beforeFailure = func() { |
| 72 | _, turnID, _ := c.currentTurnToken() |
| 73 | c.recordProviderRequest(turnID, provider.RequestObservation{ID: ^uint64(0), Phase: "request_started", |
| 74 | StartedAt: time.Date(10000, 1, 1, 0, 0, 0, 0, time.UTC)}) |
| 75 | } |
| 76 | } |
| 77 | if async { |
| 78 | c.Send("hello") |
| 79 | if terminal := waitTurnDoneEvent(t, done); terminal.Err == nil { |
| 80 | t.Fatal("expected transport failure") |
| 81 | } |
| 82 | waitIdleAdmission(t, c) |
| 83 | } else if err := c.RunTurn(t.Context(), "hello"); err == nil { |
| 84 | t.Fatal("expected transport failure") |
| 85 | } |
| 86 | if transport.calls != 1 { |
| 87 | t.Fatalf("request count=%d", transport.calls) |
| 88 | } |
| 89 | var recovery *provider.InterruptedTurnRecovery |
| 90 | for _, m := range c.sessionEventStore().Snapshot().Projection.Messages { |
| 91 | if m.InterruptedTurn != nil { |
| 92 | recovery = m.InterruptedTurn |
| 93 | } |
| 94 | } |
| 95 | if recovery == nil || recovery.TerminalStatus != "failed" || recovery.FailureDiagnostic == nil || recovery.FailureDiagnostic.TransportCode != "PROTOCOL_ERROR" { |
| 96 | t.Fatalf("cleanup lost failure: %+v", recovery) |
| 97 | } |
| 98 | commits := terminationCommitHistory(t, c) |
| 99 | count, terminals := 0, 0 |
| 100 | for _, commit := range commits { |
| 101 | for _, e := range commit.Events { |
| 102 | if e.Kind == "turn/end" { |
| 103 | terminals++ |
| 104 | } |
| 105 | if e.Kind != "diagnostic/provider" { |
| 106 | continue |
| 107 | } |
| 108 | count++ |
| 109 | if !e.Optional || commit.Events[len(commit.Events)-1].Kind != "turn/end" { |
| 110 | t.Fatal("diagnostic is not optional and atomic with termination") |
| 111 | } |
| 112 | var payload struct { |
| 113 | Failure *provider.FailureDiagnostic |
| 114 | Requests []providerDiagnostic |
| 115 | Dropped *uint64 |
| 116 | Truncated bool |
| 117 | } |
| 118 | if err := json.Unmarshal(e.Payload, &payload); err != nil { |
| 119 | t.Fatal(err) |
| 120 | } |
| 121 | if payload.Failure == nil || payload.Failure.TransportCode != "PROTOCOL_ERROR" || len(payload.Requests) != 1 { |
| 122 | t.Fatalf("incomplete durable evidence: %s", e.Payload) |
| 123 | } |
| 124 | if bytes.Contains(e.Payload, []byte(`"transportError"`)) { |
| 125 | t.Fatal("persisted duplicate transport error prose") |
| 126 | } |
| 127 | if payload.Dropped == nil || *payload.Dropped != 0 || payload.Truncated { |
| 128 | t.Fatalf("fresh turn lost completeness: %s", e.Payload) |
| 129 | } |
| 130 | r := payload.Requests[0] |
| 131 | if r.Host != "provider.test" || r.RequestPath != "/v1/chat/completions" || r.Phase != "request_error" || r.TransportCode != "PROTOCOL_ERROR" || r.RequestBytes <= 0 { |
| 132 | t.Fatalf("request evidence: %+v", r) |
| 133 | } |
| 134 | } |
| 135 | } |
| 136 | wantCount := 1 |
| 137 | if broken { |
| 138 | wantCount = 0 |
| 139 | } |
| 140 | if count != wantCount || terminals != 1 { |
| 141 | t.Fatalf("diagnostic count=%d, terminal count=%d", count, terminals) |
| 142 | } |
| 143 | ref := runtime.Ref() |
| 144 | c.Close() |
| 145 | <-c.closeFinalized |
| 146 | if err := service.CloseAll(context.Background()); err != nil { |
| 147 | t.Fatal(err) |
| 148 | } |
| 149 | snapshot, err := service.Query().CaptureDiagnosticSnapshot(t.Context(), ref) |
| 150 | if err != nil { |
| 151 | t.Fatal(err) |
| 152 | } |
| 153 | var out bytes.Buffer |
| 154 | if err := WriteColdSessionDiagnostics(t.Context(), &out, service.Query(), snapshot, GoalDiagnosticMetadata{}, nil); err != nil { |
| 155 | t.Fatal(err) |
| 156 | } |
| 157 | if !json.Valid(out.Bytes()) || !strings.Contains(out.String(), "PROTOCOL_ERROR") || strings.Contains(out.String(), "diagnostic/provider") == broken || strings.Contains(out.String(), "private-key") { |
| 158 | t.Fatalf("invalid cold export: %s", out.String()) |
| 159 | } |
| 160 | } |
| 161 | |
| 162 | func TestProviderDiagnosticEventIsBoundedAndTurnScoped(t *testing.T) { |
| 163 | c := &Controller{} |
| 164 | c.beginProviderDiagnosticTurn("current") |
| 165 | for id := uint64(1); id <= 130; id++ { |
| 166 | c.recordProviderRequest("current", provider.RequestObservation{ID: id, Phase: "request_started"}) |
| 167 | } |
| 168 | c.recordProviderRequest("other", provider.RequestObservation{ID: 131, Phase: "request_started"}) |
| 169 | e, err := c.providerDiagnosticEvent(event.Event{TurnID: "current", Err: context.Canceled}) |
| 170 | if err != nil { |
| 171 | t.Fatal(err) |
| 172 | } |
| 173 | var payload struct { |
| 174 | Failure *provider.FailureDiagnostic |
| 175 | Requests []providerDiagnostic |
| 176 | Dropped uint64 |
| 177 | Truncated bool |
| 178 | } |
| 179 | if err := json.Unmarshal(e.Payload, &payload); err != nil { |
| 180 | t.Fatal(err) |
| 181 | } |
| 182 | if len(payload.Requests) != 127 || payload.Failure.Kind != provider.FailureKindCancelled || payload.Dropped != 3 || !payload.Truncated { |
| 183 | t.Fatalf("bounded evidence: %+v", payload) |
| 184 | } |
| 185 | for _, request := range payload.Requests { |
| 186 | if request.TurnID != "current" { |
| 187 | t.Fatal("cross-turn evidence") |
| 188 | } |
| 189 | } |
| 190 | if e, err := c.providerDiagnosticEvent(event.Event{TurnID: "unused"}); err != nil || e != nil { |
| 191 | t.Fatal("invented evidence for an unobserved turn") |
| 192 | } |
| 193 | } |
| 194 | |
| 195 | func TestTerminationPreservesFailureWithoutChangingModelContext(t *testing.T) { |
| 196 | for _, status := range []string{"failed", "interrupted", ""} { |
| 197 | t.Run(status, func(t *testing.T) { |
| 198 | recovery := &provider.InterruptedTurnRecovery{Pending: true, TerminalStatus: status} |
| 199 | if status == "failed" { |
| 200 | recovery.FailureDiagnostic = &provider.FailureDiagnostic{Kind: "transport_protocol", TransportCode: "PROTOCOL_ERROR"} |
| 201 | } |
| 202 | user := provider.Message{ID: "user", Role: provider.RoleUser, Content: "hello"} |
| 203 | local := provider.Message{ID: "local", Role: provider.RoleTool, LocalOnly: true, InterruptedTurn: recovery} |
| 204 | before := []provider.Message{user, local} |
| 205 | got := planCancelledMessages(before, 0, user, time.Time{}, func(provider.Message) bool { return true }, nil) |
| 206 | oldBytes, _ := json.Marshal(provider.ModelMessages(before)) |
| 207 | newBytes, _ := json.Marshal(provider.ModelMessages(got)) |
| 208 | if !bytes.Equal(oldBytes, newBytes) { |
| 209 | t.Fatal("diagnostics changed model-visible bytes") |
| 210 | } |
| 211 | after := got[len(got)-1].InterruptedTurn |
| 212 | if after.TerminalStatus != status || (after.FailureDiagnostic == nil) != (recovery.FailureDiagnostic == nil) { |
| 213 | t.Fatalf("lost terminal metadata: %+v", after) |
| 214 | } |
| 215 | if after.FailureDiagnostic != nil { |
| 216 | after.FailureDiagnostic.Kind = "changed" |
| 217 | if recovery.FailureDiagnostic.Kind == "changed" { |
| 218 | t.Fatal("aliased original diagnostic") |
| 219 | } |
| 220 | } |
| 221 | }) |
| 222 | } |
| 223 | } |
| 224 | |
| 225 | func TestProviderDiagnosticEventOmitsErrorProse(t *testing.T) { |
| 226 | for _, err := range []error{ |
| 227 | &provider.RequestFailure{Err: &url.Error{Op: "Post", URL: "https://user:password@provider.test/v1/chat?arbitrary=private-query#private-fragment", Err: http2.ConnectionError(http2.ErrCodeProtocol)}}, |
| 228 | &provider.RequestFailure{Err: errors.New("private free-form error " + strings.Repeat("中文", 1200))}, |
| 229 | &provider.RequestFailure{Err: &provider.APIError{Status: 400, Body: "private request body"}}, |
| 230 | } { |
| 231 | c := &Controller{} |
| 232 | e, encodeErr := c.providerDiagnosticEvent(event.Event{TurnID: "turn", Status: event.TurnFailed, Err: err}) |
| 233 | if encodeErr != nil || e == nil { |
| 234 | t.Fatalf("missing structured evidence: %v", encodeErr) |
| 235 | } |
| 236 | for _, unwanted := range []string{`"transportError"`, "private", "password", "user:", "中文"} { |
| 237 | if bytes.Contains(e.Payload, []byte(unwanted)) { |
| 238 | t.Fatalf("persisted error prose %q: %s", unwanted, e.Payload) |
| 239 | } |
| 240 | } |
| 241 | } |
| 242 | } |
| 243 |