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:
Wesley Liddick
2026-08-11 14:21:45 +08:00
committed by GitHub
3 changed files with 199 additions and 0 deletions
@@ -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)