Merge pull request #4547 from heathermhuang/codex/fix-model-scoped-temp-cooldown-4527

fix(routing): isolate temporary cooldowns by model
This commit is contained in:
Wesley Liddick
2026-07-18 20:41:12 +08:00
committed by GitHub
10 changed files with 528 additions and 68 deletions
@@ -198,17 +198,18 @@ func (s *AntigravityGatewayService) getUpstreamErrorDetail(body []byte) string {
}
// checkErrorPolicy nil 安全的包装
func (s *AntigravityGatewayService) checkErrorPolicy(ctx context.Context, account *Account, statusCode int, body []byte) ErrorPolicyResult {
func (s *AntigravityGatewayService) checkErrorPolicy(ctx context.Context, account *Account, statusCode int, body []byte, requestedModel ...string) ErrorPolicyResult {
if s.rateLimitService == nil {
return ErrorPolicyNone
}
return s.rateLimitService.CheckErrorPolicy(ctx, account, statusCode, body)
return s.rateLimitService.CheckErrorPolicy(ctx, account, statusCode, body, firstRequestedModel(requestedModel))
}
// applyErrorPolicy 应用错误策略结果,返回是否应终止当前循环及应返回的状态码。
// ErrorPolicySkipped 时 outStatus 为 500(前端约定:未命中的错误返回 500)。
func (s *AntigravityGatewayService) applyErrorPolicy(p antigravityRetryLoopParams, statusCode int, headers http.Header, respBody []byte) (handled bool, outStatus int, retErr error) {
switch s.checkErrorPolicy(p.ctx, p.account, statusCode, respBody) {
modelKey := resolveFinalAntigravityModelKey(p.ctx, p.account, p.requestedModel)
switch s.checkErrorPolicy(p.ctx, p.account, statusCode, respBody, modelKey) {
case ErrorPolicySkipped:
if s.handleAntigravityModelRateLimitBeforePolicy(p, statusCode, headers, respBody) {
return true, statusCode, nil
+23 -6
View File
@@ -443,6 +443,9 @@ func TestApplyErrorPolicy(t *testing.T) {
Type: AccountTypeOAuth,
Platform: PlatformAntigravity,
Credentials: map[string]any{
"model_mapping": map[string]any{
"claude-sonnet-4-5": "claude-sonnet-4-5",
},
"temp_unschedulable_enabled": true,
"temp_unschedulable_rules": []any{
map[string]any{
@@ -472,9 +475,10 @@ func TestApplyErrorPolicy(t *testing.T) {
var handleErrorCount int
p := antigravityRetryLoopParams{
ctx: context.Background(),
prefix: "[test]",
account: tt.account,
ctx: context.Background(),
prefix: "[test]",
account: tt.account,
requestedModel: "claude-sonnet-4-5",
handleError: func(ctx context.Context, prefix string, account *Account, statusCode int, headers http.Header, body []byte, requestedModel string, groupID int64, sessionHash string, isStickySession bool) *handleModelRateLimitResult {
handleErrorCount++
return nil
@@ -492,6 +496,9 @@ func TestApplyErrorPolicy(t *testing.T) {
var switchErr *AntigravityAccountSwitchError
require.ErrorAs(t, retErr, &switchErr)
require.Equal(t, tt.account.ID, switchErr.OriginalAccountID)
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
require.Equal(t, "claude-sonnet-4-5", repo.modelRateLimitCalls[0].scope)
} else {
require.NoError(t, retErr)
}
@@ -559,9 +566,10 @@ func TestApplyErrorPolicy_GeminiRateLimitBypassesCustomSkip(t *testing.T) {
type errorPolicyRepoStub struct {
mockAccountRepoForGemini
tempCalls int
setErrCalls int
lastErrorMsg string
tempCalls int
setErrCalls int
lastErrorMsg string
modelRateLimitCalls []modelNotFoundRateLimitCall
}
func (r *errorPolicyRepoStub) SetTempUnschedulable(ctx context.Context, id int64, until time.Time, reason string) error {
@@ -574,3 +582,12 @@ func (r *errorPolicyRepoStub) SetError(ctx context.Context, id int64, errorMsg s
r.lastErrorMsg = errorMsg
return nil
}
func (r *errorPolicyRepoStub) SetModelRateLimit(_ context.Context, id int64, scope string, resetAt time.Time, reason ...string) error {
call := modelNotFoundRateLimitCall{accountID: id, scope: scope, resetAt: resetAt}
if len(reason) > 0 {
call.reason = reason[0]
}
r.modelRateLimitCalls = append(r.modelRateLimitCalls, call)
return nil
}
@@ -17,8 +17,11 @@ import (
type nonJSONTempUnschedAccountRepo struct {
AccountRepository
tempUnschedCalls int
tempReason string
tempUnschedCalls int
tempReason string
modelRateLimitCalls int
modelScope string
modelReason string
}
func (r *nonJSONTempUnschedAccountRepo) SetTempUnschedulable(_ context.Context, _ int64, _ time.Time, reason string) error {
@@ -27,6 +30,15 @@ func (r *nonJSONTempUnschedAccountRepo) SetTempUnschedulable(_ context.Context,
return nil
}
func (r *nonJSONTempUnschedAccountRepo) SetModelRateLimit(_ context.Context, _ int64, scope string, _ time.Time, reason ...string) error {
r.modelRateLimitCalls++
r.modelScope = scope
if len(reason) > 0 {
r.modelReason = reason[0]
}
return nil
}
func TestHandleNonStreamingResponse_NonJSON2xxTriggersFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -131,7 +143,7 @@ func TestHandleNonStreamingResponseAnthropicAPIKeyPassthrough_ValidJSONUnchanged
require.JSONEq(t, string(body), rec.Body.String())
}
func TestHandleNonStreamingResponse_NonJSON2xxMatchesTempUnschedulableRule(t *testing.T) {
func TestHandleNonStreamingResponse_NonJSON2xxMatchesModelScopedTempUnschedulableRule(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
@@ -171,7 +183,9 @@ func TestHandleNonStreamingResponse_NonJSON2xxMatchesTempUnschedulableRule(t *te
require.True(t, errors.As(err, &failoverErr))
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
require.Equal(t, body, failoverErr.ResponseBody)
require.Equal(t, 1, repo.tempUnschedCalls)
require.Contains(t, repo.tempReason, `"status_code":502`)
require.Contains(t, repo.tempReason, `"matched_keyword":"upstream request failed"`)
require.Zero(t, repo.tempUnschedCalls)
require.Equal(t, 1, repo.modelRateLimitCalls)
require.Equal(t, "claude-sonnet-4-6", repo.modelScope)
require.Contains(t, repo.modelReason, `"status_code":502`)
require.Contains(t, repo.modelReason, `"matched_keyword":"upstream request failed"`)
}
@@ -143,7 +143,7 @@ func (s *GeminiMessagesCompatService) forwardClaudeBodyAsChatCompletions(
return nil, s.writeChatCompletionsError(c, http.StatusBadGateway, "upstream_error", "Upstream request failed after retries: "+safeErr)
}
if matched, rebuilt := s.checkErrorPolicyInLoop(ctx, account, resp); matched {
if matched, rebuilt := s.checkErrorPolicyInLoop(ctx, account, resp, mappedModel); matched {
resp = rebuilt
break
} else {
@@ -211,7 +211,13 @@ func (s *GeminiMessagesCompatService) forwardClaudeBodyAsChatCompletions(
if resp.StatusCode >= 400 {
respBody := s.readUpstreamErrorBody(resp)
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
policy := ErrorPolicyNone
if s.rateLimitService != nil {
policy = s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody, mappedModel)
}
if policy != ErrorPolicyTempUnscheduled {
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
evBody := unwrapIfNeeded(account.Type == AccountTypeOAuth, respBody)
if s.shouldFailoverGeminiUpstreamError(resp.StatusCode) {
@@ -199,6 +199,7 @@ func TestGeminiErrorPolicyIntegration(t *testing.T) {
expectFailover bool // expect UpstreamFailoverError
expectHandleError bool // expect handleGeminiUpstreamError to be called
expectShouldFailover bool // for None path, whether shouldFailover triggers
expectModelScope string
}{
{
name: "custom_codes_matched_429_failover",
@@ -252,7 +253,8 @@ func TestGeminiErrorPolicyIntegration(t *testing.T) {
statusCode: 503,
respBody: []byte(`overloaded`),
expectFailover: true,
expectHandleError: true,
expectHandleError: false,
expectModelScope: "gemini-2.5-pro",
},
{
name: "no_policy_429_failover_via_shouldFailover",
@@ -306,17 +308,22 @@ func TestGeminiErrorPolicyIntegration(t *testing.T) {
headers := http.Header{}
if svc.rateLimitService != nil {
switch svc.rateLimitService.CheckErrorPolicy(ctx, account, statusCode, respBody) {
policy := svc.rateLimitService.CheckErrorPolicy(ctx, account, statusCode, respBody, "gemini-2.5-pro")
switch policy {
case ErrorPolicySkipped:
// Skipped → return error directly (no handleGeminiUpstreamError, no failover)
gotFailover = false
handleErrorCalled = false
goto verify
case ErrorPolicyMatched, ErrorPolicyTempUnscheduled:
case ErrorPolicyMatched:
svc.handleGeminiUpstreamError(ctx, account, statusCode, headers, respBody)
handleErrorCalled = true
gotFailover = true
goto verify
case ErrorPolicyTempUnscheduled:
handleErrorCalled = false
gotFailover = true
goto verify
}
}
@@ -330,6 +337,12 @@ func TestGeminiErrorPolicyIntegration(t *testing.T) {
verify:
require.Equal(t, tt.expectFailover, gotFailover, "failover mismatch")
require.Equal(t, tt.expectHandleError, handleErrorCalled, "handleGeminiUpstreamError call mismatch")
if tt.expectModelScope != "" {
require.Equal(t, 1, repo.setModelRateLimitedCalls)
require.Equal(t, tt.expectModelScope, repo.lastModelScope)
require.Zero(t, repo.setTempCalls)
require.Zero(t, repo.setRateLimitedCalls, "model temp rule must not be widened into an account rate limit")
}
if tt.expectShouldFailover {
require.True(t, svc.shouldFailoverGeminiUpstreamError(statusCode),
@@ -416,9 +429,11 @@ func TestHandleGeminiUpstreamError_GoogleOneCapacityExhaustedUsesTierCooldown(t
type geminiErrorPolicyRepo struct {
mockAccountRepoForGemini
setErrorCalls int
setRateLimitedCalls int
setTempCalls int
setErrorCalls int
setRateLimitedCalls int
setTempCalls int
setModelRateLimitedCalls int
lastModelScope string
}
func (r *geminiErrorPolicyRepo) SetError(_ context.Context, _ int64, _ string) error {
@@ -435,3 +450,9 @@ func (r *geminiErrorPolicyRepo) SetTempUnschedulable(_ context.Context, _ int64,
r.setTempCalls++
return nil
}
func (r *geminiErrorPolicyRepo) SetModelRateLimit(_ context.Context, _ int64, scope string, _ time.Time, _ ...string) error {
r.setModelRateLimitedCalls++
r.lastModelScope = scope
return nil
}
@@ -867,7 +867,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
}
// 错误策略优先:匹配则跳过重试直接处理。
if matched, rebuilt := s.checkErrorPolicyInLoop(ctx, account, resp); matched {
if matched, rebuilt := s.checkErrorPolicyInLoop(ctx, account, resp, mappedModel); matched {
resp = rebuilt
break
} else {
@@ -937,7 +937,8 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
respBody := s.readUpstreamErrorBody(resp)
// 统一错误策略:自定义错误码 + 临时不可调度
if s.rateLimitService != nil {
switch s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody) {
policy := s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody, mappedModel)
switch policy {
case ErrorPolicySkipped:
upstreamReqID := resp.Header.Get(requestIDHeader)
if upstreamReqID == "" {
@@ -945,7 +946,9 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
}
return nil, s.writeGeminiMappedError(c, account, http.StatusInternalServerError, upstreamReqID, respBody)
case ErrorPolicyMatched, ErrorPolicyTempUnscheduled:
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if policy == ErrorPolicyMatched {
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
upstreamReqID := resp.Header.Get(requestIDHeader)
if upstreamReqID == "" {
upstreamReqID = resp.Header.Get("x-goog-request-id")
@@ -1336,7 +1339,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
}
// 错误策略优先:匹配则跳过重试直接处理。
if matched, rebuilt := s.checkErrorPolicyInLoop(ctx, account, resp); matched {
if matched, rebuilt := s.checkErrorPolicyInLoop(ctx, account, resp, mappedModel); matched {
resp = rebuilt
break
} else {
@@ -1445,7 +1448,8 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
// 统一错误策略:自定义错误码 + 临时不可调度
if s.rateLimitService != nil {
switch s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody) {
policy := s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, respBody, mappedModel)
switch policy {
case ErrorPolicySkipped:
respBody = unwrapIfNeeded(isOAuth, respBody)
contentType := resp.Header.Get("Content-Type")
@@ -1456,7 +1460,9 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
c.Data(http.StatusInternalServerError, contentType, respBody)
return nil, fmt.Errorf("gemini upstream error: %d (skipped by error policy)", resp.StatusCode)
case ErrorPolicyMatched, ErrorPolicyTempUnscheduled:
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
if policy == ErrorPolicyMatched {
s.handleGeminiUpstreamError(ctx, account, resp.StatusCode, resp.Header, respBody)
}
evBody := unwrapIfNeeded(isOAuth, respBody)
upstreamMsg := strings.TrimSpace(extractUpstreamErrorMessage(evBody))
upstreamMsg = sanitizeUpstreamErrorMessage(upstreamMsg)
@@ -1631,7 +1637,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
// 返回 true 表示策略已匹配(调用者应 break),resp 已重建可直接使用。
// 返回 false 表示 ErrorPolicyNone,resp 已重建,调用者继续走重试逻辑。
func (s *GeminiMessagesCompatService) checkErrorPolicyInLoop(
ctx context.Context, account *Account, resp *http.Response,
ctx context.Context, account *Account, resp *http.Response, mappedModel string,
) (matched bool, rebuilt *http.Response) {
if resp.StatusCode < 400 || s.rateLimitService == nil {
return false, resp
@@ -1643,7 +1649,7 @@ func (s *GeminiMessagesCompatService) checkErrorPolicyInLoop(
Header: resp.Header.Clone(),
Body: io.NopCloser(bytes.NewReader(body)),
}
policy := s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, body)
policy := s.rateLimitService.CheckErrorPolicy(ctx, account, resp.StatusCode, body, mappedModel)
return policy != ErrorPolicyNone, rebuilt
}
@@ -62,17 +62,30 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont
return false
}
if s == nil || account == nil {
return false
}
stateCtx = withTempUnschedulableModel(stateCtx, canonicalModel)
if s.rateLimitService != nil && len(canonicalModel) > 0 && s.rateLimitService.HandleUpstreamModelNotFound(stateCtx, account, canonicalModel[0], statusCode, responseBody) {
return true
}
// Isolate a custom temporary-unschedulable match to the known upstream
// model before entering the generic account error path. This keeps the
// account available to other models and avoids the account runtime blocker.
if s.rateLimitService != nil && statusCode != http.StatusUnauthorized && len(canonicalModel) > 0 && strings.TrimSpace(canonicalModel[0]) != "" &&
s.rateLimitService.HandleTempUnschedulable(stateCtx, account, statusCode, responseBody, canonicalModel[0]) {
return true
}
if statusCode == http.StatusTooManyRequests {
s.markOpenAIOAuth429RateLimited(stateCtx, account, headers, responseBody)
}
if s == nil || account == nil || s.rateLimitService == nil {
if s.rateLimitService == nil {
return false
}
if len(canonicalModel) > 0 && s.rateLimitService.HandleUpstreamModelNotFound(stateCtx, account, canonicalModel[0], statusCode, responseBody) {
return true
}
shouldDisable := s.rateLimitService.HandleUpstreamError(stateCtx, account, statusCode, headers, responseBody)
if shouldDisable {
modelTempMatched := statusCode != http.StatusUnauthorized && tempUnschedulableModel(stateCtx, nil) != "" &&
len(matchTempUnschedulableRules(account, statusCode, responseBody)) > 0
if shouldDisable && !modelTempMatched {
s.BlockAccountScheduling(account, time.Time{}, "upstream_disable")
}
if !shouldDisable && account.Platform == PlatformOpenAI && account.Type == AccountTypeAPIKey && shouldCooldownOpenAITransientUpstreamError(statusCode, responseBody) {
@@ -4,6 +4,7 @@ package service
import (
"context"
"errors"
"net/http"
"testing"
"time"
@@ -157,6 +158,131 @@ func TestOpenAIModelNotFound_DoesNotRuntimeBlockWholeAccount(t *testing.T) {
require.Len(t, repo.modelRateLimitCalls, 1)
}
func TestOpenAIModelTempUnschedulable_DoesNotRuntimeBlockWholeAccount(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &OpenAIGatewayService{
rateLimitService: &RateLimitService{accountRepo: repo},
}
account := openAIModelNotFoundTempAccount()
shouldDisable := svc.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
"gpt-5.4",
)
require.True(t, shouldDisable)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
require.Equal(t, "gpt-5.4", repo.modelRateLimitCalls[0].scope)
}
func TestOpenAIModelTempUnschedulable_WriteFailureDoesNotRuntimeBlockWholeAccount(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{modelRateLimitErr: errors.New("write failed")}
svc := &OpenAIGatewayService{
rateLimitService: &RateLimitService{accountRepo: repo},
}
account := openAIModelNotFoundTempAccount()
shouldDisable := svc.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
"gpt-5.4",
)
require.True(t, shouldDisable)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
}
func TestOpenAIOAuth429_MatchingModelTempRuleAvoidsAccountRuntimeBlock(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &OpenAIGatewayService{
rateLimitService: &RateLimitService{accountRepo: repo},
}
account := openAIModelNotFoundTempAccount()
account.Type = AccountTypeOAuth
account.Credentials["temp_unschedulable_rules"] = []any{
map[string]any{
"error_code": float64(http.StatusTooManyRequests),
"keywords": []any{"model quota"},
"duration_minutes": float64(10),
},
}
shouldDisable := svc.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusTooManyRequests,
http.Header{},
[]byte(`{"error":{"message":"model quota exhausted"}}`),
"gpt-5.4",
)
require.True(t, shouldDisable)
require.False(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.Len(t, repo.modelRateLimitCalls, 1)
require.Equal(t, "gpt-5.4", repo.modelRateLimitCalls[0].scope)
}
func TestOpenAIOAuth429_NonmatchingModelTempRuleKeepsAccountRuntimeBlock(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &OpenAIGatewayService{
rateLimitService: &RateLimitService{accountRepo: repo},
}
account := openAIModelNotFoundTempAccount()
account.Type = AccountTypeOAuth
account.Credentials["temp_unschedulable_rules"] = []any{
map[string]any{
"error_code": float64(http.StatusTooManyRequests),
"keywords": []any{"different marker"},
"duration_minutes": float64(10),
},
}
shouldDisable := svc.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusTooManyRequests,
http.Header{},
[]byte(`{"error":{"message":"global rate limit"}}`),
"gpt-5.4",
)
require.False(t, shouldDisable)
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.Empty(t, repo.modelRateLimitCalls)
}
func TestOpenAITempUnschedulable_UnknownModelKeepsAccountRuntimeBlock(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &OpenAIGatewayService{
rateLimitService: &RateLimitService{accountRepo: repo},
}
account := openAIModelNotFoundTempAccount()
shouldDisable := svc.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
)
require.True(t, shouldDisable)
require.True(t, svc.isOpenAIAccountRuntimeBlocked(account))
require.Equal(t, 1, repo.tempCalls)
require.Empty(t, repo.modelRateLimitCalls)
}
func TestOpenAIRuntimeBlock_DoesNotShortenExistingBlock(t *testing.T) {
svc := &OpenAIGatewayService{}
account := &Account{ID: 46, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
+94 -31
View File
@@ -150,7 +150,8 @@ const (
// CheckErrorPolicy 检查自定义错误码和临时不可调度规则。
// 自定义错误码开启时覆盖后续所有逻辑(包括临时不可调度)。
func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Account, statusCode int, responseBody []byte) ErrorPolicyResult {
func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Account, statusCode int, responseBody []byte, requestedModel ...string) ErrorPolicyResult {
ctx = withTempUnschedulableModel(ctx, requestedModel)
if account.IsCustomErrorCodesEnabled() {
if account.ShouldHandleErrorCode(statusCode) {
return ErrorPolicyMatched
@@ -166,7 +167,7 @@ func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Accoun
}
return ErrorPolicySkipped
}
if s.tryTempUnschedulable(ctx, account, statusCode, responseBody) {
if s.tryTempUnschedulable(ctx, account, statusCode, responseBody, firstRequestedModel(requestedModel)) {
return ErrorPolicyTempUnscheduled
}
return ErrorPolicyNone
@@ -175,6 +176,7 @@ func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Accoun
// HandleUpstreamError 处理上游错误响应,标记账号状态
// 返回是否应该停止该账号的调度
func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte, requestedModel ...string) (shouldDisable bool) {
ctx = withTempUnschedulableModel(ctx, requestedModel)
customErrorCodesEnabled := account.IsCustomErrorCodesEnabled()
// 池模式默认不标记本地账号状态;但管理员显式配置的临时不可调度规则优先。
@@ -216,7 +218,7 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc
// 先尝试临时不可调度规则(401除外)
// 如果匹配成功,直接返回,不执行后续禁用逻辑
if statusCode != 401 {
if s.tryTempUnschedulable(ctx, account, statusCode, responseBody) {
if s.tryTempUnschedulable(ctx, account, statusCode, responseBody, firstRequestedModel(requestedModel)) {
return true
}
}
@@ -1900,14 +1902,18 @@ func (s *RateLimitService) GetTempUnschedStatus(ctx context.Context, accountID i
return state, nil
}
func (s *RateLimitService) HandleTempUnschedulable(ctx context.Context, account *Account, statusCode int, responseBody []byte) bool {
func (s *RateLimitService) HandleTempUnschedulable(ctx context.Context, account *Account, statusCode int, responseBody []byte, requestedModel ...string) bool {
if account == nil {
return false
}
if account.IsPoolMode() && !account.IsCustomErrorCodesEnabled() {
return false
}
if !account.ShouldHandleErrorCode(statusCode) {
return false
}
return s.tryTempUnschedulable(ctx, account, statusCode, responseBody)
ctx = withTempUnschedulableModel(ctx, requestedModel)
return s.tryTempUnschedulable(ctx, account, statusCode, responseBody, firstRequestedModel(requestedModel))
}
func (s *RateLimitService) HandleOpenAIImageRateLimit(ctx context.Context, account *Account, statusCode int, headers http.Header, responseBody []byte) bool {
@@ -2075,7 +2081,71 @@ func modelRateLimitKeyForUpstreamModelNotFound(ctx context.Context, account *Acc
return modelKey
}
func (s *RateLimitService) tryTempUnschedulable(ctx context.Context, account *Account, statusCode int, responseBody []byte) bool {
func firstRequestedModel(requestedModel []string) string {
if len(requestedModel) == 0 {
return ""
}
return strings.TrimSpace(requestedModel[0])
}
type tempUnschedulableModelContextKey struct{}
func withTempUnschedulableModel(ctx context.Context, requestedModel []string) context.Context {
model := firstRequestedModel(requestedModel)
if model == "" {
return ctx
}
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, tempUnschedulableModelContextKey{}, model)
}
func tempUnschedulableModel(ctx context.Context, requestedModel []string) string {
if model := firstRequestedModel(requestedModel); model != "" {
return model
}
if ctx == nil {
return ""
}
model, _ := ctx.Value(tempUnschedulableModelContextKey{}).(string)
return strings.TrimSpace(model)
}
type tempUnschedulableRuleMatch struct {
rule TempUnschedulableRule
ruleIndex int
matchedKeyword string
}
func matchTempUnschedulableRules(account *Account, statusCode int, responseBody []byte) []tempUnschedulableRuleMatch {
if account == nil || !account.IsTempUnschedulableEnabled() || statusCode <= 0 || len(responseBody) == 0 {
return nil
}
rules := account.GetTempUnschedulableRules()
if len(rules) == 0 {
return nil
}
body := responseBody
if len(body) > tempUnschedBodyMaxBytes {
body = body[:tempUnschedBodyMaxBytes]
}
bodyLower := strings.ToLower(string(body))
matches := make([]tempUnschedulableRuleMatch, 0, 1)
for idx, rule := range rules {
if rule.ErrorCode != statusCode || len(rule.Keywords) == 0 {
continue
}
matchedKeyword := matchTempUnschedKeyword(bodyLower, rule.Keywords)
if matchedKeyword == "" {
continue
}
matches = append(matches, tempUnschedulableRuleMatch{rule: rule, ruleIndex: idx, matchedKeyword: matchedKeyword})
}
return matches
}
func (s *RateLimitService) tryTempUnschedulable(ctx context.Context, account *Account, statusCode int, responseBody []byte, requestedModel ...string) bool {
if account == nil {
return false
}
@@ -2099,30 +2169,8 @@ func (s *RateLimitService) tryTempUnschedulable(ctx context.Context, account *Ac
return false
}
}
rules := account.GetTempUnschedulableRules()
if len(rules) == 0 {
return false
}
if statusCode <= 0 || len(responseBody) == 0 {
return false
}
body := responseBody
if len(body) > tempUnschedBodyMaxBytes {
body = body[:tempUnschedBodyMaxBytes]
}
bodyLower := strings.ToLower(string(body))
for idx, rule := range rules {
if rule.ErrorCode != statusCode || len(rule.Keywords) == 0 {
continue
}
matchedKeyword := matchTempUnschedKeyword(bodyLower, rule.Keywords)
if matchedKeyword == "" {
continue
}
if s.triggerTempUnschedulable(ctx, account, rule, idx, statusCode, matchedKeyword, responseBody) {
for _, match := range matchTempUnschedulableRules(account, statusCode, responseBody) {
if s.triggerTempUnschedulable(ctx, account, match.rule, match.ruleIndex, statusCode, match.matchedKeyword, responseBody, tempUnschedulableModel(ctx, requestedModel)) {
return true
}
}
@@ -2162,7 +2210,7 @@ func matchTempUnschedKeyword(bodyLower string, keywords []string) string {
return ""
}
func (s *RateLimitService) triggerTempUnschedulable(ctx context.Context, account *Account, rule TempUnschedulableRule, ruleIndex int, statusCode int, matchedKeyword string, responseBody []byte) bool {
func (s *RateLimitService) triggerTempUnschedulable(ctx context.Context, account *Account, rule TempUnschedulableRule, ruleIndex int, statusCode int, matchedKeyword string, responseBody []byte, requestedModel ...string) bool {
if account == nil {
return false
}
@@ -2190,6 +2238,21 @@ func (s *RateLimitService) triggerTempUnschedulable(ctx context.Context, account
reason = strings.TrimSpace(state.ErrorMessage)
}
// Persist known-model failures under the model key so the scheduler excludes
// only this (account, model) pair. Authentication and model-unknown failures
// retain the legacy account-wide temporary-unschedulable behavior below.
modelKey := firstRequestedModel(requestedModel)
if modelKey != "" && statusCode != http.StatusUnauthorized {
if err := s.accountRepo.SetModelRateLimit(ctx, account.ID, modelKey, until, reason); err != nil {
slog.Warn("temp_unsched_model_rate_limit_set_failed", "account_id", account.ID, "model", modelKey, "error", err)
// The rule matched, so fail over the current request even if persistence
// failed; never widen a model-scoped failure into an account-wide block.
return true
}
slog.Info("account_model_temp_unschedulable", "account_id", account.ID, "model", modelKey, "until", until, "rule_index", ruleIndex, "status_code", statusCode)
return true
}
s.notifyAccountSchedulingBlocked(account, until, "temp_unschedulable")
if err := s.accountRepo.SetTempUnschedulable(ctx, account.ID, until, reason); err != nil {
slog.Warn("temp_unsched_set_failed", "account_id", account.ID, "error", err)
@@ -4,6 +4,7 @@ package service
import (
"context"
"encoding/json"
"errors"
"net/http"
"testing"
@@ -87,7 +88,7 @@ func TestRateLimitService_HandleUpstreamError_ModelNotFoundWriteFailureDoesNotTe
require.Len(t, repo.modelRateLimitCalls, 1)
}
func TestRateLimitService_HandleUpstreamError_Bare404KeepsTempUnschedulablePath(t *testing.T) {
func TestRateLimitService_HandleUpstreamError_Bare404UsesModelScopedTempUnschedulableWhenModelKnown(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
@@ -101,11 +102,203 @@ func TestRateLimitService_HandleUpstreamError_Bare404KeepsTempUnschedulablePath(
"gpt-5.4",
)
require.True(t, handled)
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
call := repo.modelRateLimitCalls[0]
require.Equal(t, account.ID, call.accountID)
require.Equal(t, "gpt-5.4", call.scope)
require.WithinDuration(t, time.Now().Add(10*time.Minute), call.resetAt, 5*time.Second)
var state TempUnschedState
require.NoError(t, json.Unmarshal([]byte(call.reason), &state))
require.Equal(t, http.StatusNotFound, state.StatusCode)
require.Equal(t, "not found", state.MatchedKeyword)
}
func TestRateLimitService_HandleUpstreamError_Bare404WithoutModelKeepsAccountTempUnschedulable(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
handled := svc.HandleUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
)
require.True(t, handled)
require.Equal(t, 1, repo.tempCalls)
require.Empty(t, repo.modelRateLimitCalls)
}
func TestRateLimitService_HandleUpstreamError_ModelTempWriteFailureNeverWidensToAccount(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{modelRateLimitErr: errors.New("write failed")}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
handled := svc.HandleUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
"gpt-5.4",
)
require.True(t, handled)
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
}
func TestRateLimitService_HandleTempUnschedulable_PoolModeWithoutCustomPolicySkipsState(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
account.Credentials["pool_mode"] = true
handled := svc.HandleTempUnschedulable(
context.Background(),
account,
http.StatusNotFound,
[]byte(`{"error":{"message":"endpoint not found"}}`),
"gpt-5.4",
)
require.False(t, handled)
require.Zero(t, repo.tempCalls)
require.Empty(t, repo.modelRateLimitCalls)
}
func TestRateLimitService_HandleTempUnschedulable_PoolModeCustomPolicyUsesModelScope(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
account.Credentials["pool_mode"] = true
account.Credentials["custom_error_codes_enabled"] = true
account.Credentials["custom_error_codes"] = []any{float64(http.StatusNotFound)}
handled := svc.HandleTempUnschedulable(
context.Background(),
account,
http.StatusNotFound,
[]byte(`{"error":{"message":"endpoint not found"}}`),
"gpt-5.4",
)
require.True(t, handled)
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
require.Equal(t, "gpt-5.4", repo.modelRateLimitCalls[0].scope)
}
func TestRateLimitService_TempUnschedulableContextPreservesModelForPoolDependency(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
account.Credentials["pool_mode"] = true
ctx := withTempUnschedulableModel(context.Background(), []string{"gpt-5.4"})
// #4496 calls tryTempUnschedulable from the pool-mode branch without an
// explicit model argument. The request context must preserve the canonical
// model so that combined behavior remains model-scoped after it lands.
handled := svc.tryTempUnschedulable(
ctx,
account,
http.StatusNotFound,
[]byte(`{"error":{"message":"endpoint not found"}}`),
)
require.True(t, handled)
require.Zero(t, repo.tempCalls)
require.Len(t, repo.modelRateLimitCalls, 1)
require.Equal(t, "gpt-5.4", repo.modelRateLimitCalls[0].scope)
}
func TestRateLimitService_HandleUpstreamError_CustomPolicyExclusionSkipsAllState(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
account.Credentials["custom_error_codes_enabled"] = true
account.Credentials["custom_error_codes"] = []any{float64(http.StatusServiceUnavailable)}
handled := svc.HandleUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
"gpt-5.4",
)
require.False(t, handled)
require.Zero(t, repo.tempCalls)
require.Empty(t, repo.modelRateLimitCalls)
}
func TestRateLimitService_HandleTempUnschedulable_AuthenticationFailureStaysAccountScoped(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
account.TempUnschedulableReason = "legacy non-JSON reason"
account.Credentials["temp_unschedulable_rules"] = []any{
map[string]any{
"error_code": float64(http.StatusUnauthorized),
"keywords": []any{"unauthorized"},
"duration_minutes": float64(10),
},
}
handled := svc.HandleTempUnschedulable(
context.Background(),
account,
http.StatusUnauthorized,
[]byte(`{"error":{"message":"unauthorized"}}`),
"gpt-5.4",
)
require.True(t, handled)
require.Equal(t, 1, repo.tempCalls)
require.Empty(t, repo.modelRateLimitCalls)
}
func TestRateLimitService_ModelTempUnschedulableIsolatesSchedulerByModel(t *testing.T) {
repo := &modelNotFoundAccountRepoStub{}
svc := &RateLimitService{accountRepo: repo}
account := openAIModelNotFoundTempAccount()
account.Credentials["model_mapping"] = map[string]any{
"public-a": "upstream-a",
"upstream-a": "upstream-b",
}
handled := svc.HandleUpstreamError(
context.Background(),
account,
http.StatusNotFound,
http.Header{},
[]byte(`{"error":{"message":"endpoint not found"}}`),
"upstream-a",
)
require.True(t, handled)
require.Len(t, repo.modelRateLimitCalls, 1)
call := repo.modelRateLimitCalls[0]
require.Equal(t, "upstream-a", call.scope, "canonical upstream model must not be mapped a second time")
account.Extra = map[string]any{
modelRateLimitsKey: map[string]any{
call.scope: map[string]any{
"rate_limit_reset_at": call.resetAt.UTC().Format(time.RFC3339),
},
},
}
require.False(t, account.IsSchedulableForModelWithContext(context.Background(), "public-a"))
require.True(t, account.IsSchedulableForModelWithContext(context.Background(), "gpt-5.6-sol"))
}
func openAIModelNotFoundTempAccount() *Account {
return &Account{
ID: 101,