From 17455a60a67610c5314ed1ec55c57174e4b30822 Mon Sep 17 00:00:00 2001 From: feitianbubu Date: Wed, 29 Jul 2026 21:14:40 +0800 Subject: [PATCH] fix: exclude failed requests from memory rate limit success count --- common/rate-limit.go | 12 ++++++++++ middleware/model-rate-limit.go | 14 +++++------- middleware/model_rate_limit_test.go | 34 +++++++++++++++++++++++++++++ 3 files changed, 51 insertions(+), 9 deletions(-) diff --git a/common/rate-limit.go b/common/rate-limit.go index 301c101c9748..a6c020e536de 100644 --- a/common/rate-limit.go +++ b/common/rate-limit.go @@ -41,6 +41,18 @@ func (l *InMemoryRateLimiter) clearExpiredItems() { } } +// CanRequest reports whether Request would allow key right now, without +// recording anything. Parameter duration's unit is seconds. +func (l *InMemoryRateLimiter) CanRequest(key string, maxRequestNum int, duration int64) bool { + l.mutex.Lock() + defer l.mutex.Unlock() + queue, ok := l.store[key] + if !ok || len(*queue) < maxRequestNum { + return true + } + return time.Now().Unix()-(*queue)[0] >= duration +} + // Request parameter duration's unit is seconds func (l *InMemoryRateLimiter) Request(key string, maxRequestNum int, duration int64) bool { l.mutex.Lock() diff --git a/middleware/model-rate-limit.go b/middleware/model-rate-limit.go index 9f1d94039685..258e5e0909ae 100644 --- a/middleware/model-rate-limit.go +++ b/middleware/model-rate-limit.go @@ -140,24 +140,20 @@ func memoryRateLimitHandler(duration int64, totalMaxCount, successMaxCount int) // 1. 检查总请求数限制(当totalMaxCount为0时跳过) if totalMaxCount > 0 && !inMemoryRateLimiter.Request(totalKey, totalMaxCount, duration) { - c.Status(http.StatusTooManyRequests) - c.Abort() + abortWithOpenAiMessage(c, http.StatusTooManyRequests, fmt.Sprintf("您已达到总请求数限制:%d分钟内最多请求%d次,包括失败次数,请检查您的请求是否正确", setting.ModelRequestRateLimitDurationMinutes, totalMaxCount)) return } - // 2. 检查成功请求数限制 - // 使用一个临时key来检查限制,这样可以避免实际记录 - checkKey := successKey + "_check" - if !inMemoryRateLimiter.Request(checkKey, successMaxCount, duration) { - c.Status(http.StatusTooManyRequests) - c.Abort() + // 2. 检查成功请求数限制:只检查不计数,成功后才在步骤4计入 + if !inMemoryRateLimiter.CanRequest(successKey, successMaxCount, duration) { + abortWithOpenAiMessage(c, http.StatusTooManyRequests, fmt.Sprintf("您已达到请求数限制:%d分钟内最多请求%d次", setting.ModelRequestRateLimitDurationMinutes, successMaxCount)) return } // 3. 处理请求 c.Next() - // 4. 如果请求成功,记录到实际的成功请求计数中 + // 4. 如果请求成功,记录到成功请求计数中 if c.Writer.Status() < 400 { inMemoryRateLimiter.Request(successKey, successMaxCount, duration) } diff --git a/middleware/model_rate_limit_test.go b/middleware/model_rate_limit_test.go index 3e9923fdac15..648f10997fd6 100644 --- a/middleware/model_rate_limit_test.go +++ b/middleware/model_rate_limit_test.go @@ -2,9 +2,12 @@ package middleware import ( "context" + "net/http" + "net/http/httptest" "testing" "time" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -32,3 +35,34 @@ func TestModelRedisRateLimitUsesUTCRegardlessOfLocalTimezone(t *testing.T) { require.NoError(t, err) assert.False(t, allowed, "an existing UTC timestamp inside the window must remain limited on a non-UTC host") } + +// 成功数限制只统计成功请求,失败请求不占配额;拒绝时返回 JSON 错误体而非空 429。 +func TestModelMemoryRateLimitSuccessCountIgnoresFailedRequests(t *testing.T) { + gin.SetMode(gin.TestMode) + + // 限流器是进程级全局且无重置接口,user id 每次唯一才能保证 -count=2 复跑 + userID := int(time.Now().UnixNano() % 1_000_000_000) + downstreamStatus := http.StatusOK + router := gin.New() + router.GET("/limited", func(c *gin.Context) { + c.Set("id", userID) + }, memoryRateLimitHandler(60, 0, 2), func(c *gin.Context) { + c.Status(downstreamStatus) + }) + + do := func(status int) *httptest.ResponseRecorder { + downstreamStatus = status + return performRateLimitRequest(router, "/limited", "192.0.2.70:12345") + } + + for range 5 { + assert.Equal(t, http.StatusInternalServerError, do(http.StatusInternalServerError).Code, "failed requests must not consume the success quota") + } + + require.Equal(t, http.StatusOK, do(http.StatusOK).Code) + require.Equal(t, http.StatusOK, do(http.StatusOK).Code) + + limited := do(http.StatusOK) + require.Equal(t, http.StatusTooManyRequests, limited.Code) + assert.Contains(t, limited.Body.String(), "您已达到请求数限制", "429 must carry the same error message as the Redis path") +}