From 1429e8f714e7c875f4cb6f1fd3c1656b928df7f3 Mon Sep 17 00:00:00 2001 From: IanShaw <131567472+IanShaw027@users.noreply.github.com> Date: Thu, 20 Aug 2026 22:01:20 -0700 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=20PR=205888=20=E5=89=A9?= =?UTF-8?q?=E4=BD=99=E5=AE=A1=E8=AE=A1=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../handler/openai_gateway_count_tokens.go | 76 +---------------- .../handler/openai_gateway_handler.go | 4 + .../openai_responses_input_tokens_test.go | 83 ------------------ .../handler/ops_capture_writer_nil_test.go | 11 +-- backend/internal/handler/ops_error_logger.go | 85 ++++++++++--------- .../internal/handler/ops_error_logger_test.go | 50 +++++++++++ backend/internal/service/error_policy_test.go | 32 ++++++- .../service/openai_account_scheduler_test.go | 53 ++++++++++++ .../service/openai_codex_transform_test.go | 8 +- .../service/openai_gateway_scheduling.go | 27 ++++++ .../service/openai_ws_forwarder_ingress.go | 25 +----- .../service/openai_ws_session_preemption.go | 44 ++++++++++ .../openai_ws_session_preemption_test.go | 41 +++++++++ .../service/overload_cooldown_test.go | 32 ++++++- backend/internal/service/ratelimit_service.go | 29 ++++--- 15 files changed, 351 insertions(+), 249 deletions(-) diff --git a/backend/internal/handler/openai_gateway_count_tokens.go b/backend/internal/handler/openai_gateway_count_tokens.go index 126100b349..2df4f4f2d0 100644 --- a/backend/internal/handler/openai_gateway_count_tokens.go +++ b/backend/internal/handler/openai_gateway_count_tokens.go @@ -96,18 +96,12 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) { requestPlatform := openAICompatibleRequestPlatform(c.Request.Context(), apiKey) sessionHash := h.gatewayService.GenerateSessionHash(c, body) requestStart := time.Now() - selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( + account, err := h.gatewayService.SelectAccountForTokenCount( c.Request.Context(), apiKey.GroupID, - "", sessionHash, routingModel, - nil, - service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, - false, - false, - false, requestPlatform, ) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) @@ -120,7 +114,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) { h.errorResponse(c, cls.Status, cls.ErrType, cls.Message) return } - if selection == nil || selection.Account == nil { + if account == nil { cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, routingModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimited(c) @@ -129,16 +123,7 @@ func (h *OpenAIGatewayHandler) ResponsesInputTokens(c *gin.Context) { return } - account := selection.Account setOpsSelectedAccount(c, account.ID, account.Platform) - accountRelease, acquired := h.acquireCountTokensAccountSlot(c, apiKey.GroupID, sessionHash, selection, false, reqLog) - if !acquired { - return - } - if accountRelease != nil { - defer accountRelease() - } - account = selection.Account if err := h.gatewayService.ForwardResponsesInputTokens(c.Request.Context(), c, account, forwardBody); err != nil { reqLog.Error("openai_input_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } @@ -283,18 +268,12 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { if preferredMappedModel != "" { currentRoutingModel = preferredMappedModel } - selection, _, err := h.gatewayService.SelectAccountWithSchedulerForCapability( + account, err := h.gatewayService.SelectAccountForTokenCount( c.Request.Context(), apiKey.GroupID, - "", sessionHash, currentRoutingModel, - nil, - service.OpenAIUpstreamTransportAny, service.OpenAIEndpointCapabilityChatCompletions, - false, - false, - false, openAICompatibleRequestPlatform(c.Request.Context(), apiKey), ) service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) @@ -308,7 +287,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { h.anthropicErrorResponse(c, cls.Status, cls.ErrType, cls.Message) return } - if selection == nil || selection.Account == nil { + if account == nil { cls := classifyOpenAICompatibleNoAccountErrorFromGin(c, h.gatewayService, apiKey, currentRoutingModel, reqModel) if !cls.ModelNotFound { markOpsRoutingCapacityLimited(c) @@ -317,16 +296,7 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { return } - account := selection.Account setOpsSelectedAccount(c, account.ID, account.Platform) - accountRelease, acquired := h.acquireCountTokensAccountSlot(c, apiKey.GroupID, sessionHash, selection, true, reqLog) - if !acquired { - return - } - if accountRelease != nil { - defer accountRelease() - } - account = selection.Account forwardBody := mappedBodyForMessages(channelMapping.Mapped, channelMapping.MappedModel) defaultMappedModel := preferredMappedModel @@ -334,41 +304,3 @@ func (h *OpenAIGatewayHandler) CountTokens(c *gin.Context) { reqLog.Error("openai_count_tokens.forward_failed", zap.Int64("account_id", account.ID), zap.Error(err)) } } - -func (h *OpenAIGatewayHandler) acquireCountTokensAccountSlot( - c *gin.Context, - groupID *int64, - sessionHash string, - selection *service.AccountSelectionResult, - anthropicResponse bool, - reqLog *zap.Logger, -) (func(), bool) { - writeError := func(status int, errType, message string) { - if anthropicResponse { - h.anthropicErrorResponse(c, status, errType, message) - return - } - h.errorResponse(c, status, errType, message) - } - streamStarted := false - release, result := h.acquireOpenAIAccountSlot( - c, - groupID, - sessionHash, - selection, - false, - &streamStarted, - reqLog, - writeError, - ) - if result == openAISlotAcquireOK { - return release, true - } - // Token-count requests suppress the profit gate before selection, so this - // is defensive only. Never forward without a slot if a stale gate appears. - if result == openAISlotAcquireProfitVetoed { - markOpsRoutingCapacityLimited(c) - writeError(http.StatusServiceUnavailable, "api_error", "No available accounts") - } - return nil, false -} diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 45144479c7..9c8ea82a08 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -2466,6 +2466,10 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { // WebSocket 首包可能很大,hash 必须在 hooks 外算成字符串,避免 AfterTurn 闭包保活请求体。 requestPayloadHash = service.HashUsageRequestPayload(wsFirstMessage) + if preemptCtx, cleanupPreempt, armed := h.gatewayService.BeginOpenAIWSIngressSessionPreemption(ctx, c, account, wsFirstMessage); armed { + ctx = preemptCtx + defer cleanupPreempt() + } for { err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks) diff --git a/backend/internal/handler/openai_responses_input_tokens_test.go b/backend/internal/handler/openai_responses_input_tokens_test.go index c9c8aaa625..b3260fc2ad 100644 --- a/backend/internal/handler/openai_responses_input_tokens_test.go +++ b/backend/internal/handler/openai_responses_input_tokens_test.go @@ -1,17 +1,12 @@ package handler import ( - "context" "net/http" "net/http/httptest" - "sync/atomic" "testing" - "time" - "github.com/Wei-Shaw/sub2api/internal/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" - "go.uber.org/zap" ) func TestResponsesInputTokensIsClassifiedAsTokenCountRequest(t *testing.T) { @@ -24,81 +19,3 @@ func TestResponsesInputTokensIsClassifiedAsTokenCountRequest(t *testing.T) { require.True(t, isTokenCountRequestPath("/responses/input_tokens")) require.False(t, isTokenCountRequestPath("/v1/responses")) } - -func TestCountTokensAccountSlot_CancellationStopsBeforeForward(t *testing.T) { - gin.SetMode(gin.TestMode) - - for _, tt := range []struct { - name string - anthropic bool - }{ - {name: "responses input tokens", anthropic: false}, - {name: "anthropic count tokens", anthropic: true}, - } { - t.Run(tt.name, func(t *testing.T) { - cache := &concurrencyCacheMock{ - acquireAccountSlotFn: func(context.Context, int64, int, string) (bool, error) { - return false, nil - }, - } - h := &OpenAIGatewayHandler{ - gatewayService: &service.OpenAIGatewayService{}, - concurrencyHelper: NewConcurrencyHelper(service.NewConcurrencyService(cache), SSEPingFormatNone, time.Second), - } - ctx, cancel := context.WithCancel(context.Background()) - cancel() - recorder := httptest.NewRecorder() - c, _ := gin.CreateTestContext(recorder) - c.Request = httptest.NewRequest(http.MethodPost, "/v1/count_tokens", nil).WithContext(ctx) - groupID := int64(41) - selection := &service.AccountSelectionResult{ - Account: &service.Account{ID: 42, Platform: service.PlatformOpenAI}, - WaitPlan: &service.AccountWaitPlan{ - AccountID: 42, - MaxConcurrency: 1, - MaxWaiting: 1, - Timeout: time.Second, - }, - } - - release, acquired := h.acquireCountTokensAccountSlot(c, &groupID, "", selection, tt.anthropic, zap.NewNop()) - forwarded := false - if acquired { - forwarded = true - if release != nil { - release() - } - } - - require.False(t, acquired) - require.Nil(t, release) - require.False(t, forwarded, "a canceled WaitPlan must stop before count-token forwarding") - require.Zero(t, atomic.LoadInt32(&cache.releaseAccountCalled)) - }) - } -} - -func TestCountTokensAccountSlot_SelectionReleaseRunsExactlyOnce(t *testing.T) { - gin.SetMode(gin.TestMode) - - var released atomic.Int32 - ctx, cancel := context.WithCancel(context.Background()) - recorder := httptest.NewRecorder() - c, _ := gin.CreateTestContext(recorder) - c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses/input_tokens", nil).WithContext(ctx) - h := &OpenAIGatewayHandler{gatewayService: &service.OpenAIGatewayService{}} - groupID := int64(51) - selection := &service.AccountSelectionResult{ - Account: &service.Account{ID: 52, Platform: service.PlatformOpenAI}, - Acquired: true, - ReleaseFunc: func() { released.Add(1) }, - } - - release, acquired := h.acquireCountTokensAccountSlot(c, &groupID, "", selection, false, zap.NewNop()) - require.True(t, acquired) - require.NotNil(t, release) - release() - cancel() - require.Eventually(t, func() bool { return released.Load() == 1 }, time.Second, 10*time.Millisecond) - require.Equal(t, int32(1), released.Load()) -} diff --git a/backend/internal/handler/ops_capture_writer_nil_test.go b/backend/internal/handler/ops_capture_writer_nil_test.go index 38d3f1433b..8b6cd6533f 100644 --- a/backend/internal/handler/ops_capture_writer_nil_test.go +++ b/backend/internal/handler/ops_capture_writer_nil_test.go @@ -176,17 +176,10 @@ func TestOpsCaptureWriter_ReleaseWaitsForDelegatedWriteWithoutHoldingStateMutex( }() <-inner.writeStarted - mutexAvailable := make(chan struct{}) - go func() { - w.state.mu.Lock() - w.state.mu.Unlock() - close(mutexAvailable) - }() - select { - case <-mutexAvailable: - case <-time.After(time.Second): + if !w.state.mu.TryLock() { t.Fatal("state mutex remained held across the delegated network write") } + w.state.mu.Unlock() releaseDone := make(chan struct{}) go func() { diff --git a/backend/internal/handler/ops_error_logger.go b/backend/internal/handler/ops_error_logger.go index 41cd2119b9..41a1ea0583 100644 --- a/backend/internal/handler/ops_error_logger.go +++ b/backend/internal/handler/ops_error_logger.go @@ -1280,6 +1280,18 @@ func logOpsRecoveredUpstream(c *gin.Context, ops *service.OpsService, finalStatu entry := &service.OpsInsertErrorLogInput{StatusCode: finalStatus} applyOpsUpstreamFieldsFromContext(c, entry) + if len(entry.UpstreamErrors) > 0 { + visibleEvents := make([]*service.OpsUpstreamErrorEvent, 0, len(entry.UpstreamErrors)) + for _, event := range entry.UpstreamErrors { + if event != nil && !event.SkipMonitoring { + visibleEvents = append(visibleEvents, event) + } + } + if len(visibleEvents) == 0 { + return + } + applyOpsUpstreamErrorEvents(entry, visibleEvents) + } if entry.UpstreamStatusCode == nil && entry.UpstreamErrorMessage == nil && entry.UpstreamErrorDetail == nil && len(entry.UpstreamErrors) == 0 { return @@ -1350,7 +1362,7 @@ func logOpsRecoveredUpstream(c *gin.Context, ops *service.OpsService, finalStatu apiKey := getOpsAPIKey(c) fallbackPlatform := guessPlatformFromPath(entry.RequestPath) - var requestContext context.Context = context.Background() + requestContext := context.Background() if c.Request != nil { requestContext = c.Request.Context() } @@ -1668,49 +1680,42 @@ func applyOpsUpstreamFieldsFromContext(c *gin.Context, entry *service.OpsInsertE } if v, ok := c.Get(service.OpsUpstreamErrorsKey); ok { if events, ok := v.([]*service.OpsUpstreamErrorEvent); ok && len(events) > 0 { - entry.UpstreamErrors = events - var last *service.OpsUpstreamErrorEvent - for i := len(events) - 1; i >= 0; i-- { - if events[i] != nil { - last = events[i] - break - } - } - if last == nil { - return - } - if last.Stage == string(service.GatewayFailureStageAccountAuth) { - code := 0 - entry.UpstreamStatusCode = &code - entry.UpstreamErrorMessage = nil - if message := strings.TrimSpace(last.Message); message != "" { - entry.UpstreamErrorMessage = &message - } - entry.UpstreamErrorDetail = nil - if detail := strings.TrimSpace(last.Detail); detail != "" { - entry.UpstreamErrorDetail = &detail - } - } else { - entry.UpstreamStatusCode = nil - if last.UpstreamStatusCode > 0 { - code := last.UpstreamStatusCode - entry.UpstreamStatusCode = &code - } - entry.UpstreamErrorMessage = nil - if strings.TrimSpace(last.Message) != "" { - message := strings.TrimSpace(last.Message) - entry.UpstreamErrorMessage = &message - } - entry.UpstreamErrorDetail = nil - if strings.TrimSpace(last.Detail) != "" { - detail := strings.TrimSpace(last.Detail) - entry.UpstreamErrorDetail = &detail - } - } + applyOpsUpstreamErrorEvents(entry, events) } } } +func applyOpsUpstreamErrorEvents(entry *service.OpsInsertErrorLogInput, events []*service.OpsUpstreamErrorEvent) { + entry.UpstreamErrors = events + var last *service.OpsUpstreamErrorEvent + for i := len(events) - 1; i >= 0; i-- { + if events[i] != nil { + last = events[i] + break + } + } + if last == nil { + return + } + + entry.UpstreamStatusCode = nil + entry.UpstreamErrorMessage = nil + entry.UpstreamErrorDetail = nil + if last.Stage == string(service.GatewayFailureStageAccountAuth) { + code := 0 + entry.UpstreamStatusCode = &code + } else if last.UpstreamStatusCode > 0 { + code := last.UpstreamStatusCode + entry.UpstreamStatusCode = &code + } + if message := strings.TrimSpace(last.Message); message != "" { + entry.UpstreamErrorMessage = &message + } + if detail := strings.TrimSpace(last.Detail); detail != "" { + entry.UpstreamErrorDetail = &detail + } +} + func suppressOpsUpstreamAttributionForLocalModelConfiguration(c *gin.Context, entry *service.OpsInsertErrorLogInput) { if entry == nil || !service.HasOpsClientBusinessLimited(c) || service.OpsClientBusinessLimitedReason(c) != service.OpsClientBusinessLimitedReasonLocalModelConfiguration { return diff --git a/backend/internal/handler/ops_error_logger_test.go b/backend/internal/handler/ops_error_logger_test.go index 38176cb1b0..640e8efa91 100644 --- a/backend/internal/handler/ops_error_logger_test.go +++ b/backend/internal/handler/ops_error_logger_test.go @@ -375,6 +375,56 @@ func TestOpsErrorLoggerMiddleware_RecordsRecoveredUpstreamTelemetryOutsideFailur require.Equal(t, http.StatusTooManyRequests, persistedEvents[0].UpstreamStatusCode) } +func TestOpsErrorLoggerMiddleware_RecoveredTelemetryFiltersSkipMonitoringAttempts(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.POST("/v1/responses", func(c *gin.Context) { + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{ + {UpstreamStatusCode: http.StatusTooManyRequests, Message: "visible retry"}, + {UpstreamStatusCode: http.StatusBadGateway, Message: "hidden retry", SkipMonitoring: true}, + }) + c.JSON(http.StatusOK, gin.H{"status": "completed"}) + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) + + require.Equal(t, int64(1), OpsErrorLogQueueLength()) + job := <-opsErrorLogQueue + require.Equal(t, "Recovered upstream error 429: visible retry", job.entry.ErrorMessage) + require.NotNil(t, job.entry.UpstreamErrorsJSON) + events, err := service.ParseOpsUpstreamErrors(*job.entry.UpstreamErrorsJSON) + require.NoError(t, err) + require.Len(t, events, 1) + require.Equal(t, "visible retry", events[0].Message) +} + +func TestOpsErrorLoggerMiddleware_RecoveredTelemetrySkipsAllHiddenAttempts(t *testing.T) { + setupOpsErrorLogTestQueue(t, 2) + gin.SetMode(gin.TestMode) + + ops := service.NewOpsService(nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil) + router := gin.New() + router.Use(OpsErrorLoggerMiddleware(ops)) + router.POST("/v1/responses", func(c *gin.Context) { + c.Set(service.OpsUpstreamErrorsKey, []*service.OpsUpstreamErrorEvent{{ + UpstreamStatusCode: http.StatusTooManyRequests, + Message: "hidden retry", + SkipMonitoring: true, + }}) + c.JSON(http.StatusOK, gin.H{"status": "completed"}) + }) + + recorder := httptest.NewRecorder() + router.ServeHTTP(recorder, httptest.NewRequest(http.MethodPost, "/v1/responses", nil)) + + require.Equal(t, int64(0), OpsErrorLogQueueLength()) +} + func TestOpsErrorLoggerMiddleware_IntermediateSkipMonitoringDoesNotHideFinalVisibleFailure(t *testing.T) { setupOpsErrorLogTestQueue(t, 2) gin.SetMode(gin.TestMode) diff --git a/backend/internal/service/error_policy_test.go b/backend/internal/service/error_policy_test.go index de3d03300b..74cbacf9a2 100644 --- a/backend/internal/service/error_policy_test.go +++ b/backend/internal/service/error_policy_test.go @@ -67,7 +67,7 @@ func TestCheckErrorPolicy(t *testing.T) { expected: ErrorPolicySkipped, }, { - name: "global_529_bypasses_custom_error_code_filter", + name: "custom_error_codes_excluding_529_skip_global_cooldown", account: &Account{ ID: 33, Type: AccountTypeAPIKey, @@ -79,10 +79,10 @@ func TestCheckErrorPolicy(t *testing.T) { }, statusCode: 529, body: []byte(`{"error":{"message":"overloaded"}}`), - expected: ErrorPolicyMatched, + expected: ErrorPolicySkipped, }, { - name: "global_529_bypasses_pool_mode", + name: "pool_mode_skips_global_529_cooldown", account: &Account{ ID: 34, Type: AccountTypeAPIKey, @@ -93,6 +93,32 @@ func TestCheckErrorPolicy(t *testing.T) { }, statusCode: 529, body: []byte(`{"error":{"message":"overloaded"}}`), + expected: ErrorPolicySkipped, + }, + { + name: "ordinary_account_uses_global_529_cooldown", + account: &Account{ + ID: 35, + Type: AccountTypeAPIKey, + Platform: PlatformOpenAI, + }, + statusCode: 529, + body: []byte(`{"error":{"message":"overloaded"}}`), + expected: ErrorPolicyMatched, + }, + { + name: "custom_error_codes_including_529_take_precedence", + account: &Account{ + ID: 36, + Type: AccountTypeAPIKey, + Platform: PlatformOpenAI, + Credentials: map[string]any{ + "custom_error_codes_enabled": true, + "custom_error_codes": []any{float64(529)}, + }, + }, + statusCode: 529, + body: []byte(`{"error":{"message":"overloaded"}}`), expected: ErrorPolicyMatched, }, { diff --git a/backend/internal/service/openai_account_scheduler_test.go b/backend/internal/service/openai_account_scheduler_test.go index 62dde1ee8c..3e8c241b2a 100644 --- a/backend/internal/service/openai_account_scheduler_test.go +++ b/backend/internal/service/openai_account_scheduler_test.go @@ -620,6 +620,59 @@ func TestOpenAIGatewayService_SelectAccountWithScheduler_DefaultDisabled_Embeddi require.Equal(t, openAIAccountScheduleLayerLoadBalance, decision.Layer) } +func TestOpenAIGatewayService_SelectAccountForTokenCount_DoesNotAcquireGenerationSlot(t *testing.T) { + ctx := context.Background() + groupID := int64(10115) + acquiredIDs := make([]int64, 0) + accounts := []Account{ + { + ID: 36501, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 0, + Credentials: map[string]any{"openai_capabilities": []any{"chat_completions"}}, + }, + { + ID: 36502, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 5, + Credentials: map[string]any{"openai_capabilities": []any{"embeddings"}}, + }, + { + ID: 36503, Platform: PlatformGrok, Type: AccountTypeAPIKey, + Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 10, + Credentials: map[string]any{"openai_capabilities": []any{"chat_completions"}}, + }, + { + ID: 36504, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, + Status: StatusActive, Schedulable: true, Concurrency: 1, Priority: 15, + Credentials: map[string]any{ + "openai_capabilities": []any{"chat_completions"}, + "model_mapping": map[string]any{"gpt-4o": "gpt-4o"}, + }, + }, + } + svc := &OpenAIGatewayService{ + accountRepo: schedulerTestOpenAIAccountRepo{accounts: accounts}, + cache: &schedulerTestGatewayCache{}, + cfg: &config.Config{}, + concurrencyService: NewConcurrencyService(schedulerTestConcurrencyCache{ + acquireResults: map[int64]bool{36501: false}, + acquiredIDs: &acquiredIDs, + }), + } + + account, err := svc.SelectAccountForTokenCount( + ctx, + &groupID, + "", + "gpt-5.1", + OpenAIEndpointCapabilityChatCompletions, + PlatformOpenAI, + ) + require.NoError(t, err) + require.NotNil(t, account) + require.Equal(t, int64(36501), account.ID) + require.Empty(t, acquiredIDs, "token counting must not acquire a generation slot") +} + // 生图意图的 /v1/responses 请求要求 OpenAIEndpointCapabilityResponses:探测确认 // 不支持 Responses API 的 APIKey 账号必须被排除,避免 forward 阶段降级为无法生图 // 的 Chat Completions 直转(#4417)。 diff --git a/backend/internal/service/openai_codex_transform_test.go b/backend/internal/service/openai_codex_transform_test.go index 7895b03cc3..96e7715d5d 100644 --- a/backend/internal/service/openai_codex_transform_test.go +++ b/backend/internal/service/openai_codex_transform_test.go @@ -142,9 +142,11 @@ func TestApplyCodexOAuthTransform_NormalizesIsolatedLegacyReferenceAcrossTurns(t input, ok := reqBody["input"].([]any) require.True(t, ok) require.Len(t, input, 3) - require.Equal(t, "fc_previous_turn", input[0].(map[string]any)["id"]) - require.Equal(t, "fc_remote_item", input[1].(map[string]any)["id"]) - require.Equal(t, "vendor_remote_item", input[2].(map[string]any)["id"]) + for i, expectedID := range []string{"fc_previous_turn", "fc_remote_item", "vendor_remote_item"} { + item, itemOK := input[i].(map[string]any) + require.True(t, itemOK) + require.Equal(t, expectedID, item["id"]) + } } func TestApplyCodexOAuthTransform_BoundsLongCallIDsAndPreservesPairing(t *testing.T) { diff --git a/backend/internal/service/openai_gateway_scheduling.go b/backend/internal/service/openai_gateway_scheduling.go index 8ecdd45eb4..435ec6545a 100644 --- a/backend/internal/service/openai_gateway_scheduling.go +++ b/backend/internal/service/openai_gateway_scheduling.go @@ -257,6 +257,33 @@ func (s *OpenAIGatewayService) SelectAccountForModelWithExclusions(ctx context.C return s.selectAccountForModelWithExclusions(s.withOpenAIQuotaAutoPauseContext(ctx), groupID, PlatformOpenAI, sessionHash, requestedModel, excludedIDs, false, 0, "", false) } +// SelectAccountForTokenCount selects an account for a non-billable token-count +// request. It applies the normal platform, model, capability, and runtime +// eligibility checks without acquiring or waiting for a generation slot. +func (s *OpenAIGatewayService) SelectAccountForTokenCount( + ctx context.Context, + groupID *int64, + sessionHash string, + requestedModel string, + requiredCapability OpenAIEndpointCapability, + platform string, +) (*Account, error) { + ctx = WithOpenAIProfitControlSuppressed(ctx) + ctx = s.withOpenAIQuotaAutoPauseContext(ctx) + return s.selectAccountForModelWithExclusions( + ctx, + groupID, + platform, + sessionHash, + requestedModel, + nil, + false, + 0, + requiredCapability, + false, + ) +} + // NormalizeOpenAICompatiblePlatform 保留 grok 与国产 OpenAI 兼容供应商(kimi/zhipu/ // deepseek)的原值,其他值一律归一为 openai。调度器据此对账号与请求做精确平台匹配: // kimi 分组请求只命中 kimi 账号,语义与 openai/grok 一致。 diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index 81e5847fa5..19de862623 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -99,22 +99,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( } } - // Only persistent inbound WebSocket sessions participate in preemption. - // HTTP ingress that opportunistically uses an upstream WS is handled by - // forwardOpenAIWSV2 and deliberately never reaches this registration. - preemptSessionHash := "" - preemptGroupID := getOpenAIGroupIDFromContext(c) - if account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth { - preemptSessionHash = s.GenerateSessionHash(c, firstClientMessage) - } - if preemptCtx, cleanupPreempt, armed, preemptedPrevious := s.beginOpenAIWSSessionPreemptContext( - ctx, - account, - preemptGroupID, - getAPIKeyIDFromContext(c), - preemptSessionHash, - false, - ); armed { + // The handler normally owns this registration across retry attempts. Direct + // callers still get the same session-scoped preemption behavior here. + if preemptCtx, cleanupPreempt, armed := s.BeginOpenAIWSIngressSessionPreemption(ctx, c, account, firstClientMessage); armed { ctx = preemptCtx defer cleanupPreempt() defer func() { @@ -122,12 +109,6 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( returnErr = errOpenAIWSSessionPreempted } }() - if preemptedPrevious { - if stateStore := s.getOpenAIWSStateStore(); stateStore != nil { - stateStore.DeleteSessionTurnState(preemptGroupID, preemptSessionHash) - stateStore.DeleteSessionConn(preemptGroupID, preemptSessionHash) - } - } } wsDecision := s.getOpenAIWSProtocolResolver().Resolve(account) diff --git a/backend/internal/service/openai_ws_session_preemption.go b/backend/internal/service/openai_ws_session_preemption.go index 1f69426da4..54da7f7a1d 100644 --- a/backend/internal/service/openai_ws_session_preemption.go +++ b/backend/internal/service/openai_ws_session_preemption.go @@ -8,6 +8,7 @@ import ( "sync" "time" + "github.com/gin-gonic/gin" "github.com/google/uuid" ) @@ -38,6 +39,49 @@ type openAIWSSessionPreemptKey struct { sessionHash string } +type openAIWSSessionPreemptContextKey struct{} + +// BeginOpenAIWSIngressSessionPreemption keeps a persistent inbound WS session +// registered across upstream retry attempts. Nested forwarding calls reuse the +// registration so returning from one attempt cannot create a preemption gap. +func (s *OpenAIGatewayService) BeginOpenAIWSIngressSessionPreemption( + ctx context.Context, + c *gin.Context, + account *Account, + firstClientMessage []byte, +) (context.Context, func(), bool) { + if ctx == nil { + ctx = context.Background() + } + if armed, _ := ctx.Value(openAIWSSessionPreemptContextKey{}).(bool); armed { + return ctx, func() {}, true + } + + preemptSessionHash := "" + preemptGroupID := getOpenAIGroupIDFromContext(c) + if account != nil && account.Platform == PlatformOpenAI && account.Type == AccountTypeOAuth { + preemptSessionHash = s.GenerateSessionHash(c, firstClientMessage) + } + preemptCtx, cleanup, armed, preemptedPrevious := s.beginOpenAIWSSessionPreemptContext( + ctx, + account, + preemptGroupID, + getAPIKeyIDFromContext(c), + preemptSessionHash, + false, + ) + if !armed { + return ctx, func() {}, false + } + if preemptedPrevious { + if stateStore := s.getOpenAIWSStateStore(); stateStore != nil { + stateStore.DeleteSessionTurnState(preemptGroupID, preemptSessionHash) + stateStore.DeleteSessionConn(preemptGroupID, preemptSessionHash) + } + } + return context.WithValue(preemptCtx, openAIWSSessionPreemptContextKey{}, true), cleanup, true +} + func newOpenAIWSSessionPreemptKey(groupID, apiKeyID int64, sessionHash string) (openAIWSSessionPreemptKey, bool) { sessionHash = strings.TrimSpace(sessionHash) if groupID <= 0 || apiKeyID <= 0 || sessionHash == "" { diff --git a/backend/internal/service/openai_ws_session_preemption_test.go b/backend/internal/service/openai_ws_session_preemption_test.go index 42b5b308bf..9df9465944 100644 --- a/backend/internal/service/openai_ws_session_preemption_test.go +++ b/backend/internal/service/openai_ws_session_preemption_test.go @@ -5,10 +5,13 @@ package service import ( "context" "fmt" + "net/http" + "net/http/httptest" "sync" "testing" "time" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) @@ -110,6 +113,44 @@ func TestOpenAIWSSessionPreemptContextEligibilityAndLocalCancellation(t *testing secondCleanup() } +func TestOpenAIWSIngressSessionPreemptionSurvivesNestedForwardCleanup(t *testing.T) { + gin.SetMode(gin.TestMode) + groupID := int64(7) + newContext := func() *gin.Context { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + c.Set("api_key", &APIKey{ID: 11, GroupID: &groupID}) + return c + } + + svc := &OpenAIGatewayService{} + account := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth} + firstMessage := []byte(`{"type":"response.create","prompt_cache_key":"session-1","input":"hello"}`) + + firstCtx, firstCleanup, armed := svc.BeginOpenAIWSIngressSessionPreemption( + context.Background(), newContext(), account, firstMessage, + ) + require.True(t, armed) + defer firstCleanup() + + // ProxyResponsesWebSocketFromClient enters the same helper for each upstream + // attempt. Its cleanup must not release the handler-owned registration. + nestedCtx, nestedCleanup, armed := svc.BeginOpenAIWSIngressSessionPreemption( + firstCtx, newContext(), account, firstMessage, + ) + require.True(t, armed) + require.Equal(t, firstCtx, nestedCtx) + nestedCleanup() + require.NoError(t, firstCtx.Err()) + + _, secondCleanup, armed := svc.BeginOpenAIWSIngressSessionPreemption( + context.Background(), newContext(), account, firstMessage, + ) + require.True(t, armed) + defer secondCleanup() + require.True(t, IsOpenAIWSSessionPreemptedError(context.Cause(firstCtx))) +} + func TestOpenAIWSSessionPreemptRemoteClaimAndStaleReleaseAreAtomic(t *testing.T) { cache := &openAIWSSessionPreemptCacheStub{} svc := &OpenAIGatewayService{cache: cache} diff --git a/backend/internal/service/overload_cooldown_test.go b/backend/internal/service/overload_cooldown_test.go index 58dc13177c..92fc743d1e 100644 --- a/backend/internal/service/overload_cooldown_test.go +++ b/backend/internal/service/overload_cooldown_test.go @@ -36,10 +36,16 @@ func (r *errSettingRepo) Get(_ context.Context, _ string) (*Setting, error) { type overloadAccountRepoStub struct { mockAccountRepoForGemini overloadCalls int + errorCalls int lastOverloadID int64 lastOverloadEnd time.Time } +func (r *overloadAccountRepoStub) SetError(_ context.Context, _ int64, _ string) error { + r.errorCalls++ + return nil +} + func (r *overloadAccountRepoStub) SetOverloaded(_ context.Context, id int64, until time.Time) error { r.overloadCalls++ r.lastOverloadID = id @@ -269,7 +275,7 @@ func TestHandle529_DBReadError_FallsBackToConfig(t *testing.T) { require.WithinDuration(t, before.Add(7*time.Minute), accountRepo.lastOverloadEnd, 2*time.Second) } -func TestHandleUpstreamError_529BypassesPoolAndCustomCodeGates(t *testing.T) { +func TestHandleUpstreamError_529RespectsAccountPolicies(t *testing.T) { tests := []struct { name string credentials map[string]any @@ -301,12 +307,32 @@ func TestHandleUpstreamError_529BypassesPoolAndCustomCodeGates(t *testing.T) { shouldDisable := svc.HandleUpstreamError(context.Background(), account, 529, nil, []byte(`{"error":{"message":"overloaded"}}`)) require.False(t, shouldDisable) - require.Equal(t, 1, repo.overloadCalls) - require.Equal(t, account.ID, repo.lastOverloadID) + require.Zero(t, repo.overloadCalls) + require.Zero(t, repo.errorCalls) }) } } +func TestHandleUpstreamError_529CustomCodeDisablesInsteadOfOverloadCooldown(t *testing.T) { + repo := &overloadAccountRepoStub{} + svc := NewRateLimitService(repo, nil, &config.Config{}, nil, nil) + account := &Account{ + ID: 102, + Platform: PlatformOpenAI, + Type: AccountTypeAPIKey, + Credentials: map[string]any{ + "custom_error_codes_enabled": true, + "custom_error_codes": []any{float64(529)}, + }, + } + + shouldDisable := svc.HandleUpstreamError(context.Background(), account, 529, nil, []byte(`{"error":{"message":"overloaded"}}`)) + + require.True(t, shouldDisable) + require.Equal(t, 1, repo.errorCalls) + require.Zero(t, repo.overloadCalls) +} + // =========================================================================== // Model: defaults & JSON round-trip // =========================================================================== diff --git a/backend/internal/service/ratelimit_service.go b/backend/internal/service/ratelimit_service.go index 4e4068eb98..2c42b810b6 100644 --- a/backend/internal/service/ratelimit_service.go +++ b/backend/internal/service/ratelimit_service.go @@ -252,11 +252,6 @@ const ( // 自定义错误码开启时覆盖后续所有逻辑(包括临时不可调度)。 func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Account, statusCode int, responseBody []byte, requestedModel ...string) ErrorPolicyResult { ctx = withTempUnschedulableModel(ctx, requestedModel) - // 529 is governed by the global overload cooldown. Return Matched before - // local pool/custom-code filters so every caller reaches HandleUpstreamError. - if statusCode == 529 { - return ErrorPolicyMatched - } if account.IsCustomErrorCodesEnabled() { if account.ShouldHandleErrorCode(statusCode) { return ErrorPolicyMatched @@ -272,6 +267,11 @@ func (s *RateLimitService) CheckErrorPolicy(ctx context.Context, account *Accoun } return ErrorPolicySkipped } + // The global overload cooldown is the default for ordinary accounts. Explicit + // account policies above retain precedence over this fallback. + if statusCode == 529 { + return ErrorPolicyMatched + } if s.tryTempUnschedulable(ctx, account, statusCode, responseBody, firstRequestedModel(requestedModel)) { return ErrorPolicyTempUnscheduled } @@ -287,14 +287,6 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc s.maybeHandleOpenAITeamLinkedError(ctx, account, statusCode, responseBody) customErrorCodesEnabled := account.IsCustomErrorCodesEnabled() - // The configured 529 cooldown is a global overload policy. Apply it before - // pool-mode and custom-code gates so those local retry policies cannot leave - // an overloaded account eligible for new requests. - if statusCode == 529 { - s.handle529(ctx, account) - return false - } - // 池模式默认不标记本地账号状态;但管理员显式配置的临时不可调度规则优先。 // 401 保留现有认证错误语义,不在这里改变池模式的认证处理。 if account.IsPoolMode() && !customErrorCodesEnabled { @@ -312,6 +304,15 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc return false } + if statusCode == 529 { + if customErrorCodesEnabled { + s.handleCustomErrorCode(ctx, account, statusCode, extractUpstreamErrorMessage(responseBody)) + return true + } + s.handle529(ctx, account) + return false + } + if len(requestedModel) > 0 && s.HandleUpstreamModelNotFound(ctx, account, requestedModel[0], statusCode, responseBody) { return true } @@ -499,7 +500,7 @@ func (s *RateLimitService) HandleUpstreamError(ctx context.Context, account *Acc s.handle429(ctx, account, headers, responseBody) shouldDisable = false case 529: - // Handled before pool/custom-code policy gates above. + // Handled after pool/custom-code policy gates above. shouldDisable = false default: // 自定义错误码启用时:在列表中的错误码都应该停止调度