mengstack-api/internal/app/middleware/ratelimit_redis.go
MengStack Dev 42527b616d
Some checks are pending
CI / Build & Test (push) Waiting to run
fix: increase anonymous rate limit to 100/min for demo
2026-10-03 04:11:55 +08:00

112 lines
2.6 KiB
Go

package middleware
import (
"context"
"fmt"
"net/http"
"strconv"
"time"
"mengstack/internal/kernel/errors"
"mengstack/internal/kernel/response"
"github.com/gin-gonic/gin"
"github.com/redis/go-redis/v9"
)
const (
rateLimitHeader = "X-RateLimit-Limit"
rateLimitRemainHeader = "X-RateLimit-Remaining"
rateLimitResetHeader = "X-RateLimit-Reset"
)
type RateLimitTier struct {
Limit int
Window time.Duration
}
var (
TierAnonymous = RateLimitTier{Limit: 100, Window: time.Minute}
TierAuth = RateLimitTier{Limit: 100, Window: time.Minute}
TierAdmin = RateLimitTier{Limit: 500, Window: time.Minute}
)
type redisRateLimiter struct {
rdb *redis.Client
}
func newRedisRateLimiter(rdb *redis.Client) *redisRateLimiter {
return &redisRateLimiter{rdb: rdb}
}
func (rl *redisRateLimiter) allow(ctx context.Context, key string, tier RateLimitTier) (allowed bool, limit, remaining int, resetSec int) {
now := time.Now()
windowStart := now.Add(-tier.Window).UnixMicro()
pipe := rl.rdb.Pipeline()
pipe.ZRemRangeByScore(ctx, key, "0", fmt.Sprintf("%d", windowStart))
pipe.ZAdd(ctx, key, redis.Z{Score: float64(now.UnixMicro()), Member: fmt.Sprintf("%d", now.UnixNano())})
pipe.ZCard(ctx, key)
pipe.Expire(ctx, key, tier.Window+time.Second)
results, err := pipe.Exec(ctx)
if err != nil {
return true, tier.Limit, tier.Limit - 1, int(tier.Window.Seconds())
}
count := results[2].(*redis.IntCmd).Val()
limit = tier.Limit
resetSec = int(tier.Window.Seconds())
if count > int64(limit) {
rl.rdb.ZRem(ctx, key, fmt.Sprintf("%d", now.UnixNano()))
return false, limit, 0, resetSec
}
remaining = limit - int(count)
if remaining < 0 {
remaining = 0
}
return true, limit, remaining, resetSec
}
func RateLimitRedis(rdb *redis.Client) gin.HandlerFunc {
rl := newRedisRateLimiter(rdb)
return func(c *gin.Context) {
ip := c.ClientIP()
userID := c.GetString("user_id")
role := c.GetString("role")
var key string
var tier RateLimitTier
if userID != "" {
switch role {
case "admin":
tier = TierAdmin
default:
tier = TierAuth
}
key = fmt.Sprintf("ratelimit:user:%s", userID)
} else {
tier = TierAnonymous
key = fmt.Sprintf("ratelimit:ip:%s", ip)
}
allowed, limit, remaining, resetSec := rl.allow(c.Request.Context(), key, tier)
c.Header(rateLimitHeader, strconv.Itoa(limit))
c.Header(rateLimitRemainHeader, strconv.Itoa(remaining))
c.Header(rateLimitResetHeader, strconv.Itoa(resetSec))
if !allowed {
response.Fail(c, errors.New("RATE_LIMITED", "too many requests", http.StatusTooManyRequests))
c.Abort()
return
}
c.Next()
}
}