package middleware import ( "fmt" "net/http" "strings" "sync" "time" "golang.org/x/time/rate" ) // RateLimiter hanterar rate limiting per endpoint type RateLimiter struct { limiters map[string]*rate.Limiter mu sync.RWMutex // Default rates defaultRate rate.Limit defaultBurst int } // NewRateLimiter skapar en ny rate limiter func NewRateLimiter() *RateLimiter { return &RateLimiter{ limiters: make(map[string]*rate.Limiter), defaultRate: rate.Every(time.Second), // 1 request per second defaultBurst: 10, } } // getLimiter hämtar eller skapar en limiter för en given nyckel func (rl *RateLimiter) getLimiter(key string, r rate.Limit, b int) *rate.Limiter { rl.mu.RLock() limiter, exists := rl.limiters[key] rl.mu.RUnlock() if exists { return limiter } rl.mu.Lock() defer rl.mu.Unlock() // Dubbelkolla efter lås limiter, exists = rl.limiters[key] if exists { return limiter } limiter = rate.NewLimiter(r, b) rl.limiters[key] = limiter return limiter } // RateLimit middleware med anpassade gränser per endpoint func RateLimit(rl *RateLimiter) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // Skapa nyckel baserat på IP + endpoint clientIP := getClientIP(r) endpoint := r.Method + " " + r.URL.Path key := clientIP + ":" + endpoint // Anpassade gränser per endpoint var limit rate.Limit var burst int switch { // Login: 5 försök per minut case strings.Contains(endpoint, "/auth/login"): limit = rate.Every(12 * time.Second) burst = 5 // API endpoints: 100 per minut case strings.HasPrefix(r.URL.Path, "/api/v1/"): limit = rate.Every(600 * time.Millisecond) burst = 100 // Default: 10 per sekund default: limit = rl.defaultRate burst = rl.defaultBurst } limiter := rl.getLimiter(key, limit, burst) if !limiter.Allow() { w.Header().Set("Retry-After", "60") w.Header().Set("X-RateLimit-Limit", fmt.Sprintf("%v", limit)) w.Header().Set("X-RateLimit-Remaining", "0") writeError(w, http.StatusTooManyRequests, "rate limit exceeded") return } next.ServeHTTP(w, r) }) } } // getClientIP hämtar klientens IP-adress func getClientIP(r *http.Request) string { // Kolla X-Forwarded-For header (för reverse proxies) xff := r.Header.Get("X-Forwarded-For") if xff != "" { // Ta första IP:et parts := strings.Split(xff, ",") if len(parts) > 0 { return strings.TrimSpace(parts[0]) } } // Kolla X-Real-IP xri := r.Header.Get("X-Real-IP") if xri != "" { return xri } // Fallback till RemoteAddr ip := r.RemoteAddr // Ta bort port if idx := strings.LastIndex(ip, ":"); idx != -1 { ip = ip[:idx] } return ip }