ratelimit_redis.go 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123
  1. // Package middleware
  2. package middleware
  3. import (
  4. "context"
  5. "net/http"
  6. "strconv"
  7. "time"
  8. "github.com/redis/go-redis/v9"
  9. "github.com/ulule/limiter/v3"
  10. "github.com/ulule/limiter/v3/drivers/store/redis"
  11. chimiddleware "github.com/go-chi/chi/v5/middleware"
  12. )
  13. type RedisRateLimiter struct {
  14. instance *limiter.Limiter
  15. keyFunc func(*http.Request) string
  16. }
  17. type RateLimitConfig struct {
  18. Rate limiter.Rate
  19. KeyFunc func(*http.Request) string
  20. }
  21. func NewRedisRateLimiter(redisClient *redis.Client, config RateLimitConfig) (*RedisRateLimiter, error) {
  22. store, err := redis.NewStoreWithOptions(redisClient, limiter.StoreOptions{
  23. Prefix: "ratelimit",
  24. MaxRetry: 3,
  25. })
  26. if err != nil {
  27. return nil, err
  28. }
  29. instance := limiter.New(store, config.Rate)
  30. return &RedisRateLimiter{
  31. instance: instance,
  32. keyFunc: config.KeyFunc,
  33. }, nil
  34. }
  35. func (r *RedisRateLimiter) Middleware() func(http.Handler) http.Handler {
  36. return func(next http.Handler) http.Handler {
  37. return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
  38. key := r.keyFunc(req)
  39. ctx, cancel := context.WithTimeout(req.Context(), 100*time.Millisecond)
  40. defer cancel()
  41. limit, err := r.instance.Get(ctx, key)
  42. if err != nil {
  43. // Fail open - log error but allow request through
  44. next.ServeHTTP(w, req)
  45. return
  46. }
  47. w.Header().Set("X-RateLimit-Limit", strconv.FormatInt(limit.Limit, 10))
  48. w.Header().Set("X-RateLimit-Remaining", strconv.FormatInt(limit.Remaining, 10))
  49. w.Header().Set("X-RateLimit-Reset", strconv.FormatInt(limit.Reset, 10))
  50. if limit.Reached {
  51. writeRateLimitError(w, limit.Reset)
  52. return
  53. }
  54. next.ServeHTTP(w, req)
  55. })
  56. }
  57. }
  58. func writeRateLimitError(w http.ResponseWriter, reset int64) {
  59. w.Header().Set("Content-Type", "application/json")
  60. w.Header().Set("Retry-After", strconv.FormatInt(reset, 10))
  61. w.WriteHeader(http.StatusTooManyRequests)
  62. _, _ = w.Write([]byte(`{"error":"rate limit exceeded"}`))
  63. }
  64. func GetClientIP(r *http.Request) string {
  65. return chimiddleware.GetIP(r)
  66. }
  67. func KeyByIP(prefix string) func(*http.Request) string {
  68. return func(r *http.Request) string {
  69. return prefix + ":" + GetClientIP(r)
  70. }
  71. }
  72. func KeyByUserID(prefix string) func(*http.Request) string {
  73. return func(r *http.Request) string {
  74. userID := GetUserID(r.Context())
  75. if userID != "" {
  76. return prefix + ":user:" + userID
  77. }
  78. return prefix + ":ip:" + GetClientIP(r)
  79. }
  80. }
  81. func RateLimitAuthEndpoints(redisClient *redis.Client) (*RedisRateLimiter, error) {
  82. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  83. Rate: limiter.Rate{Period: time.Minute, Limit: 10},
  84. KeyFunc: KeyByIP("auth"),
  85. })
  86. }
  87. func RateLimitAPIRead(redisClient *redis.Client) (*RedisRateLimiter, error) {
  88. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  89. Rate: limiter.Rate{Period: time.Minute, Limit: 60},
  90. KeyFunc: KeyByUserID("api_read"),
  91. })
  92. }
  93. func RateLimitAPIWrite(redisClient *redis.Client) (*RedisRateLimiter, error) {
  94. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  95. Rate: limiter.Rate{Period: time.Minute, Limit: 10},
  96. KeyFunc: KeyByUserID("api_write"),
  97. })
  98. }
  99. func RateLimitAdmin(redisClient *redis.Client) (*RedisRateLimiter, error) {
  100. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  101. Rate: limiter.Rate{Period: time.Minute, Limit: 100},
  102. KeyFunc: KeyByUserID("admin"),
  103. })
  104. }