From ef2a22be70d0b42ea10556d5bf75c6ee9ba1d7b3 Mon Sep 17 00:00:00 2001 From: Heatherm Huang Date: Sat, 18 Jul 2026 15:30:15 +0800 Subject: [PATCH] fix(routing): isolate temporary cooldowns by model --- .../service/antigravity_gateway_service.go | 7 +- backend/internal/service/error_policy_test.go | 29 ++- .../gateway_non_streaming_response_test.go | 26 ++- .../gemini_chat_completions_compat_service.go | 10 +- .../service/gemini_error_policy_test.go | 33 ++- .../service/gemini_messages_compat_service.go | 22 +- .../openai_account_runtime_block_fastpath.go | 23 ++- ...nai_account_runtime_block_fastpath_test.go | 126 +++++++++++ backend/internal/service/ratelimit_service.go | 125 ++++++++--- .../ratelimit_service_model_not_found_test.go | 195 +++++++++++++++++- 10 files changed, 528 insertions(+), 68 deletions(-) diff --git a/backend/internal/service/antigravity_gateway_service.go b/backend/internal/service/antigravity_gateway_service.go index 9e0cc804d..b835a8410 100644 --- a/backend/internal/service/antigravity_gateway_service.go +++ b/backend/internal/service/antigravity_gateway_service.go @@ -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 diff --git a/backend/internal/service/error_policy_test.go b/backend/internal/service/error_policy_test.go index 2aa7a421d..35bebb94e 100644 --- a/backend/internal/service/error_policy_test.go +++ b/backend/internal/service/error_policy_test.go @@ -333,6 +333,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{ @@ -362,9 +365,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 @@ -382,6 +386,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) } @@ -449,9 +456,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 { @@ -464,3 +472,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 +} diff --git a/backend/internal/service/gateway_non_streaming_response_test.go b/backend/internal/service/gateway_non_streaming_response_test.go index 2416e3e0d..a812a62e4 100644 --- a/backend/internal/service/gateway_non_streaming_response_test.go +++ b/backend/internal/service/gateway_non_streaming_response_test.go @@ -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"`) } diff --git a/backend/internal/service/gemini_chat_completions_compat_service.go b/backend/internal/service/gemini_chat_completions_compat_service.go index 2dc71934c..7db000d26 100644 --- a/backend/internal/service/gemini_chat_completions_compat_service.go +++ b/backend/internal/service/gemini_chat_completions_compat_service.go @@ -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) { diff --git a/backend/internal/service/gemini_error_policy_test.go b/backend/internal/service/gemini_error_policy_test.go index 84f9a706b..deae912c9 100644 --- a/backend/internal/service/gemini_error_policy_test.go +++ b/backend/internal/service/gemini_error_policy_test.go @@ -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 +} diff --git a/backend/internal/service/gemini_messages_compat_service.go b/backend/internal/service/gemini_messages_compat_service.go index a50a7fc80..2aaba767e 100644 --- a/backend/internal/service/gemini_messages_compat_service.go +++ b/backend/internal/service/gemini_messages_compat_service.go @@ -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 } diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index 7b3df522f..6b68b4e6f 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -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) { diff --git a/backend/internal/service/openai_account_runtime_block_fastpath_test.go b/backend/internal/service/openai_account_runtime_block_fastpath_test.go index 7d7cfbda5..dc1de42f1 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath_test.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath_test.go @@ -4,6 +4,7 @@ package service import ( "context" + "errors" "net/http" "testing" "time" @@ -105,6 +106,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} diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 507e02d87..1d53f3791 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -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 @@ -161,7 +162,7 @@ func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Accoun if account.IsPoolMode() { return ErrorPolicySkipped } - if s.tryTempUnschedulable(ctx, account, statusCode, responseBody) { + if s.tryTempUnschedulable(ctx, account, statusCode, responseBody, firstRequestedModel(requestedModel)) { return ErrorPolicyTempUnscheduled } return ErrorPolicyNone @@ -170,6 +171,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() // 池模式默认不标记本地账号状态;仅当用户显式配置自定义错误码时按本地策略处理。 @@ -207,7 +209,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 } } @@ -1891,14 +1893,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 { @@ -2066,7 +2072,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 } @@ -2090,30 +2160,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 } } @@ -2153,7 +2201,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 } @@ -2181,6 +2229,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) diff --git a/backend/internal/service/ratelimit_service_model_not_found_test.go b/backend/internal/service/ratelimit_service_model_not_found_test.go index 51bd8a607..55becb28a 100644 --- a/backend/internal/service/ratelimit_service_model_not_found_test.go +++ b/backend/internal/service/ratelimit_service_model_not_found_test.go @@ -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,