#!/usr/bin/env python3
"""Rate limiter unit tests (Phase 2 — 2d)
Tests the Redis-backed sliding window rate limiter:
- Allow under limit
- Deny over limit
- Window expiry resets counter
- Redis unavailable → fail-open
- get_remaining() accuracy
- flush() clears keys
"""
import os
import sys
import time
import pytest
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
from app.services.ratelimit import RateLimiter
class FakeRedis:
"""Minimal Redis mock for unit testing without a real Redis instance."""
def __init__(self):
self._store = {}
self._ttls = {}
def incr(self, key):
# Check TTL expiry before incrementing
if key in self._ttls and time.time() > self._ttls[key]:
del self._store[key]
del self._ttls[key]
val = self._store.get(key, 0) + 1
self._store[key] = val
return val
def expire(self, key, ttl):
self._ttls[key] = time.time() + ttl
return True
def get(self, key):
if key in self._ttls and time.time() > self._ttls[key]:
del self._store[key]
del self._ttls[key]
return None
return str(self._store.get(key, "")) or None
def ttl(self, key):
if key in self._ttls:
remaining = self._ttls[key] - time.time()
return max(0, int(remaining))
return -2
def delete(self, *keys):
for key in keys:
self._store.pop(key, None)
self._ttls.pop(key, None)
def keys(self, pattern="*"):
import fnmatch
return [k for k in self._store if fnmatch.fnmatch(k, pattern)]
def pipeline(self):
return _FakePipeline(self)
class _FakePipeline:
def __init__(self, redis):
self._redis = redis
self._ops = []
def incr(self, key):
self._ops.append(("incr", key))
return self
def expire(self, key, ttl):
self._ops.append(("expire", key, ttl))
return self
def execute(self):
results = []
for op in self._ops:
if op[0] == "incr":
results.append(self._redis.incr(op[1]))
elif op[0] == "expire":
results.append(self._redis.expire(op[1], op[2]))
self._ops = []
return results
class TestRateLimiterAllow:
"""Test basic allow/deny logic."""
@pytest.fixture
def limiter(self):
r = FakeRedis()
lim = RateLimiter(r)
return lim
def test_allow_under_limit(self, limiter):
result = limiter.allow("test:allow", max_requests=5, window=60)
assert result is True
def test_deny_over_limit(self, limiter):
# Fill up to the limit
for _ in range(5):
limiter.allow("test:deny", max_requests=5, window=60)
# 6th request should be denied
result = limiter.allow("test:deny", max_requests=5, window=60)
assert result is False
def test_exact_limit_allowed(self, limiter):
# Exactly at the limit should still be allowed
for i in range(3):
result = limiter.allow("test:exact", max_requests=3, window=60)
assert result is True
# One over
result = limiter.allow("test:exact", max_requests=3, window=60)
assert result is False
class TestRateLimiterCheck:
"""Test check() return values."""
@pytest.fixture
def limiter(self):
r = FakeRedis()
lim = RateLimiter(r)
return lim
def test_check_returns_tuple(self, limiter):
allowed, remaining, reset_at = limiter.check("test:check", max_requests=10, window=60)
assert isinstance(allowed, bool)
assert isinstance(remaining, int)
assert isinstance(reset_at, float)
def test_check_remaining_decrements(self, limiter):
_, remaining, _ = limiter.check("test:rem", max_requests=5, window=60)
assert remaining == 4
for _ in range(4):
limiter.check("test:rem", max_requests=5, window=60)
_, remaining, _ = limiter.check("test:rem", max_requests=5, window=60)
assert remaining == 0
def test_check_reset_at_future(self, limiter):
_, _, reset_at = limiter.check("test:reset", max_requests=10, window=60)
assert reset_at > time.time()
assert reset_at <= time.time() + 61
class TestRateLimiterWindowExpiry:
"""Test that counters reset after window expires."""
@pytest.fixture
def limiter(self):
r = FakeRedis()
lim = RateLimiter(r)
return lim
def test_window_expiry_resets(self, limiter):
# Fill up with a 1-second window
for _ in range(5):
limiter.allow("test:expiry", max_requests=5, window=1)
assert limiter.allow("test:expiry", max_requests=5, window=1) is False
# Wait for expiry
time.sleep(1.1)
# Should be allowed again
assert limiter.allow("test:expiry", max_requests=5, window=1) is True
class TestRateLimiterFailOpen:
"""Test fail-open behavior when Redis is unavailable."""
def test_no_redis_allow(self):
lim = RateLimiter() # No redis_client
result = lim.allow("test:noredis")
assert result is True
def test_no_redis_check(self):
lim = RateLimiter()
allowed, remaining, reset_at = lim.check("test:noredis", max_requests=10, window=60)
assert allowed is True
assert remaining == 10
assert reset_at == 0
def test_no_redis_get_remaining(self):
lim = RateLimiter()
remaining, reset_at = lim.get_remaining("test:noredis", max_requests=10, window=60)
assert remaining == 10
assert reset_at == 0
def test_no_redis_flush(self):
lim = RateLimiter()
# Should not raise
lim.flush("test:*")
class TestRateLimiterGetRemaining:
"""Test get_remaining() accuracy."""
@pytest.fixture
def limiter(self):
r = FakeRedis()
lim = RateLimiter(r)
return lim
def test_get_remaining_fresh(self, limiter):
remaining, _ = limiter.get_remaining("test:fresh", max_requests=10, window=60)
assert remaining == 10
def test_get_remaining_after_usage(self, limiter):
limiter.allow("test:used", max_requests=10, window=60)
limiter.allow("test:used", max_requests=10, window=60)
remaining, _ = limiter.get_remaining("test:used", max_requests=10, window=60)
assert remaining == 8
def test_get_remaining_zero(self, limiter):
for _ in range(5):
limiter.allow("test:zero", max_requests=5, window=60)
remaining, _ = limiter.get_remaining("test:zero", max_requests=5, window=60)
assert remaining == 0
class TestRateLimiterFlush:
"""Test flush() clears keys."""
@pytest.fixture
def limiter(self):
r = FakeRedis()
lim = RateLimiter(r)
return lim
def test_flush_clears_keys(self, limiter):
limiter.allow("flush:test:a", max_requests=10, window=60)
limiter.allow("flush:test:b", max_requests=10, window=60)
limiter.flush("flush:test:*")
# Keys should be cleared — fresh counts
remaining, _ = limiter.get_remaining("flush:test:a", max_requests=10, window=60)
assert remaining == 10
def test_flush_no_match(self, limiter):
limiter.allow("flush:other", max_requests=10, window=60)
limiter.flush("flush:nomatch:*")
# Should not raise on no matches
remaining, _ = limiter.get_remaining("flush:other", max_requests=10, window=60)
assert remaining == 9
if __name__ == "__main__":
pytest.main([__file__, "-v", "--tb=short"])