Files
oci-portal/internal/api/ratelimit.go
T

83 lines
2.0 KiB
Go

package api
import (
"net/http"
"sync"
"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 time.Time
}
func newIPRateLimiter() *ipRateLimiter {
l := &ipRateLimiter{limiters: make(map[string]*ipEntry)}
go l.evictLoop()
return l
}
// get 返回 ip 对应的令牌桶,参数与当前配置不一致时就地调整。
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 = time.Now()
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)
for ip, e := range l.limiters {
if e.lastSeen.Before(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()
}
}