Add internal/middleware/ratelimit.go
This commit is contained in:
59
internal/middleware/ratelimit.go
Normal file
59
internal/middleware/ratelimit.go
Normal file
@@ -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()
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user