| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160 |
- // 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,
- })
- }
|