Files
Mycontent/soul-api/internal/middleware/ratelimit.go
2026-04-13 14:32:32 +08:00

120 lines
2.6 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package middleware
import (
"net/http"
"sync"
"time"
"golang.org/x/time/rate"
"github.com/gin-gonic/gin"
)
// RateLimiter 按 IP 的限流器
type RateLimiter struct {
mu sync.Mutex
clients map[string]*rate.Limiter
r rate.Limit
b int
}
// NewRateLimiter 创建限流中间件r 每秒请求数b 突发容量
func NewRateLimiter(r rate.Limit, b int) *RateLimiter {
return &RateLimiter{
clients: make(map[string]*rate.Limiter),
r: r,
b: b,
}
}
// getLimiter 获取或创建该 key 的 limiter
func (rl *RateLimiter) getLimiter(key string) *rate.Limiter {
rl.mu.Lock()
defer rl.mu.Unlock()
if lim, ok := rl.clients[key]; ok {
return lim
}
lim := rate.NewLimiter(rl.r, rl.b)
rl.clients[key] = lim
return lim
}
// Middleware 返回 Gin 限流中间件(按客户端 IP
func (rl *RateLimiter) Middleware() gin.HandlerFunc {
return func(c *gin.Context) {
key := c.ClientIP()
lim := rl.getLimiter(key)
if !lim.Allow() {
c.AbortWithStatus(http.StatusTooManyRequests)
return
}
c.Next()
}
}
// Cleanup 定期清理过期 limiter可选避免 map 无限增长)
func (rl *RateLimiter) Cleanup(interval time.Duration) {
ticker := time.NewTicker(interval)
go func() {
for range ticker.C {
rl.mu.Lock()
rl.clients = make(map[string]*rate.Limiter)
rl.mu.Unlock()
}
}()
}
// MatchRateLimiter 针对匹配接口的独立限流:按 IP+UserID 组合键,每分钟 maxPerMin 次
type MatchRateLimiter struct {
mu sync.Mutex
clients map[string]*rate.Limiter
r rate.Limit
b int
}
func NewMatchRateLimiter(maxPerMin int) *MatchRateLimiter {
return &MatchRateLimiter{
clients: make(map[string]*rate.Limiter),
r: rate.Limit(float64(maxPerMin) / 60.0),
b: maxPerMin,
}
}
func (ml *MatchRateLimiter) getLimiter(key string) *rate.Limiter {
ml.mu.Lock()
defer ml.mu.Unlock()
if lim, ok := ml.clients[key]; ok {
return lim
}
lim := rate.NewLimiter(ml.r, ml.b)
ml.clients[key] = lim
return lim
}
func (ml *MatchRateLimiter) Middleware() gin.HandlerFunc {
return func(c *gin.Context) {
key := c.ClientIP()
lim := ml.getLimiter(key)
if !lim.Allow() {
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{
"success": false,
"code": "RATE_LIMITED",
"message": "操作过于频繁,请稍后再试",
})
return
}
c.Next()
}
}
func (ml *MatchRateLimiter) Cleanup(interval time.Duration) {
ticker := time.NewTicker(interval)
go func() {
for range ticker.C {
ml.mu.Lock()
ml.clients = make(map[string]*rate.Limiter)
ml.mu.Unlock()
}
}()
}