112 lines
2.6 KiB
Go
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()
|
|
}
|
|
}
|