diff --git a/backend/internal/service/gateway_forward_as_chat_completions.go b/backend/internal/service/gateway_forward_as_chat_completions.go index d03cf53531..7a7ad51f2b 100644 --- a/backend/internal/service/gateway_forward_as_chat_completions.go +++ b/backend/internal/service/gateway_forward_as_chat_completions.go @@ -165,12 +165,14 @@ func (s *GatewayService) ForwardAsChatCompletions( Kind: "failover", Message: upstreamMsg, }) + shouldDisable := false if s.rateLimitService != nil { - s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, mappedModel) + shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, mappedModel) } return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: !shouldDisable && account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } diff --git a/backend/internal/service/gateway_forward_as_responses.go b/backend/internal/service/gateway_forward_as_responses.go index da125eff91..cdd0197eab 100644 --- a/backend/internal/service/gateway_forward_as_responses.go +++ b/backend/internal/service/gateway_forward_as_responses.go @@ -178,12 +178,14 @@ func (s *GatewayService) ForwardAsResponses( Kind: "failover", Message: upstreamMsg, }) + shouldDisable := false if s.rateLimitService != nil { - s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, mappedModel) + shouldDisable = s.rateLimitService.HandleUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody, mappedModel) } return nil, &UpstreamFailoverError{ - StatusCode: resp.StatusCode, - ResponseBody: respBody, + StatusCode: resp.StatusCode, + ResponseBody: respBody, + RetryableOnSameAccount: !shouldDisable && account.IsPoolMode() && account.IsPoolModeRetryableStatus(resp.StatusCode), } } diff --git a/backend/internal/service/gateway_pool_mode_retry_test.go b/backend/internal/service/gateway_pool_mode_retry_test.go new file mode 100644 index 0000000000..83455e5b95 --- /dev/null +++ b/backend/internal/service/gateway_pool_mode_retry_test.go @@ -0,0 +1,78 @@ +package service + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestGatewayCompatPoolMode429AllowsSameAccountRetry(t *testing.T) { + gin.SetMode(gin.TestMode) + + tests := []struct { + name string + path string + body []byte + call func(*GatewayService, context.Context, *gin.Context, *Account, []byte) (*ForwardResult, error) + }{ + { + name: "chat completions", + path: "/v1/chat/completions", + body: []byte(`{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"hello"}]}`), + call: func(svc *GatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsChatCompletions(ctx, c, account, body, nil) + }, + }, + { + name: "responses", + path: "/v1/responses", + body: []byte(`{"model":"claude-sonnet-4-5","input":"hello"}`), + call: func(svc *GatewayService, ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) { + return svc.ForwardAsResponses(ctx, c, account, body, nil) + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + upstream := &queuedHTTPUpstreamStub{responses: []*http.Response{{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"X-Request-Id": []string{"pool-429"}}, + Body: io.NopCloser(http.NoBody), + }}} + svc := &GatewayService{ + cfg: &config.Config{}, + httpUpstream: upstream, + tlsFPProfileService: &TLSFingerprintProfileService{}, + } + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, tt.path, nil) + account := &Account{ + ID: 1, + Name: "pool-account", + Platform: PlatformAnthropic, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "api_key": "test-key", + "pool_mode": true, + }, + } + + result, err := tt.call(svc, context.Background(), c, account, tt.body) + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.True(t, failoverErr.RetryableOnSameAccount) + require.Equal(t, 1, upstream.callCount) + require.Empty(t, recorder.Body.String()) + }) + } +}