ratelimit_redis.go 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160
  1. // Package middleware
  2. package middleware
  3. import (
  4. "context"
  5. "log/slog"
  6. "net/http"
  7. "strconv"
  8. "strings"
  9. "time"
  10. "github.com/redis/go-redis/v9"
  11. "github.com/ulule/limiter/v3"
  12. limiterRedis "github.com/ulule/limiter/v3/drivers/store/redis"
  13. )
  14. type RedisRateLimiter struct {
  15. instance *limiter.Limiter
  16. keyFunc func(*http.Request) string
  17. failOpen bool
  18. logger *slog.Logger
  19. }
  20. type RateLimitConfig struct {
  21. Rate limiter.Rate
  22. KeyFunc func(*http.Request) string
  23. FailOpen bool
  24. Logger *slog.Logger
  25. }
  26. func NewRedisRateLimiter(redisClient *redis.Client, config RateLimitConfig) (*RedisRateLimiter, error) {
  27. store, err := limiterRedis.NewStoreWithOptions(redisClient, limiter.StoreOptions{
  28. Prefix: "ratelimit",
  29. MaxRetry: 3,
  30. })
  31. if err != nil {
  32. return nil, err
  33. }
  34. instance := limiter.New(store, config.Rate)
  35. return &RedisRateLimiter{
  36. instance: instance,
  37. keyFunc: config.KeyFunc,
  38. failOpen: config.FailOpen,
  39. logger: config.Logger,
  40. }, nil
  41. }
  42. func (r *RedisRateLimiter) Middleware() func(http.Handler) http.Handler {
  43. return func(next http.Handler) http.Handler {
  44. return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
  45. key := r.keyFunc(req)
  46. ctx, cancel := context.WithTimeout(req.Context(), 100*time.Millisecond)
  47. defer cancel()
  48. limit, err := r.instance.Get(ctx, key)
  49. if err != nil {
  50. if r.logger != nil {
  51. r.logger.ErrorContext(req.Context(), "rate limiter error", "error", err, "key", key)
  52. }
  53. if r.failOpen {
  54. next.ServeHTTP(w, req)
  55. return
  56. }
  57. // Fail closed - return 503 Service Unavailable
  58. w.Header().Set("Content-Type", "application/json")
  59. w.WriteHeader(http.StatusServiceUnavailable)
  60. _, _ = w.Write([]byte(`{"error":"rate limiter unavailable"}`))
  61. return
  62. }
  63. w.Header().Set("X-RateLimit-Limit", strconv.FormatInt(limit.Limit, 10))
  64. w.Header().Set("X-RateLimit-Remaining", strconv.FormatInt(limit.Remaining, 10))
  65. w.Header().Set("X-RateLimit-Reset", strconv.FormatInt(limit.Reset, 10))
  66. if limit.Reached {
  67. writeRateLimitError(w, limit.Reset)
  68. return
  69. }
  70. next.ServeHTTP(w, req)
  71. })
  72. }
  73. }
  74. func writeRateLimitError(w http.ResponseWriter, reset int64) {
  75. w.Header().Set("Content-Type", "application/json")
  76. w.Header().Set("Retry-After", strconv.FormatInt(reset, 10))
  77. w.WriteHeader(http.StatusTooManyRequests)
  78. _, _ = w.Write([]byte(`{"error":"rate limit exceeded"}`))
  79. }
  80. // GetClientIP извлекает реальный IP клиента из заголовков X-Forwarded-For,
  81. // X-Real-IP или RemoteAddr. Используется для ключей rate limiter по IP.
  82. func GetClientIP(r *http.Request) string {
  83. if fwd := r.Header.Get("X-Forwarded-For"); fwd != "" {
  84. parts := strings.Split(fwd, ",")
  85. return strings.TrimSpace(parts[0])
  86. }
  87. if realIP := r.Header.Get("X-Real-IP"); realIP != "" {
  88. return realIP
  89. }
  90. addr := r.RemoteAddr
  91. if idx := strings.LastIndex(addr, ":"); idx != -1 {
  92. return addr[:idx]
  93. }
  94. return addr
  95. }
  96. func KeyByIP(prefix string) func(*http.Request) string {
  97. return func(r *http.Request) string {
  98. return prefix + ":" + GetClientIP(r)
  99. }
  100. }
  101. func KeyByUserID(prefix string) func(*http.Request) string {
  102. return func(r *http.Request) string {
  103. userID := GetUserID(r.Context())
  104. if userID != "" {
  105. return prefix + ":user:" + userID
  106. }
  107. return prefix + ":ip:" + GetClientIP(r)
  108. }
  109. }
  110. func RateLimitAuthEndpoints(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) {
  111. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  112. Rate: limiter.Rate{Period: time.Minute, Limit: 30},
  113. KeyFunc: KeyByIP("auth"),
  114. FailOpen: failOpen,
  115. Logger: logger,
  116. })
  117. }
  118. func RateLimitAPIRead(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) {
  119. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  120. Rate: limiter.Rate{Period: time.Minute, Limit: 60},
  121. KeyFunc: KeyByUserID("api_read"),
  122. FailOpen: failOpen,
  123. Logger: logger,
  124. })
  125. }
  126. func RateLimitAPIWrite(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) {
  127. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  128. Rate: limiter.Rate{Period: time.Minute, Limit: 60},
  129. KeyFunc: KeyByUserID("api_write"),
  130. FailOpen: failOpen,
  131. Logger: logger,
  132. })
  133. }
  134. func RateLimitAdmin(redisClient *redis.Client, logger *slog.Logger, failOpen bool) (*RedisRateLimiter, error) {
  135. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  136. Rate: limiter.Rate{Period: time.Minute, Limit: 100},
  137. KeyFunc: KeyByUserID("admin"),
  138. FailOpen: failOpen,
  139. Logger: logger,
  140. })
  141. }