| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677 |
- // Package middleware
- package middleware
- import (
- "net/http"
- "sync"
- "time"
- "golang.org/x/time/rate"
- )
- type RateLimiter struct {
- visitors map[string]*rate.Limiter
- mu sync.Mutex
- rate rate.Limit
- burst int
- lastSeen map[string]time.Time
- stopCh chan struct{}
- }
- func NewRateLimiter(r rate.Limit, burst int) *RateLimiter {
- rl := &RateLimiter{
- visitors: make(map[string]*rate.Limiter),
- rate: r,
- burst: burst,
- lastSeen: make(map[string]time.Time),
- stopCh: make(chan struct{}),
- }
- go rl.cleanup()
- return rl
- }
- func (rl *RateLimiter) Stop() {
- close(rl.stopCh)
- }
- func (rl *RateLimiter) cleanup() {
- ticker := time.NewTicker(10 * time.Minute)
- defer ticker.Stop()
- for {
- select {
- case <-rl.stopCh:
- return
- case <-ticker.C:
- rl.mu.Lock()
- for ip, last := range rl.lastSeen {
- if time.Since(last) > 30*time.Minute {
- delete(rl.visitors, ip)
- delete(rl.lastSeen, ip)
- }
- }
- rl.mu.Unlock()
- }
- }
- }
- func (rl *RateLimiter) Middleware() func(http.Handler) http.Handler {
- return func(next http.Handler) http.Handler {
- return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- ip := GetClientIP(r)
- rl.mu.Lock()
- defer rl.mu.Unlock()
- limiter, ok := rl.visitors[ip]
- if !ok {
- limiter = rate.NewLimiter(rl.rate, rl.burst)
- rl.visitors[ip] = limiter
- }
- rl.lastSeen[ip] = time.Now()
- if !limiter.Allow() {
- http.Error(w, `{"error":"rate limit exceeded"}`, http.StatusTooManyRequests)
- return
- }
- next.ServeHTTP(w, r)
- })
- }
- }
|