diff --git a/internal/middleware/ratelimit.go b/internal/middleware/ratelimit.go new file mode 100644 index 0000000..1ac8a57 --- /dev/null +++ b/internal/middleware/ratelimit.go @@ -0,0 +1,59 @@ +package middleware + +import ( + "net/http" + "sync" + + "github.com/gin-gonic/gin" + "golang.org/x/time/rate" + + "gotest/internal/config" +) + +// RateLimit 限流中间件,基于令牌桶算法按客户端 IP 进行限流 +func RateLimit(cfg *config.RateLimitConfig) gin.HandlerFunc { + // 未启用限流直接放行 + if !cfg.Enabled { + return func(c *gin.Context) { + c.Next() + } + } + + // 计算限流参数,对异常配置提供默认值兜底 + rps := rate.Limit(cfg.RequestsPerSecond) + if rps <= 0 { + rps = 1 + } + burst := cfg.BurstSize + if burst <= 0 { + burst = cfg.RequestsPerSecond + } + + var ( + mu sync.Mutex + limiters = make(map[string]*rate.Limiter) + ) + + // getLimiter 获取或创建指定键(客户端 IP)对应的令牌桶限流器 + getLimiter := func(key string) *rate.Limiter { + mu.Lock() + defer mu.Unlock() + if l, ok := limiters[key]; ok { + return l + } + l := rate.NewLimiter(rps, burst) + limiters[key] = l + return l + } + + return func(c *gin.Context) { + key := c.ClientIP() + // 令牌不足即拒绝请求 + if !getLimiter(key).Allow() { + c.JSON(http.StatusTooManyRequests, gin.H{"code": 1, "message": "请求过于频繁,请稍后再试", "data": nil}) + c.Abort() + return + } + c.Next() + } +}