mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 15:38:22 +08:00
Merge pull request #5054 from wucm667/fix/issue-5029-openai-passthrough-pool-auth-retry
fix(openai): retry pool auth failures before failover
This commit is contained in:
@@ -1836,6 +1836,37 @@ func (u *openAIHTTPPassthroughFailoverUpstream) calls() []int64 {
|
||||
return append([]int64(nil), u.accountIDs...)
|
||||
}
|
||||
|
||||
type openAIHTTPPassthroughAuthFailoverUpstream struct {
|
||||
service.HTTPUpstream
|
||||
mu sync.Mutex
|
||||
accountIDs []int64
|
||||
statusCode int
|
||||
}
|
||||
|
||||
func (u *openAIHTTPPassthroughAuthFailoverUpstream) Do(_ *http.Request, _ string, accountID int64, _ int) (*http.Response, error) {
|
||||
u.mu.Lock()
|
||||
u.accountIDs = append(u.accountIDs, accountID)
|
||||
u.mu.Unlock()
|
||||
if accountID == 9911 {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"id":"resp_healthy","object":"response","model":"gpt-5.2","status":"completed","usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}`)),
|
||||
}, nil
|
||||
}
|
||||
return &http.Response{
|
||||
StatusCode: u.statusCode,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(`{"error":{"message":"upstream credential rejected"}}`)),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (u *openAIHTTPPassthroughAuthFailoverUpstream) calls() []int64 {
|
||||
u.mu.Lock()
|
||||
defer u.mu.Unlock()
|
||||
return append([]int64(nil), u.accountIDs...)
|
||||
}
|
||||
|
||||
type openAIHTTPPassthroughSSERateLimitUpstream struct {
|
||||
service.HTTPUpstream
|
||||
mu sync.Mutex
|
||||
@@ -2070,6 +2101,108 @@ func TestOpenAIResponses_APIKeyPassthroughPool5xxRetriesThenExhaustsMaxSwitches(
|
||||
require.Equal(t, "Upstream service temporarily unavailable", gjson.GetBytes(rec.Body.Bytes(), "error.message").String())
|
||||
}
|
||||
|
||||
func TestOpenAIResponses_APIKeyPassthroughPoolAuthFailureRetriesThenSwitchesToHealthyAccount(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
}{
|
||||
{name: "401", statusCode: http.StatusUnauthorized},
|
||||
{name: "403", statusCode: http.StatusForbidden},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
groupID := int64(4203)
|
||||
accounts := []service.Account{
|
||||
{
|
||||
ID: 9910, Name: "pool-api-key", Platform: service.PlatformOpenAI,
|
||||
Type: service.AccountTypeAPIKey, Status: service.StatusActive, Schedulable: true, Priority: 1,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "sk-pool",
|
||||
"base_url": "https://api.example.test",
|
||||
"pool_mode": true,
|
||||
"pool_mode_retry_count": float64(1),
|
||||
"pool_mode_retry_status_codes": []any{float64(tt.statusCode)},
|
||||
},
|
||||
Extra: map[string]any{"openai_passthrough": true},
|
||||
},
|
||||
{
|
||||
ID: 9911, Name: "fallback-api-key", Platform: service.PlatformOpenAI,
|
||||
Type: service.AccountTypeAPIKey, Status: service.StatusActive, Schedulable: true, Priority: 2,
|
||||
Credentials: map[string]any{
|
||||
"api_key": "sk-fallback",
|
||||
"base_url": "https://api.example.test",
|
||||
},
|
||||
Extra: map[string]any{"openai_passthrough": true},
|
||||
},
|
||||
}
|
||||
cfg := &config.Config{RunMode: config.RunModeSimple}
|
||||
cfg.Default.RateMultiplier = 1
|
||||
cfg.Security.URLAllowlist.Enabled = false
|
||||
cfg.Gateway.MaxAccountSwitches = 1
|
||||
|
||||
accountRepo := &openAIWSFailoverHandlerAccountRepoStub{accounts: accounts}
|
||||
upstream := &openAIHTTPPassthroughAuthFailoverUpstream{statusCode: tt.statusCode}
|
||||
rateLimitSvc := service.NewRateLimitService(accountRepo, nil, cfg, nil, nil)
|
||||
billingCacheSvc := service.NewBillingCacheService(nil, nil, nil, nil, nil, nil, cfg, nil)
|
||||
t.Cleanup(billingCacheSvc.Stop)
|
||||
gatewaySvc := service.NewOpenAIGatewayService(
|
||||
accountRepo,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
cfg,
|
||||
nil,
|
||||
nil,
|
||||
service.NewBillingService(cfg, nil),
|
||||
rateLimitSvc,
|
||||
billingCacheSvc,
|
||||
upstream,
|
||||
&service.DeferredService{},
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
)
|
||||
h := NewOpenAIGatewayHandler(
|
||||
gatewaySvc,
|
||||
service.NewConcurrencyService(nil),
|
||||
billingCacheSvc,
|
||||
service.NewAPIKeyService(nil, nil, nil, nil, nil, nil, cfg),
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
nil,
|
||||
cfg,
|
||||
)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", strings.NewReader(`{"model":"gpt-5.2","input":"hello","stream":false}`))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
c.Set(string(middleware.ContextKeyAPIKey), &service.APIKey{
|
||||
ID: 1803, GroupID: &groupID,
|
||||
User: &service.User{ID: 1703, Status: service.StatusActive},
|
||||
Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI, Status: service.StatusActive},
|
||||
})
|
||||
c.Set(string(middleware.ContextKeyUser), middleware.AuthSubject{UserID: 1703, Concurrency: 0})
|
||||
|
||||
h.Responses(c)
|
||||
|
||||
require.Equal(t, []int64{9910, 9910, 9911}, upstream.calls())
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
require.Equal(t, "resp_healthy", gjson.GetBytes(rec.Body.Bytes(), "id").String())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIResponses_APIKeyPassthroughSSERateLimitUsesConfiguredPoolRetry(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
groupID := int64(4204)
|
||||
|
||||
@@ -503,6 +503,9 @@ func shouldFailoverOpenAIPassthroughResponse(account *Account, statusCode int, r
|
||||
if isOpenAIRequestBodyTooLargeError(statusCode, "", responseBody) {
|
||||
return true
|
||||
}
|
||||
if account != nil && account.IsPoolMode() && account.IsPoolModeRetryableStatus(statusCode) {
|
||||
return true
|
||||
}
|
||||
switch statusCode {
|
||||
case http.StatusTooManyRequests, 529:
|
||||
return true
|
||||
|
||||
@@ -1676,6 +1676,69 @@ func TestOpenAIGatewayService_APIKeyPassthrough_PoolModeConfigured5xxRetriesSame
|
||||
require.False(t, c.Writer.Written())
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_APIKeyPassthrough_PoolModeAuthErrorsTriggerFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
tests := []struct {
|
||||
name string
|
||||
statusCode int
|
||||
credentials map[string]any
|
||||
}{
|
||||
{
|
||||
name: "configured_401",
|
||||
statusCode: http.StatusUnauthorized,
|
||||
credentials: map[string]any{
|
||||
"pool_mode_retry_status_codes": []any{float64(http.StatusUnauthorized)},
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "default_403",
|
||||
statusCode: http.StatusForbidden,
|
||||
credentials: map[string]any{},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(nil))
|
||||
|
||||
upstreamBody := `{"error":{"message":"upstream credential rejected"}}`
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: &config.Config{Gateway: config.GatewayConfig{ForceCodexCLI: false}},
|
||||
rateLimitService: NewRateLimitService(transientCooldownAccountRepo{}, nil, &config.Config{}, nil, nil),
|
||||
httpUpstream: &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: tt.statusCode,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}},
|
||||
}
|
||||
credentials := map[string]any{
|
||||
"api_key": "sk-test",
|
||||
"base_url": "https://api.example.test",
|
||||
"pool_mode": true,
|
||||
}
|
||||
for key, value := range tt.credentials {
|
||||
credentials[key] = value
|
||||
}
|
||||
account := &Account{
|
||||
ID: 129, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1,
|
||||
Credentials: credentials,
|
||||
Extra: map[string]any{"openai_passthrough": true}, Status: StatusActive, Schedulable: true,
|
||||
}
|
||||
|
||||
_, err := svc.Forward(context.Background(), c, account, []byte(`{"model":"gpt-5.2","input":"hello"}`))
|
||||
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.Equal(t, tt.statusCode, failoverErr.StatusCode)
|
||||
require.True(t, failoverErr.RetryableOnSameAccount)
|
||||
require.False(t, c.Writer.Written(), "pool-mode auth failure must fail over before committing a response")
|
||||
require.False(t, IsResponseCommitted(c))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayService_OpenAIPassthrough_CompactNetworkErrorsTriggerFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user