package auth import ( "context" "crypto/rand" "encoding/hex" "fmt" "time" "github.com/redis/go-redis/v9" ) // RateLimiter implements per-scope sliding-window rate limiting and failure // counters on top of Redis. It stores nothing but short-lived counters — no // connection or behaviour logs — per the no-log baseline (doc/06 §4). // // Sliding-window keys are ZSETs named `rl:{scope}:{key}` whose members are // timestamped attempts; failure counters are plain integer keys. type RateLimiter struct { rdb *redis.Client now func() time.Time } // NewRateLimiter builds a RateLimiter. now may be nil (defaults to time.Now). func NewRateLimiter(rdb *redis.Client, now func() time.Time) *RateLimiter { if now == nil { now = time.Now } return &RateLimiter{rdb: rdb, now: now} } // slidingWindow is a single atomic Lua script: // - drops members older than the window, // - if the live count already reached the limit, returns {0, oldestScore}, // - otherwise records the attempt and returns {1, 0}. // // All times are unix-milliseconds. var slidingWindow = redis.NewScript(` local key = KEYS[1] local now = tonumber(ARGV[1]) local window = tonumber(ARGV[2]) local limit = tonumber(ARGV[3]) local member = ARGV[4] redis.call('ZREMRANGEBYSCORE', key, 0, now - window) local count = redis.call('ZCARD', key) if count >= limit then local oldest = redis.call('ZRANGE', key, 0, 0, 'WITHSCORES') return {0, oldest[2]} end redis.call('ZADD', key, now, member) redis.call('PEXPIRE', key, window) return {1, 0} `) // rlKey builds the canonical `rl:{scope}:{key}` Redis key. func rlKey(scope, key string) string { return fmt.Sprintf("rl:%s:%s", scope, key) } // Allow records an attempt in the sliding window for (scope, key). It returns // allowed=false together with the Retry-After duration when the limit within // window has been reached. limit is the maximum number of attempts permitted // inside window. func (rl *RateLimiter) Allow(ctx context.Context, scope, key string, limit int, window time.Duration) (allowed bool, retryAfter time.Duration, err error) { nowMs := rl.now().UnixMilli() winMs := window.Milliseconds() member, err := uniqueMember(nowMs) if err != nil { return false, 0, err } res, err := slidingWindow.Run(ctx, rl.rdb, []string{rlKey(scope, key)}, nowMs, winMs, limit, member).Result() if err != nil { return false, 0, fmt.Errorf("auth: ratelimit run: %w", err) } vals, ok := res.([]interface{}) if !ok || len(vals) != 2 { return false, 0, fmt.Errorf("auth: ratelimit unexpected result %v", res) } ok1, _ := vals[0].(int64) if ok1 == 1 { return true, 0, nil } // Denied: compute how long until the oldest attempt leaves the window. oldest := toInt64(vals[1]) ra := time.Duration(oldest+winMs-nowMs) * time.Millisecond if ra < time.Second { ra = time.Second } return false, ra, nil } // failKey builds the failure-counter key, e.g. `rl:login:{email}`. func failKey(scope, key string) string { return rlKey(scope, key) } // RecordFailure increments the failure counter for (scope, key), (re)setting // its TTL to window on every increment so the lock slides forward while abuse // continues. It returns the new count. func (rl *RateLimiter) RecordFailure(ctx context.Context, scope, key string, window time.Duration) (int64, error) { k := failKey(scope, key) pipe := rl.rdb.Pipeline() incr := pipe.Incr(ctx, k) pipe.Expire(ctx, k, window) if _, err := pipe.Exec(ctx); err != nil { return 0, fmt.Errorf("auth: record failure: %w", err) } return incr.Val(), nil } // FailureCount returns the current failure count and the remaining TTL (lock // window) for (scope, key). count is 0 when no counter exists. func (rl *RateLimiter) FailureCount(ctx context.Context, scope, key string) (count int64, ttl time.Duration, err error) { k := failKey(scope, key) pipe := rl.rdb.Pipeline() get := pipe.Get(ctx, k) pttl := pipe.PTTL(ctx, k) if _, err := pipe.Exec(ctx); err != nil && err != redis.Nil { return 0, 0, fmt.Errorf("auth: failure count: %w", err) } n, _ := get.Int64() d := pttl.Val() if d < 0 { d = 0 } return n, d, nil } // ClearFailures removes the failure counter after a successful authentication. func (rl *RateLimiter) ClearFailures(ctx context.Context, scope, key string) error { return rl.rdb.Del(ctx, failKey(scope, key)).Err() } // uniqueMember returns a ZSET member that is unique even for attempts that share // the same millisecond timestamp (timestamp prefix keeps ordering stable). func uniqueMember(nowMs int64) (string, error) { var b [8]byte if _, err := rand.Read(b[:]); err != nil { return "", fmt.Errorf("auth: ratelimit member: %w", err) } return fmt.Sprintf("%d-%s", nowMs, hex.EncodeToString(b[:])), nil } // toInt64 coerces a Redis Lua return value (which may be int64 or string) to int64. func toInt64(v interface{}) int64 { switch t := v.(type) { case int64: return t case string: var n int64 _, _ = fmt.Sscanf(t, "%d", &n) return n default: return 0 } }