mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-06 15:03:54 +08:00
Merge pull request #5685 from Monster-DP/main
fix(gateway): 修复池模式同账号错误重试
This commit is contained in:
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user