fix(auth): keep login state on rate-limited or failing token refresh

When the dashboard token refresh endpoint returned 429 (shared
critical rate limit) the frontend classified it as out_of_sync,
cleared local auth state, and redirected to /sign-in. The rate limit
itself is working as intended; the bug is that a temporary rejection
was treated as a terminal auth failure.

- Treat 429 refresh responses as transient errors on the frontend,
  keeping the session retryable instead of clearing it. Only explicit
  401 or confirmed session mismatch/race exhaustion clears auth state.
- Return Retry-After on all rate-limited responses (remaining TTL on
  Redis, window duration on the in-memory limiter) so clients can
  back off.
- Log the underlying error with request context when auth session
  errors map to 500 AUTH_INTERNAL_ERROR, and replace fmt.Println with
  request-scoped logging in the Redis rate limiter error paths.

Fixes #6361
This commit is contained in:
CaIon
2026-07-21 12:39:29 +08:00
parent e0d5156115
commit 1721144221
5 changed files with 57 additions and 14 deletions
+8
View File
@@ -2,9 +2,11 @@ package controller
import ( import (
"errors" "errors"
"fmt"
"net/http" "net/http"
"strings" "strings"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/middleware" "github.com/QuantumNous/new-api/middleware"
"github.com/QuantumNous/new-api/model" "github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service" "github.com/QuantumNous/new-api/service"
@@ -164,6 +166,12 @@ func writeAuthSessionError(c *gin.Context, err error) {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
status, code = http.StatusUnauthorized, "AUTH_UNAUTHORIZED" status, code = http.StatusUnauthorized, "AUTH_UNAUTHORIZED"
} }
if status == http.StatusInternalServerError {
// The response body only carries the generic AUTH_INTERNAL_ERROR
// code; without this log the underlying Redis/database/session
// failure is indistinguishable from the client side.
logger.LogError(c.Request.Context(), fmt.Sprintf("auth session internal error (%s %s): %v", c.Request.Method, c.Request.URL.Path, err))
}
c.JSON(status, gin.H{"success": false, "code": code, "message": http.StatusText(status)}) c.JSON(status, gin.H{"success": false, "code": code, "message": http.StatusText(status)})
} }
+21 -12
View File
@@ -8,6 +8,7 @@ import (
"strconv" "strconv"
"github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/logger"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@@ -106,33 +107,43 @@ func redisFixedWindowTake(ctx context.Context, key string, maxRequestNum int, du
} }
func redisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark string) { func redisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark string) {
allowed, _, _, err := redisFixedWindowTake( allowed, _, ttlSeconds, err := redisFixedWindowTake(
c.Request.Context(), c.Request.Context(),
redisIPRateLimitKey(mark, c.ClientIP()), redisIPRateLimitKey(mark, c.ClientIP()),
maxRequestNum, maxRequestNum,
duration, duration,
) )
if err != nil { if err != nil {
fmt.Println(err.Error()) logger.LogError(c.Request.Context(), fmt.Sprintf("rate limit check failed (mark=%s): %v", mark, err))
c.Status(http.StatusInternalServerError) c.Status(http.StatusInternalServerError)
c.Abort() c.Abort()
return return
} }
if !allowed { if !allowed {
c.Status(http.StatusTooManyRequests) writeRateLimited(c, ttlSeconds)
c.Abort()
} }
} }
func memoryRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark string) { func memoryRateLimiter(c *gin.Context, maxRequestNum int, duration int64, mark string) {
key := mark + c.ClientIP() key := mark + c.ClientIP()
if !inMemoryRateLimiter.Request(key, maxRequestNum, duration) { if !inMemoryRateLimiter.Request(key, maxRequestNum, duration) {
c.Status(http.StatusTooManyRequests) writeRateLimited(c, duration)
c.Abort()
return return
} }
} }
// writeRateLimited rejects the request with 429 and a Retry-After hint so
// clients can back off instead of treating the rejection as a fatal error.
// The in-memory limiter cannot report the remaining window, so callers
// without a TTL pass the full window duration as a conservative upper bound.
func writeRateLimited(c *gin.Context, retryAfterSeconds int64) {
if retryAfterSeconds > 0 {
c.Header("Retry-After", strconv.FormatInt(retryAfterSeconds, 10))
}
c.Status(http.StatusTooManyRequests)
c.Abort()
}
func rateLimitFactory(maxRequestNum int, duration int64, mark string) func(c *gin.Context) { func rateLimitFactory(maxRequestNum int, duration int64, mark string) func(c *gin.Context) {
if common.RedisEnabled { if common.RedisEnabled {
return func(c *gin.Context) { return func(c *gin.Context) {
@@ -201,8 +212,7 @@ func userRateLimitFactory(maxRequestNum int, duration int64, mark string) func(c
} }
key := fmt.Sprintf("%s:user:%d", mark, userID) key := fmt.Sprintf("%s:user:%d", mark, userID)
if !inMemoryRateLimiter.Request(key, maxRequestNum, duration) { if !inMemoryRateLimiter.Request(key, maxRequestNum, duration) {
c.Status(http.StatusTooManyRequests) writeRateLimited(c, duration)
c.Abort()
return return
} }
} }
@@ -211,16 +221,15 @@ func userRateLimitFactory(maxRequestNum int, duration int64, mark string) func(c
// userRedisRateLimiter is like redisRateLimiter but accepts a pre-built key // userRedisRateLimiter is like redisRateLimiter but accepts a pre-built key
// (to support user-ID-based keys). // (to support user-ID-based keys).
func userRedisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, key string) { func userRedisRateLimiter(c *gin.Context, maxRequestNum int, duration int64, key string) {
allowed, _, _, err := redisFixedWindowTake(c.Request.Context(), key, maxRequestNum, duration) allowed, _, ttlSeconds, err := redisFixedWindowTake(c.Request.Context(), key, maxRequestNum, duration)
if err != nil { if err != nil {
fmt.Println(err.Error()) logger.LogError(c.Request.Context(), fmt.Sprintf("rate limit check failed (key=%s): %v", key, err))
c.Status(http.StatusInternalServerError) c.Status(http.StatusInternalServerError)
c.Abort() c.Abort()
return return
} }
if !allowed { if !allowed {
c.Status(http.StatusTooManyRequests) writeRateLimited(c, ttlSeconds)
c.Abort()
} }
} }
+3 -1
View File
@@ -61,7 +61,9 @@ func TestRedisIPRateLimiterThresholdTTLAndNamespace(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, http.StatusNoContent, performRateLimitRequest(router, "/limited", remoteAddr).Code) assert.Equal(t, http.StatusNoContent, performRateLimitRequest(router, "/limited", remoteAddr).Code)
assert.Equal(t, http.StatusNoContent, performRateLimitRequest(router, "/limited", remoteAddr).Code) assert.Equal(t, http.StatusNoContent, performRateLimitRequest(router, "/limited", remoteAddr).Code)
assert.Equal(t, http.StatusTooManyRequests, performRateLimitRequest(router, "/limited", remoteAddr).Code) limitedResponse := performRateLimitRequest(router, "/limited", remoteAddr)
assert.Equal(t, http.StatusTooManyRequests, limitedResponse.Code)
assert.Equal(t, "37", limitedResponse.Header().Get("Retry-After"))
key := redisIPRateLimitKey("TEST", "192.0.2.10") key := redisIPRateLimitKey("TEST", "192.0.2.10")
count, err := redisServer.Get(key) count, err := redisServer.Get(key)
+24
View File
@@ -147,6 +147,30 @@ describe('authentication session coordination', () => {
assert.equal(transientCount, 1) assert.equal(transientCount, 1)
}) })
test('a rate limited refresh remains retryable without clearing the session', async () => {
let transientCount = 0
let clearCount = 0
const runtime: AuthRefreshRuntime = {
request: async () => ({ status: 429 }),
getExpectedSID: () => bundle.session.sid,
parseBundle: () => null,
acceptBundle: () => undefined,
clear: () => {
clearCount += 1
},
markTransient: () => {
transientCount += 1
},
wait: async () => undefined,
}
const outcome = await createRefreshRunner(runtime)()
assert.equal(outcome.kind, 'transient_error')
assert.equal(clearCount, 0)
assert.equal(transientCount, 1)
})
test('an exhausted refresh race clears the unusable local session', async () => { test('an exhausted refresh race clears the unusable local session', async () => {
const requestedDelays: number[] = [] const requestedDelays: number[] = []
const clears: Array<[boolean, string | undefined]> = [] const clears: Array<[boolean, string | undefined]> = []
+1 -1
View File
@@ -253,7 +253,7 @@ export function createRefreshRunner(
return { kind: 'anonymous' } return { kind: 'anonymous' }
} }
if (!response.status || response.status >= 500) { if (!response.status || response.status >= 500 || response.status === 429) {
runtime.markTransient() runtime.markTransient()
return { return {
kind: 'transient_error', kind: 'transient_error',