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() } }