// Package middleware package middleware import ( "context" "log/slog" "net/http" "strconv" "strings" "time" "github.com/redis/go-redis/v9" "github.com/ulule/limiter/v3" limiterRedis "github.com/ulule/limiter/v3/drivers/store/redis" ) type RedisRateLimiter struct { instance *limiter.Limiter keyFunc func(*http.Request) string failOpen bool logger *slog.Logger } type RateLimitConfig struct { Rate limiter.Rate KeyFunc func(*http.Request) string FailOpen bool Logger *slog.Logger } func NewRedisRateLimiter(redisClient *redis.Client, config RateLimitConfig) (*RedisRateLimiter, error) { store, err := limiterRedis.NewStoreWithOptions(redisClient, limiter.StoreOptions{ Prefix: "ratelimit", MaxRetry: 3, }) if err != nil { return nil, err } instance := limiter.New(store, config.Rate) return &RedisRateLimiter{ instance: instance, keyFunc: config.KeyFunc, failOpen: config.FailOpen, logger: config.Logger, }, nil } func (r *RedisRateLimiter) Middleware() func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { key := r.keyFunc(req) ctx, cancel := context.WithTimeout(req.Context(), 100*time.Millisecond) defer cancel() limit, err := r.instance.Get(ctx, key) if err != nil { if r.logger != nil { r.logger.ErrorContext(req.Context(), "rate limiter error", "error", err, "key", key) } if r.failOpen { next.ServeHTTP(w, req) return } // Fail closed - return 503 Service Unavailable w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusServiceUnavailable) _, _ = w.Write([]byte(`{"error":"rate limiter unavailable"}`)) return } w.Header().Set("X-RateLimit-Limit", strconv.FormatInt(limit.Limit, 10)) w.Header().Set("X-RateLimit-Remaining", strconv.FormatInt(limit.Remaining, 10)) w.Header().Set("X-RateLimit-Reset", strconv.FormatInt(limit.Reset, 10)) if limit.Reached { writeRateLimitError(w, limit.Reset) return } next.ServeHTTP(w, req) }) } } func writeRateLimitError(w http.ResponseWriter, reset int64) { w.Header().Set("Content-Type", "application/json") w.Header().Set("Retry-After", strconv.FormatInt(reset, 10)) w.WriteHeader(http.StatusTooManyRequests) _, _ = w.Write([]byte(`{"error":"rate limit exceeded"}`)) } // GetClientIP извлекает реальный IP клиента из заголовков X-Forwarded-For, // X-Real-IP или RemoteAddr. Используется для ключей rate limiter по IP. func GetClientIP(r *http.Request) string { if fwd := r.Header.Get("X-Forwarded-For"); fwd != "" { parts := strings.Split(fwd, ",") return strings.TrimSpace(parts[0]) } if realIP := r.Header.Get("X-Real-IP"); realIP != "" { return realIP } addr := r.RemoteAddr if idx := strings.LastIndex(addr, ":"); idx != -1 { return addr[:idx] } return addr } func KeyByIP(prefix string) func(*http.Request) string { return func(r *http.Request) string { return prefix + ":" + GetClientIP(r) } } func KeyByUserID(prefix string) func(*http.Request) string { return func(r *http.Request) string { userID := GetUserID(r.Context()) if userID != "" { return prefix + ":user:" + userID } return prefix + ":ip:" + GetClientIP(r) } } func RateLimitAuthEndpoints(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 30}, KeyFunc: KeyByIP("auth"), FailOpen: failOpen, Logger: logger, }) } func RateLimitAPIRead(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 60}, KeyFunc: KeyByUserID("api_read"), FailOpen: failOpen, Logger: logger, }) } func RateLimitAPIWrite(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 60}, KeyFunc: KeyByUserID("api_write"), FailOpen: failOpen, Logger: logger, }) } func RateLimitAdmin(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 100}, KeyFunc: KeyByUserID("admin"), FailOpen: failOpen, Logger: logger, }) }