60 lines
1.2 KiB
Go
60 lines
1.2 KiB
Go
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()
|
|
}
|
|
}
|