修复 PR 5888 剩余审计问题

This commit is contained in:
IanShaw
2026-08-20 22:01:20 -07:00
parent b2b2adcf8d
commit 1429e8f714
15 changed files with 351 additions and 249 deletions
@@ -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
}
@@ -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)
@@ -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())
}
@@ -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() {
+45 -40
View File
@@ -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
@@ -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)
+29 -3
View File
@@ -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,
},
{
@@ -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)。
@@ -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) {
@@ -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 一致。
@@ -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)
@@ -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 == "" {
@@ -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}
@@ -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
// ===========================================================================
+15 -14
View File
@@ -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:
// 自定义错误码启用时:在列表中的错误码都应该停止调度