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:
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user