Files
oci-portal/internal/api/ratelimit.go
T
2026-07-09 19:18:04 +08:00

86 lines
2.2 KiB
Go

package api
import (
"net/http"
"sync"
"sync/atomic"
"time"
"github.com/gin-gonic/gin"
"golang.org/x/time/rate"
"oci-portal/internal/service"
)
// ipEvictInterval 是陈旧限速条目的回收周期。
const ipEvictInterval = 10 * time.Minute
// ipRateLimiter 维护每 IP 令牌桶;速率与突发取安全设置快照,
// 设置变更后既有桶就地调参;陈旧条目由后台 ticker 惰性回收。
type ipRateLimiter struct {
mu sync.RWMutex
limiters map[string]*ipEntry
}
type ipEntry struct {
limiter *rate.Limiter
// lastSeen 是 UnixNano 时间戳;请求路径无锁写、回收循环无锁读
lastSeen atomic.Int64
}
func newIPRateLimiter() *ipRateLimiter {
l := &ipRateLimiter{limiters: make(map[string]*ipEntry)}
go l.evictLoop()
return l
}
// get 返回 ip 对应的令牌桶,参数与当前配置不一致时就地调整。
// lastSeen 用原子时间戳:并发请求与回收循环两侧无锁读写,避免数据竞态。
func (l *ipRateLimiter) get(ip string, rps, burst int) *rate.Limiter {
l.mu.RLock()
entry, ok := l.limiters[ip]
l.mu.RUnlock()
if !ok {
l.mu.Lock()
if entry, ok = l.limiters[ip]; !ok {
entry = &ipEntry{limiter: rate.NewLimiter(rate.Limit(rps), burst)}
l.limiters[ip] = entry
}
l.mu.Unlock()
}
entry.lastSeen.Store(time.Now().UnixNano())
if entry.limiter.Limit() != rate.Limit(rps) || entry.limiter.Burst() != burst {
entry.limiter.SetLimit(rate.Limit(rps))
entry.limiter.SetBurst(burst)
}
return entry.limiter
}
func (l *ipRateLimiter) evictLoop() {
ticker := time.NewTicker(ipEvictInterval)
defer ticker.Stop()
for range ticker.C {
l.mu.Lock()
cutoff := time.Now().Add(-ipEvictInterval).UnixNano()
for ip, e := range l.limiters {
if e.lastSeen.Load() < cutoff {
delete(l.limiters, ip)
}
}
l.mu.Unlock()
}
}
// IPRateMiddleware 返回全局 IP 限速中间件;超限返回 429,参数随安全设置热更新。
func IPRateMiddleware(settings *service.SettingService) gin.HandlerFunc {
lm := newIPRateLimiter()
return func(c *gin.Context) {
sec := settings.SecurityCached()
if !lm.get(requestIP(c), sec.IPRateRPS, sec.IPRateBurst).Allow() {
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{"error": "rate limit exceeded"})
return
}
c.Next()
}
}