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() } }