返回 DeepSeek-Reasonix
retry_test.go
根目录 / internal / provider / retry_test.go
1 package provider
2
3 import (
4 "context"
5 "errors"
6 "fmt"
7 "io"
8 "net"
9 "net/http"
10 "net/http/httptest"
11 "strings"
12 "sync"
13 "syscall"
14 "testing"
15 "time"
16 )
17
18 type rtFunc func(*http.Request) (*http.Response, error)
19
20 func (f rtFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }
21
22 func statusResp(status int, hdr map[string]string) *http.Response {
23 h := http.Header{}
24 for k, v := range hdr {
25 h.Set(k, v)
26 }
27 return &http.Response{StatusCode: status, Body: io.NopCloser(strings.NewReader("body")), Header: h}
28 }
29
30 func newDummyReq(ctx context.Context) (*http.Request, error) {
31 return http.NewRequestWithContext(ctx, http.MethodPost, "http://x/y", nil)
32 }
33
34 func TestRetryableStatus(t *testing.T) {
35 for _, s := range []int{408, 429, 500, 502, 503, 504, 529, 599} {
36 if !RetryableStatus(s) {
37 t.Errorf("status %d should be retryable", s)
38 }
39 }
40 for _, s := range []int{200, 400, 401, 402, 403, 404, 422} {
41 if RetryableStatus(s) {
42 t.Errorf("status %d should not be retryable", s)
43 }
44 }
45 }
46
47 func TestSendWithRetryCarriesDisplayIdentityAndSanitizedRequestPath(t *testing.T) {
48 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
49 http.NotFound(w, r)
50 }))
51 defer server.Close()
52 _, err := SendWithRetry(context.Background(), server.Client(), SendOptions{
53 Provider: "deepseek-anthropic", ProviderDisplayName: "Deepseek2", Protocol: "openai",
54 }, func(ctx context.Context) (*http.Request, error) {
55 return http.NewRequestWithContext(ctx, http.MethodPost, server.URL+"/anthropic/v1/chat/completions?token=secret", nil)
56 })
57 var apiErr *APIError
58 if !errors.As(err, &apiErr) {
59 t.Fatalf("error = %T %v", err, err)
60 }
61 if apiErr.Provider != "deepseek-anthropic" || apiErr.ProviderDisplayName != "Deepseek2" || apiErr.Protocol != "openai" || apiErr.RequestPath != "/anthropic/v1/chat/completions" {
62 t.Fatalf("API error identity = %+v", apiErr)
63 }
64 if strings.Contains(apiErr.RequestPath, "secret") {
65 t.Fatalf("query leaked into request path: %q", apiErr.RequestPath)
66 }
67 }
68
69 func TestIsConnReset(t *testing.T) {
70 if IsConnReset(nil) {
71 t.Error("nil is not a conn reset")
72 }
73 if IsConnReset(context.Canceled) || IsConnReset(context.DeadlineExceeded) {
74 t.Error("ctx cancel/deadline must not look like a recoverable reset")
75 }
76 if IsConnReset(errors.New("decode stream: invalid character")) {
77 t.Error("a plain protocol error must not be treated as a conn reset")
78 }
79 for _, err := range []error{
80 io.ErrUnexpectedEOF,
81 &net.OpError{Op: "read", Err: syscall.ECONNRESET},
82 fmt.Errorf("read stream: %w", &net.OpError{Op: "read", Err: errors.New("wsarecv: forcibly closed")}),
83 } {
84 if !IsConnReset(err) {
85 t.Errorf("want conn reset for %v", err)
86 }
87 }
88 }
89
90 func TestParseRetryAfterAcceptsHTTPDate(t *testing.T) {
91 resp := &http.Response{Header: http.Header{}}
92 resp.Header.Set("Retry-After", time.Now().Add(30*time.Second).UTC().Format(http.TimeFormat))
93 if d := parseRetryAfter(resp); d < 25*time.Second || d > 31*time.Second {
94 t.Errorf("http-date Retry-After = %v, want ~30s", d)
95 }
96
97 resp.Header.Set("Retry-After", time.Now().Add(-time.Minute).UTC().Format(http.TimeFormat))
98 if d := parseRetryAfter(resp); d != 0 {
99 t.Errorf("elapsed http-date Retry-After = %v, want 0", d)
100 }
101 }
102
103 func TestSendWithRetryFailsFastOnClientErrors(t *testing.T) {
104 for _, status := range []int{400, 402, 422} {
105 calls := 0
106 cl := &http.Client{Transport: rtFunc(func(r *http.Request) (*http.Response, error) {
107 calls++
108 return statusResp(status, nil), nil
109 })}
110 _, err := SendWithRetry(context.Background(), cl, SendOptions{Provider: "p", KeyEnv: "KEY"}, newDummyReq)
111 if calls != 1 {
112 t.Errorf("status %d retried (%d calls), should fail fast", status, calls)
113 }
114 var apiErr *APIError
115 if !errors.As(err, &apiErr) || apiErr.Status != status {
116 t.Errorf("status %d: want *APIError with Status=%d, got %v", status, status, err)
117 }
118 }
119 }
120
121 func TestSendWithRetryPreservesProviderTraceID(t *testing.T) {
122 cl := &http.Client{Transport: rtFunc(func(r *http.Request) (*http.Response, error) {
123 return statusResp(422, map[string]string{"trace_id": "minimax-trace-123"}), nil
124 })}
125 _, err := SendWithRetry(context.Background(), cl, SendOptions{Provider: "minimax-cn-api"}, newDummyReq)
126 var apiErr *APIError
127 if !errors.As(err, &apiErr) {
128 t.Fatalf("want *APIError, got %T: %v", err, err)
129 }
130 if apiErr.TraceID != "minimax-trace-123" {
131 t.Fatalf("TraceID = %q, want minimax-trace-123", apiErr.TraceID)
132 }
133 }
134
135 func TestSendWithRetryAuthError(t *testing.T) {
136 calls := 0
137 cl := &http.Client{Transport: rtFunc(func(r *http.Request) (*http.Response, error) {
138 calls++
139 return statusResp(401, nil), nil
140 })}
141 _, err := SendWithRetry(context.Background(), cl, SendOptions{Provider: "deepseek", KeyEnv: "DEEPSEEK_API_KEY", KeyPresent: true}, newDummyReq)
142 if calls != 1 {
143 t.Errorf("401 retried (%d calls), should fail fast for a never-authed key", calls)
144 }
145 var authErr *AuthError
146 if !errors.As(err, &authErr) || authErr.KeyEnv != "DEEPSEEK_API_KEY" {
147 t.Errorf("want *AuthError naming the key env, got %v", err)
148 }
149 if authErr != nil && authErr.Body != "body" {
150 t.Errorf("AuthError should carry the response body, got %q", authErr.Body)
151 }
152 }
153
154 func TestSendWithRetryKnownKeyFailsWithoutRetry(t *testing.T) {
155 for _, status := range []int{401, 403} {
156 calls := 0
157 cl := &http.Client{Transport: rtFunc(func(*http.Request) (*http.Response, error) { calls++; return statusResp(status, nil), nil })}
158 _, err := SendWithRetry(t.Context(), cl, SendOptions{Provider: "mimo", KeyPresent: true, RetryAuth: true}, newDummyReq)
159 var auth *AuthError
160 if calls != 1 || !errors.As(err, &auth) || auth.Status != status || !auth.HasKey {
161 t.Fatalf("calls=%d err=%v", calls, err)
162 }
163 }
164 }
165
166 // stallingBody sends headers' worth of promise and then never delivers: Read
167 // blocks until Close, mimicking a half-open 502/524 gateway that stalls after
168 // the status line. Close is what the errorBodyReadTimeout timer fires.
169 type stallingBody struct {
170 closeOnce sync.Once
171 closed chan struct{}
172 }
173
174 func newStallingBody() *stallingBody { return &stallingBody{closed: make(chan struct{})} }
175
176 func (b *stallingBody) Read(p []byte) (int, error) {
177 <-b.closed
178 return 0, errors.New("body closed")
179 }
180
181 func (b *stallingBody) Close() error {
182 b.closeOnce.Do(func() { close(b.closed) })
183 return nil
184 }
185
186 // TestSendWithRetryUnblocksStalledErrorBody locks in the #6607 freeze fix: a
187 // failed response whose body never arrives must not wedge error reporting —
188 // the deadline closes the body and the original HTTP failure is returned.
189 // Without the timer in readErrorBody this test hangs on
190 // the first 502 body and fails via the watchdog below.
191 func TestSendWithRetryUnblocksStalledErrorBody(t *testing.T) {
192 prev := errorBodyReadTimeout
193 errorBodyReadTimeout = 50 * time.Millisecond
194 defer func() { errorBodyReadTimeout = prev }()
195
196 calls := 0
197 cl := &http.Client{Transport: rtFunc(func(r *http.Request) (*http.Response, error) {
198 calls++
199 if calls == 1 {
200 return &http.Response{StatusCode: http.StatusBadGateway, Body: newStallingBody(), Header: http.Header{}}, nil
201 }
202 return statusResp(200, nil), nil
203 })}
204
205 type result struct {
206 resp *http.Response
207 err error
208 }
209 done := make(chan result, 1)
210 go func() {
211 resp, err := SendWithRetry(context.Background(), cl, SendOptions{Provider: "p", KeyEnv: "KEY"}, newDummyReq)
212 done <- result{resp, err}
213 }()
214
215 select {
216 case r := <-done:
217 var api *APIError
218 if r.resp != nil || !errors.As(r.err, &api) || api.Status != http.StatusBadGateway || calls != 1 {
219 t.Fatalf("calls=%d err=%v, want original 502 without retry", calls, r.err)
220 }
221 case <-time.After(5 * time.Second):
222 t.Fatal("SendWithRetry wedged on a stalled error body — read deadline did not fire")
223 }
224 }
225
226 func TestSendWithRetryReturnsUpstreamFailureWithoutNotification(t *testing.T) {
227 for _, status := range []int{408, 409, 429, 500, 503, 529} {
228 calls, notifications := 0, 0
229 cl := &http.Client{Transport: rtFunc(func(*http.Request) (*http.Response, error) {
230 calls++
231 return statusResp(status, map[string]string{"Retry-After": "120"}), nil
232 })}
233 ctx, cancel := context.WithCancel(t.Context())
234 ctx = WithRequestAttemptCounter(WithRetryNotify(ctx, func(RetryInfo) { notifications++; cancel() }))
235 _, err := SendWithRetry(ctx, cl, SendOptions{}, newDummyReq)
236 cancel()
237 var api *APIError
238 if !errors.As(err, &api) || api.Status != status || calls != 1 || notifications != 0 || RequestAttemptCount(ctx) != 1 {
239 t.Fatalf("status=%d calls=%d notifications=%d err=%v", status, calls, notifications, err)
240 }
241 if api.RetryAfter != 2*time.Minute {
242 t.Fatalf("lost upstream diagnostic: %+v", api)
243 }
244 }
245 }
246
247 func TestRequestAttemptCountTracksExplicitRequests(t *testing.T) {
248 calls := 0
249 cl := &http.Client{Transport: rtFunc(func(r *http.Request) (*http.Response, error) {
250 calls++
251 if calls < 3 {
252 return statusResp(http.StatusServiceUnavailable, nil), nil
253 }
254 return statusResp(http.StatusBadRequest, nil), nil
255 })}
256 ctx := WithRequestAttemptCounter(context.Background())
257 providerCtx := WithRequestAttemptCounter(ctx)
258
259 for range 3 {
260 if _, err := SendWithRetry(providerCtx, cl, SendOptions{Provider: "p"}, newDummyReq); err == nil {
261 t.Fatal("expected terminal provider error")
262 }
263 }
264 if calls != 3 {
265 t.Fatalf("calls=%d, want one per explicit request", calls)
266 }
267 if got := RequestAttemptCount(ctx); got != 3 {
268 t.Fatalf("request attempt count = %d, want 3", got)
269 }
270 usage := UsageWithRequestAttemptCount(ctx, nil)
271 if usage == nil || usage.TotalTokens != 0 || usage.RequestCount != 3 {
272 t.Fatalf("failed request usage = %+v, want tokens=0 requests=3", usage)
273 }
274 }
275
276 func TestIndependentRequestAttemptCounter(t *testing.T) {
277 parent := WithRequestAttemptCounter(context.Background())
278 recordRequestAttempt(parent)
279 child := WithIndependentRequestAttemptCounter(parent)
280 recordRequestAttempt(child)
281 recordRequestAttempt(child)
282 if RequestAttemptCount(parent) != 1 || RequestAttemptCount(child) != 2 {
283 t.Fatalf("auxiliary and main request counts leaked: parent=%d child=%d", RequestAttemptCount(parent), RequestAttemptCount(child))
284 }
285 }
286
286 lines GO