ratelimit_redis.go 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136
  1. // Package middleware
  2. package middleware
  3. import (
  4. "context"
  5. "net/http"
  6. "strconv"
  7. "strings"
  8. "time"
  9. "github.com/redis/go-redis/v9"
  10. "github.com/ulule/limiter/v3"
  11. limiterRedis "github.com/ulule/limiter/v3/drivers/store/redis"
  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 := limiterRedis.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. // GetClientIP извлекает реальный IP клиента из заголовков X-Forwarded-For,
  65. // X-Real-IP или RemoteAddr. Используется для ключей rate limiter по IP.
  66. func GetClientIP(r *http.Request) string {
  67. if fwd := r.Header.Get("X-Forwarded-For"); fwd != "" {
  68. parts := strings.Split(fwd, ",")
  69. return strings.TrimSpace(parts[0])
  70. }
  71. if realIP := r.Header.Get("X-Real-IP"); realIP != "" {
  72. return realIP
  73. }
  74. addr := r.RemoteAddr
  75. if idx := strings.LastIndex(addr, ":"); idx != -1 {
  76. return addr[:idx]
  77. }
  78. return addr
  79. }
  80. func KeyByIP(prefix string) func(*http.Request) string {
  81. return func(r *http.Request) string {
  82. return prefix + ":" + GetClientIP(r)
  83. }
  84. }
  85. func KeyByUserID(prefix string) func(*http.Request) string {
  86. return func(r *http.Request) string {
  87. userID := GetUserID(r.Context())
  88. if userID != "" {
  89. return prefix + ":user:" + userID
  90. }
  91. return prefix + ":ip:" + GetClientIP(r)
  92. }
  93. }
  94. func RateLimitAuthEndpoints(redisClient *redis.Client) (*RedisRateLimiter, error) {
  95. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  96. Rate: limiter.Rate{Period: time.Minute, Limit: 10},
  97. KeyFunc: KeyByIP("auth"),
  98. })
  99. }
  100. func RateLimitAPIRead(redisClient *redis.Client) (*RedisRateLimiter, error) {
  101. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  102. Rate: limiter.Rate{Period: time.Minute, Limit: 60},
  103. KeyFunc: KeyByUserID("api_read"),
  104. })
  105. }
  106. func RateLimitAPIWrite(redisClient *redis.Client) (*RedisRateLimiter, error) {
  107. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  108. Rate: limiter.Rate{Period: time.Minute, Limit: 10},
  109. KeyFunc: KeyByUserID("api_write"),
  110. })
  111. }
  112. func RateLimitAdmin(redisClient *redis.Client) (*RedisRateLimiter, error) {
  113. return NewRedisRateLimiter(redisClient, RateLimitConfig{
  114. Rate: limiter.Rate{Period: time.Minute, Limit: 100},
  115. KeyFunc: KeyByUserID("admin"),
  116. })
  117. }