From b0579c489179c812e8cc9ab9ff964fe039ab2dff Mon Sep 17 00:00:00 2001 From: jjaw Date: Sun, 14 Jun 2026 02:40:23 +0800 Subject: [PATCH] fix: move user wait queue accounting off hot path --- .../handler/concurrency_error_response.go | 6 + backend/internal/handler/gateway_handler.go | 27 ----- .../gateway_handler_chat_completions.go | 23 ---- .../handler/gateway_handler_responses.go | 23 ---- backend/internal/handler/gateway_helper.go | 27 ++++- .../handler/gateway_helper_hotpath_test.go | 112 ++++++++++++++++++ .../internal/handler/gemini_v1beta_handler.go | 24 ---- .../handler/openai_gateway_handler.go | 36 +----- 8 files changed, 145 insertions(+), 133 deletions(-) diff --git a/backend/internal/handler/concurrency_error_response.go b/backend/internal/handler/concurrency_error_response.go index 52abf73524..911f5a54d7 100644 --- a/backend/internal/handler/concurrency_error_response.go +++ b/backend/internal/handler/concurrency_error_response.go @@ -10,6 +10,12 @@ import ( const statusClientClosedRequest = 499 func concurrencyErrorResponse(err error, slotType string) (int, string, string) { + var waitQueueFullErr *WaitQueueFullError + if errors.As(err, &waitQueueFullErr) { + return http.StatusTooManyRequests, "rate_limit_error", + "Too many pending requests, please retry later" + } + var concurrencyErr *ConcurrencyError if errors.As(err, &concurrencyErr) { if concurrencyErr.SlotType != "" { diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index d0e2c6b730..5c909dc6fb 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -211,28 +211,6 @@ func (h *GatewayHandler) Messages(c *gin.Context) { // 获取订阅信息(可能为nil)- 提前获取用于后续检查 subscription, _ := middleware2.GetSubscriptionFromContext(c) - // 0. 检查wait队列是否已满 - maxWait := service.CalculateMaxWait(subject.Concurrency) - canWait, err := h.concurrencyHelper.IncrementWaitCount(c.Request.Context(), subject.UserID, maxWait) - waitCounted := false - if err != nil { - reqLog.Warn("gateway.user_wait_counter_increment_failed", zap.Error(err)) - // On error, allow request to proceed - } else if !canWait { - reqLog.Info("gateway.user_wait_queue_full", zap.Int("max_wait", maxWait)) - h.errorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later") - return - } - if err == nil && canWait { - waitCounted = true - } - // Ensure we decrement if we exit before acquiring the user slot. - defer func() { - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(c.Request.Context(), subject.UserID) - } - }() - // 1. 首先获取用户并发槽位 userReleaseFunc, err := h.concurrencyHelper.AcquireUserSlotWithWait(c, subject.UserID, subject.Concurrency, reqStream, &streamStarted) if err != nil { @@ -240,11 +218,6 @@ func (h *GatewayHandler) Messages(c *gin.Context) { h.handleConcurrencyError(c, err, "user", streamStarted) return } - // User slot acquired: no longer waiting in the queue. - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(c.Request.Context(), subject.UserID) - waitCounted = false - } // 在请求结束或 Context 取消时确保释放槽位,避免客户端断开造成泄漏 userReleaseFunc = wrapReleaseOnDone(c.Request.Context(), userReleaseFunc) if userReleaseFunc != nil { diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index 49347d02f0..712c2b9fb4 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -109,35 +109,12 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) { service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) - // 1. Acquire user concurrency slot - maxWait := service.CalculateMaxWait(subject.Concurrency) - canWait, err := h.concurrencyHelper.IncrementWaitCount(c.Request.Context(), subject.UserID, maxWait) - waitCounted := false - if err != nil { - reqLog.Warn("gateway.cc.user_wait_counter_increment_failed", zap.Error(err)) - } else if !canWait { - h.chatCompletionsErrorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later") - return - } - if err == nil && canWait { - waitCounted = true - } - defer func() { - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(c.Request.Context(), subject.UserID) - } - }() - userReleaseFunc, err := h.concurrencyHelper.AcquireUserSlotWithWait(c, subject.UserID, subject.Concurrency, reqStream, &streamStarted) if err != nil { reqLog.Warn("gateway.cc.user_slot_acquire_failed", zap.Error(err)) h.handleConcurrencyError(c, err, "user", streamStarted) return } - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(c.Request.Context(), subject.UserID) - waitCounted = false - } userReleaseFunc = wrapReleaseOnDone(c.Request.Context(), userReleaseFunc) if userReleaseFunc != nil { defer userReleaseFunc() diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index 141c85ae61..a813f5f767 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -118,35 +118,12 @@ func (h *GatewayHandler) Responses(c *gin.Context) { service.SetOpsLatencyMs(c, service.OpsAuthLatencyMsKey, time.Since(requestStart).Milliseconds()) - // 1. Acquire user concurrency slot - maxWait := service.CalculateMaxWait(subject.Concurrency) - canWait, err := h.concurrencyHelper.IncrementWaitCount(c.Request.Context(), subject.UserID, maxWait) - waitCounted := false - if err != nil { - reqLog.Warn("gateway.responses.user_wait_counter_increment_failed", zap.Error(err)) - } else if !canWait { - h.responsesErrorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later") - return - } - if err == nil && canWait { - waitCounted = true - } - defer func() { - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(c.Request.Context(), subject.UserID) - } - }() - userReleaseFunc, err := h.concurrencyHelper.AcquireUserSlotWithWait(c, subject.UserID, subject.Concurrency, reqStream, &streamStarted) if err != nil { reqLog.Warn("gateway.responses.user_slot_acquire_failed", zap.Error(err)) h.handleConcurrencyError(c, err, "user", streamStarted) return } - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(c.Request.Context(), subject.UserID) - waitCounted = false - } userReleaseFunc = wrapReleaseOnDone(c.Request.Context(), userReleaseFunc) if userReleaseFunc != nil { defer userReleaseFunc() diff --git a/backend/internal/handler/gateway_helper.go b/backend/internal/handler/gateway_helper.go index 4b6a47eb2f..b948ac8fc7 100644 --- a/backend/internal/handler/gateway_helper.go +++ b/backend/internal/handler/gateway_helper.go @@ -127,6 +127,14 @@ func (e *ConcurrencyError) Error() string { return fmt.Sprintf("%s concurrency limit reached", e.SlotType) } +type WaitQueueFullError struct { + SlotType string +} + +func (e *WaitQueueFullError) Error() string { + return "Too many pending requests, please retry later" +} + // ConcurrencyHelper provides common concurrency slot management for gateway handlers type ConcurrencyHelper struct { concurrencyService *service.ConcurrencyService @@ -220,6 +228,10 @@ func (h *ConcurrencyHelper) TryAcquireAccountSlot(ctx context.Context, accountID // For streaming requests, sends ping events during the wait. // streamStarted is updated if streaming response has begun. func (h *ConcurrencyHelper) AcquireUserSlotWithWait(c *gin.Context, userID int64, maxConcurrency int, isStream bool, streamStarted *bool) (func(), error) { + return h.acquireUserSlotWithWaitTimeout(c, userID, maxConcurrency, maxConcurrencyWait, isStream, streamStarted) +} + +func (h *ConcurrencyHelper) acquireUserSlotWithWaitTimeout(c *gin.Context, userID int64, maxConcurrency int, timeout time.Duration, isStream bool, streamStarted *bool) (func(), error) { ctx := c.Request.Context() // Try to acquire immediately @@ -232,8 +244,21 @@ func (h *ConcurrencyHelper) AcquireUserSlotWithWait(c *gin.Context, userID int64 return releaseFunc, nil } + queueLimit := service.CalculateMaxWait(maxConcurrency) - maxConcurrency + if queueLimit < 1 { + queueLimit = 1 + } + canWait, err := h.IncrementWaitCount(ctx, userID, queueLimit) + if err != nil { + return nil, err + } + if !canWait { + return nil, &WaitQueueFullError{SlotType: "user"} + } + defer h.DecrementWaitCount(ctx, userID) + // Need to wait - handle streaming ping if needed - return h.waitForSlotWithPing(c, "user", userID, maxConcurrency, isStream, streamStarted) + return h.waitForSlotWithPingTimeout(c, "user", userID, maxConcurrency, timeout, isStream, streamStarted, false) } // AcquireAccountSlotWithWait acquires an account concurrency slot, waiting if necessary. diff --git a/backend/internal/handler/gateway_helper_hotpath_test.go b/backend/internal/handler/gateway_helper_hotpath_test.go index a6b6a429e9..65dc849683 100644 --- a/backend/internal/handler/gateway_helper_hotpath_test.go +++ b/backend/internal/handler/gateway_helper_hotpath_test.go @@ -24,6 +24,11 @@ type helperConcurrencyCacheStub struct { userAcquireCalls int accountReleaseCalls int userReleaseCalls int + waitAllowed bool + waitIncrementCalls int + waitDecrementCalls int + waitMaxWait int + waitIncrementHook func() } func (s *helperConcurrencyCacheStub) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) { @@ -93,10 +98,25 @@ func (s *helperConcurrencyCacheStub) GetUserConcurrency(ctx context.Context, use } func (s *helperConcurrencyCacheStub) IncrementWaitCount(ctx context.Context, userID int64, maxWait int) (bool, error) { + s.mu.Lock() + s.waitIncrementCalls++ + s.waitMaxWait = maxWait + waitAllowed := s.waitAllowed + hook := s.waitIncrementHook + s.mu.Unlock() + if hook != nil { + hook() + } + if !waitAllowed { + return false, nil + } return true, nil } func (s *helperConcurrencyCacheStub) DecrementWaitCount(ctx context.Context, userID int64) error { + s.mu.Lock() + defer s.mu.Unlock() + s.waitDecrementCalls++ return nil } @@ -226,6 +246,98 @@ func TestWaitForSlotWithPingTimeout_AccountAndUserAcquire(t *testing.T) { }) } +func TestAcquireUserSlotWithWait_ImmediateAcquireSkipsWaitQueue(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{true}, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + streamStarted := false + + release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, time.Second, false, &streamStarted) + require.NoError(t, err) + require.NotNil(t, release) + release() + + require.Equal(t, 1, cache.userAcquireCalls) + require.Equal(t, 0, cache.waitIncrementCalls) + require.Equal(t, 0, cache.waitDecrementCalls) + require.Equal(t, 1, cache.userReleaseCalls) +} + +func TestAcquireUserSlotWithWait_WaitSuccessDecrementsBeforeReturn(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{false, true}, + waitAllowed: true, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + streamStarted := false + + release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, time.Second, false, &streamStarted) + require.NoError(t, err) + require.NotNil(t, release) + + require.Equal(t, 2, cache.userAcquireCalls) + require.Equal(t, 1, cache.waitIncrementCalls) + require.Equal(t, 20, cache.waitMaxWait) + require.Equal(t, 1, cache.waitDecrementCalls) + + release() + require.Equal(t, 1, cache.userReleaseCalls) +} + +func TestAcquireUserSlotWithWait_TimeoutDecrementsWaitQueue(t *testing.T) { + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{false, false, false}, + waitAllowed: true, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + streamStarted := false + + release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, 30*time.Millisecond, false, &streamStarted) + require.Nil(t, release) + var cErr *ConcurrencyError + require.ErrorAs(t, err, &cErr) + require.True(t, cErr.IsTimeout) + require.Equal(t, 1, cache.waitIncrementCalls) + require.Equal(t, 1, cache.waitDecrementCalls) + require.Equal(t, 0, cache.userReleaseCalls) +} + +func TestAcquireUserSlotWithWait_RequestCancelDecrementsWaitQueue(t *testing.T) { + cancelled := make(chan struct{}) + var cancel context.CancelFunc + cache := &helperConcurrencyCacheStub{ + userSeq: []bool{false, false}, + waitAllowed: true, + waitIncrementHook: func() { + cancel() + close(cancelled) + }, + } + concurrency := service.NewConcurrencyService(cache) + helper := NewConcurrencyHelper(concurrency, SSEPingFormatNone, 5*time.Millisecond) + c, _ := newHelperTestContext(http.MethodPost, "/v1/messages") + reqCtx, cancelFunc := context.WithCancel(c.Request.Context()) + cancel = cancelFunc + defer cancel() + c.Request = c.Request.WithContext(reqCtx) + streamStarted := false + + release, err := helper.acquireUserSlotWithWaitTimeout(c, 202, 3, time.Second, false, &streamStarted) + <-cancelled + require.Nil(t, release) + require.ErrorIs(t, err, context.Canceled) + require.Equal(t, 1, cache.waitIncrementCalls) + require.Equal(t, 1, cache.waitDecrementCalls) + require.Equal(t, 0, cache.userReleaseCalls) +} + func TestWaitForSlotWithPingTimeout_TimeoutAndStreamPing(t *testing.T) { cache := &helperConcurrencyCacheStub{ accountSeq: []bool{false, false, false}, diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index b7781eec5d..d5918e3583 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -205,26 +205,6 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { // For Gemini native API, do not send Claude-style ping frames. geminiConcurrency := NewConcurrencyHelper(h.concurrencyHelper.concurrencyService, SSEPingFormatNone, 0) - // 0) wait queue check - maxWait := service.CalculateMaxWait(authSubject.Concurrency) - canWait, err := geminiConcurrency.IncrementWaitCount(c.Request.Context(), authSubject.UserID, maxWait) - waitCounted := false - if err != nil { - reqLog.Warn("gemini.user_wait_counter_increment_failed", zap.Error(err)) - } else if !canWait { - reqLog.Info("gemini.user_wait_queue_full", zap.Int("max_wait", maxWait)) - googleError(c, http.StatusTooManyRequests, "Too many pending requests, please retry later") - return - } - if err == nil && canWait { - waitCounted = true - } - defer func() { - if waitCounted { - geminiConcurrency.DecrementWaitCount(c.Request.Context(), authSubject.UserID) - } - }() - // 1) user concurrency slot streamStarted := false if h.errorPassthroughService != nil { @@ -236,10 +216,6 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) { googleError(c, http.StatusTooManyRequests, err.Error()) return } - if waitCounted { - geminiConcurrency.DecrementWaitCount(c.Request.Context(), authSubject.UserID) - waitCounted = false - } // 确保请求取消时也会释放槽位,避免长连接被动中断造成泄漏 userReleaseFunc = wrapReleaseOnDone(c.Request.Context(), userReleaseFunc) if userReleaseFunc != nil { diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 65bf172ecd..3e142ae6c5 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -1019,46 +1019,12 @@ func (h *OpenAIGatewayHandler) acquireResponsesUserSlot( reqLog *zap.Logger, ) (func(), bool) { ctx := c.Request.Context() - userReleaseFunc, userAcquired, err := h.concurrencyHelper.TryAcquireUserSlot(ctx, userID, userConcurrency) + userReleaseFunc, err := h.concurrencyHelper.AcquireUserSlotWithWait(c, userID, userConcurrency, reqStream, streamStarted) if err != nil { reqLog.Warn("openai.user_slot_acquire_failed", zap.Error(err)) h.handleConcurrencyError(c, err, "user", *streamStarted) return nil, false } - if userAcquired { - return wrapReleaseOnDone(ctx, userReleaseFunc), true - } - - maxWait := service.CalculateMaxWait(userConcurrency) - canWait, waitErr := h.concurrencyHelper.IncrementWaitCount(ctx, userID, maxWait) - if waitErr != nil { - reqLog.Warn("openai.user_wait_counter_increment_failed", zap.Error(waitErr)) - // 按现有降级语义:等待计数异常时放行后续抢槽流程 - } else if !canWait { - reqLog.Info("openai.user_wait_queue_full", zap.Int("max_wait", maxWait)) - h.errorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Too many pending requests, please retry later") - return nil, false - } - - waitCounted := waitErr == nil && canWait - defer func() { - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(ctx, userID) - } - }() - - userReleaseFunc, err = h.concurrencyHelper.AcquireUserSlotWithWait(c, userID, userConcurrency, reqStream, streamStarted) - if err != nil { - reqLog.Warn("openai.user_slot_acquire_failed_after_wait", zap.Error(err)) - h.handleConcurrencyError(c, err, "user", *streamStarted) - return nil, false - } - - // 槽位获取成功后,立刻退出等待计数。 - if waitCounted { - h.concurrencyHelper.DecrementWaitCount(ctx, userID) - waitCounted = false - } return wrapReleaseOnDone(ctx, userReleaseFunc), true }