返回 DeepSeek-Reasonix
provider_diagnostic_event_test.go
根目录 / internal / control / provider_diagnostic_event_test.go
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
243 lines GO