// Package middleware package middleware import ( "context" "net/http" "strconv" "time" "github.com/redis/go-redis/v9" "github.com/ulule/limiter/v3" "github.com/ulule/limiter/v3/drivers/store/redis" chimiddleware "github.com/go-chi/chi/v5/middleware" ) type RedisRateLimiter struct { instance *limiter.Limiter keyFunc func(*http.Request) string } type RateLimitConfig struct { Rate limiter.Rate KeyFunc func(*http.Request) string } func NewRedisRateLimiter(redisClient *redis.Client, config RateLimitConfig) (*RedisRateLimiter, error) { store, err := redis.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, }, 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 { // Fail open - log error but allow request through next.ServeHTTP(w, req) 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"}`)) } func GetClientIP(r *http.Request) string { return chimiddleware.GetIP(r) } 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) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 10}, KeyFunc: KeyByIP("auth"), }) } func RateLimitAPIRead(redisClient *redis.Client) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 60}, KeyFunc: KeyByUserID("api_read"), }) } func RateLimitAPIWrite(redisClient *redis.Client) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 10}, KeyFunc: KeyByUserID("api_write"), }) } func RateLimitAdmin(redisClient *redis.Client) (*RedisRateLimiter, error) { return NewRedisRateLimiter(redisClient, RateLimitConfig{ Rate: limiter.Rate{Period: time.Minute, Limit: 100}, KeyFunc: KeyByUserID("admin"), }) }