返回 DeepSeek-Reasonix
reconnect_test.go
根目录 / internal / provider / openai / reconnect_test.go
1 package openai
2
3 import (
4 "context"
5 "errors"
6 "io"
7 "net"
8 "net/http"
9 "net/http/httptest"
10 "strings"
11 "sync/atomic"
12 "syscall"
13 "testing"
14 "time"
15
16 "reasonix/internal/provider"
17 )
18
19 // resetAfter answers every request with a 200 SSE head whose body yields the
20 // prelude and then fails the way a peer reset does. A real socket RST cannot
21 // pin the phase: whether net/http sees it before or after the headers is up to
22 // the OS scheduler, and a header-phase reset is retried by design (#8327).
23 func resetAfter(p provider.Provider, prelude string) *atomic.Int32 {
24 var reqs atomic.Int32
25 p.(*client).http = &http.Client{Transport: roundTripFunc(func(r *http.Request) (*http.Response, error) {
26 reqs.Add(1)
27 return &http.Response{
28 StatusCode: http.StatusOK,
29 Header: http.Header{"Content-Type": {"text/event-stream"}},
30 Body: &cutBody{r: strings.NewReader(prelude)},
31 Request: r,
32 }, nil
33 })}
34 return &reqs
35 }
36
37 type roundTripFunc func(*http.Request) (*http.Response, error)
38
39 func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
40
41 type cutBody struct{ r io.Reader }
42
43 func (b *cutBody) Read(p []byte) (int, error) {
44 if n, _ := b.r.Read(p); n > 0 {
45 return n, nil
46 }
47 return 0, &net.OpError{Op: "read", Net: "tcp", Err: syscall.ECONNRESET}
48 }
49
50 func (b *cutBody) Close() error { return nil }
51
52 // TestStreamSurfacesEarlyConnResetAsInterrupt moves body-phase replay to the
53 // Agent: a pre-output connection reset is StreamInterruptedError, not an
54 // in-provider transparent reconnect (avoids stacked retry budgets).
55 func TestStreamSurfacesEarlyConnResetAsInterrupt(t *testing.T) {
56 p, err := New(provider.Config{Name: "deepseek", BaseURL: "http://gateway.invalid", Model: "deepseek-v4", APIKey: "k"})
57 if err != nil {
58 t.Fatalf("New: %v", err)
59 }
60 reqs := resetAfter(p, ": keep-alive\n\n") // a comment line, zero model output
61 ch, err := p.Stream(context.Background(), provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}})
62 if err != nil {
63 t.Fatalf("Stream: %v", err)
64 }
65
66 var gotInterrupted bool
67 for chunk := range ch {
68 if chunk.Type == provider.ChunkError {
69 var interrupted *provider.StreamInterruptedError
70 gotInterrupted = errors.As(chunk.Err, &interrupted)
71 }
72 }
73 if !gotInterrupted {
74 t.Error("early conn reset must surface as StreamInterruptedError for Agent replay")
75 }
76 if n := reqs.Load(); n != 1 {
77 t.Errorf("server saw %d requests, want 1 (no provider body replay)", n)
78 }
79 }
80
81 func TestStreamCancelDoesNotReconnect(t *testing.T) {
82 var reqs atomic.Int32
83 ready := make(chan struct{})
84 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
85 first := reqs.Add(1) == 1
86 w.Header().Set("Content-Type", "text/event-stream")
87 _, _ = io.WriteString(w, ": keep-alive\n\n")
88 flush(w)
89 if first {
90 close(ready)
91 }
92 <-r.Context().Done()
93 }))
94 defer srv.Close()
95
96 p, err := New(provider.Config{Name: "deepseek", BaseURL: srv.URL, Model: "deepseek-v4", APIKey: "k"})
97 if err != nil {
98 t.Fatalf("New: %v", err)
99 }
100 ctx, cancel := context.WithCancel(context.Background())
101 ch, err := p.Stream(ctx, provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}})
102 if err != nil {
103 t.Fatalf("Stream: %v", err)
104 }
105 select {
106 case <-ready:
107 case <-time.After(2 * time.Second):
108 t.Fatal("server did not receive the streaming request")
109 }
110 cancel()
111
112 var got error
113 for chunk := range ch {
114 if chunk.Type == provider.ChunkError {
115 got = chunk.Err
116 }
117 }
118 // Depending on whether the server close or the client watchdog observes
119 // cancellation first, the stream may close silently or surface cancellation.
120 // The contract guarded here is that cancellation never triggers a replay.
121 if got != nil && !errors.Is(got, context.Canceled) {
122 t.Fatalf("stream error = %v, want nil or context.Canceled", got)
123 }
124 if reqs.Load() != 1 {
125 t.Fatalf("cancelled stream reconnected; server saw %d requests, want 1", reqs.Load())
126 }
127 }
128
129 // TestStreamTreatsCleanEOFWithoutDoneAsCut reproduces issue #3953: a proxy that
130 // idle-closes the SSE connection with a clean FIN ends the scan with no error,
131 // which used to commit the turn as complete. Body-phase cuts surface as
132 // StreamInterruptedError so the Agent can replay the frozen request.
133 func TestStreamTreatsCleanEOFWithoutDoneAsCut(t *testing.T) {
134 var reqs int
135 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
136 reqs++
137 w.Header().Set("Content-Type", "text/event-stream")
138 _, _ = io.WriteString(w, ": keep-alive\n\n") // clean close, no [DONE], no finish_reason
139 }))
140 defer srv.Close()
141
142 p, err := New(provider.Config{Name: "deepseek", BaseURL: srv.URL, Model: "deepseek-v4", APIKey: "k"})
143 if err != nil {
144 t.Fatalf("New: %v", err)
145 }
146 ch, err := p.Stream(context.Background(), provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}})
147 if err != nil {
148 t.Fatalf("Stream: %v", err)
149 }
150
151 var gotInterrupted bool
152 for chunk := range ch {
153 switch chunk.Type {
154 case provider.ChunkToolCall:
155 t.Fatalf("incomplete stream must not emit tool calls: %+v", chunk.ToolCall)
156 case provider.ChunkError:
157 var interrupted *provider.StreamInterruptedError
158 gotInterrupted = errors.As(chunk.Err, &interrupted)
159 }
160 }
161 if !gotInterrupted {
162 t.Error("clean EOF before terminal must surface as StreamInterruptedError")
163 }
164 if reqs != 1 {
165 t.Errorf("server saw %d requests, want 1 (no provider body replay)", reqs)
166 }
167 }
168
169 // TestStreamDropsPartialToolCallOnCleanEOF is the post-output half of #3953: the
170 // connection dies mid-tool-call after the call's start was forwarded. The partial
171 // arguments must never surface as a ChunkToolCall; the cut surfaces as a stream
172 // interruption so the agent's recovery path takes over.
173 func TestStreamDropsPartialToolCallOnCleanEOF(t *testing.T) {
174 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
175 w.Header().Set("Content-Type", "text/event-stream")
176 _, _ = io.WriteString(w, "data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"c1\",\"function\":{\"name\":\"bash\",\"arguments\":\"{\"}}]}}]}\n\n")
177 }))
178 defer srv.Close()
179
180 p, err := New(provider.Config{Name: "deepseek", BaseURL: srv.URL, Model: "deepseek-v4", APIKey: "k"})
181 if err != nil {
182 t.Fatalf("New: %v", err)
183 }
184 ch, err := p.Stream(context.Background(), provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}})
185 if err != nil {
186 t.Fatalf("Stream: %v", err)
187 }
188
189 var gotInterrupted bool
190 for chunk := range ch {
191 switch chunk.Type {
192 case provider.ChunkToolCall:
193 t.Fatalf("partial tool call surfaced: %+v", chunk.ToolCall)
194 case provider.ChunkError:
195 var interrupted *provider.StreamInterruptedError
196 gotInterrupted = errors.As(chunk.Err, &interrupted)
197 }
198 }
199 if !gotInterrupted {
200 t.Error("a cut after the tool-call start should surface as a stream interruption")
201 }
202 }
203
204 // TestStreamAcceptsFinishReasonWithoutDone keeps gateways that omit the [DONE]
205 // sentinel working: a finish_reason marks the turn complete on its own.
206 func TestStreamAcceptsFinishReasonWithoutDone(t *testing.T) {
207 srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
208 w.Header().Set("Content-Type", "text/event-stream")
209 _, _ = io.WriteString(w, "data: {\"choices\":[{\"delta\":{\"content\":\"hello\"},\"finish_reason\":\"stop\"}]}\n\n")
210 }))
211 defer srv.Close()
212
213 p, err := New(provider.Config{Name: "deepseek", BaseURL: srv.URL, Model: "deepseek-v4", APIKey: "k"})
214 if err != nil {
215 t.Fatalf("New: %v", err)
216 }
217 ch, err := p.Stream(context.Background(), provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}})
218 if err != nil {
219 t.Fatalf("Stream: %v", err)
220 }
221
222 var text strings.Builder
223 for chunk := range ch {
224 if chunk.Type == provider.ChunkError {
225 t.Fatalf("finish_reason without [DONE] should complete cleanly: %v", chunk.Err)
226 }
227 if chunk.Type == provider.ChunkText {
228 text.WriteString(chunk.Text)
229 }
230 }
231 if text.String() != "hello" {
232 t.Errorf("text = %q, want %q", text.String(), "hello")
233 }
234 }
235
236 // TestStreamDoesNotReplayAfterOutput guards against duplicated output: once a
237 // token has streamed, a mid-stream reset must surface as an error rather than
238 // replaying the request (which would re-emit the already-shown text).
239 func TestStreamDoesNotReplayAfterOutput(t *testing.T) {
240 p, err := New(provider.Config{Name: "deepseek", BaseURL: "http://gateway.invalid", Model: "deepseek-v4", APIKey: "k"})
241 if err != nil {
242 t.Fatalf("New: %v", err)
243 }
244 reqs := resetAfter(p, "data: {\"choices\":[{\"delta\":{\"content\":\"partial\"}}]}\n\n")
245 ch, err := p.Stream(context.Background(), provider.Request{Messages: []provider.Message{{Role: provider.RoleUser, Content: "hi"}}})
246 if err != nil {
247 t.Fatalf("Stream: %v", err)
248 }
249
250 var text strings.Builder
251 var gotErr bool
252 var gotInterrupted bool
253 for chunk := range ch {
254 switch chunk.Type {
255 case provider.ChunkText:
256 text.WriteString(chunk.Text)
257 case provider.ChunkError:
258 gotErr = true
259 var interrupted *provider.StreamInterruptedError
260 gotInterrupted = errors.As(chunk.Err, &interrupted)
261 }
262 }
263 if text.String() != "partial" {
264 t.Errorf("text = %q, want %q (the one delta that streamed)", text.String(), "partial")
265 }
266 if !gotErr {
267 t.Error("a reset after output should surface a ChunkError")
268 }
269 if !gotInterrupted {
270 t.Error("a reset after output should be marked as a stream interruption")
271 }
272 if n := reqs.Load(); n != 1 {
273 t.Errorf("server saw %d requests, want 1 (no replay after output)", n)
274 }
275 }
276
276 lines GO