| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123 |
- // Package middleware
- package middleware
- import (
- "context"
- "net/http"
- "strconv"
- "time"
- "github.com/redis/go-redis/v9"
- "github.com/ulule/limiter/v3"
- limiterRedis "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 := 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,
- }, 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"),
- })
- }
|