Coverage for app/backend/src/tests/test_ratelimit.py: 99%
178 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:53 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 13:53 +0000
1import os
2from uuid import uuid4
4import grpc
5import pytest
6import valkey
8from couchers import metrics
9from couchers.middleware import ratelimit
10from couchers.middleware.interceptors import CouchersHeaders
11from couchers.proto import api_pb2, auth_pb2
12from tests.fixtures.db import generate_user
13from tests.fixtures.sessions import auth_api_session, real_api_session
15AUTHENTICATE = "/org.couchers.auth.Auth/Authenticate"
16USERNAME_VALID = "/org.couchers.auth.Auth/UsernameValid"
18# Where the Valkey integration tests look for a server; docker-compose.test.yml publishes one on 6545.
19VALKEY_TEST_HOST = os.environ.get("VALKEY_TEST_HOST", "localhost")
20VALKEY_TEST_PORT = int(os.environ.get("VALKEY_TEST_PORT", "6545"))
23class InMemoryCounterStore:
24 """A pure-Python fixed-window store mirroring the Valkey one, for hermetic tests."""
26 def __init__(self) -> None:
27 self.counts: dict[str, int] = {}
29 def incr_and_check(self, entries: list[tuple[str, int]], ttl: int) -> list[int]:
30 tripped = []
31 for i, (key, limit) in enumerate(entries):
32 self.counts[key] = self.counts.get(key, 0) + 1
33 if self.counts[key] > limit:
34 tripped.append(i)
35 return tripped
38class AlwaysTripStore:
39 def incr_and_check(self, entries: list[tuple[str, int]], ttl: int) -> list[int]:
40 return list(range(len(entries)))
43class BrokenStore:
44 def incr_and_check(self, entries: list[tuple[str, int]], ttl: int) -> list[int]:
45 raise RuntimeError("valkey down")
48def _use_store(monkeypatch, store):
49 """Point the rate limiter at this counter store (None meaning none configured)."""
50 monkeypatch.setattr(ratelimit, "_get_store", lambda: store)
53@pytest.fixture
54def store(monkeypatch):
55 """Inject an in-memory counter store, bypassing Valkey."""
56 s = InMemoryCounterStore()
57 _use_store(monkeypatch, s)
58 return s
61@pytest.fixture
62def valkey_client():
63 """A raw client against the test Valkey, skipping the test if there isn't one running."""
64 client = valkey.Valkey(host=VALKEY_TEST_HOST, port=VALKEY_TEST_PORT, socket_connect_timeout=1, socket_timeout=1)
65 try:
66 client.ping()
67 except valkey.ConnectionError as e:
68 pytest.skip(
69 f"no Valkey at {VALKEY_TEST_HOST}:{VALKEY_TEST_PORT} ({e}); "
70 f"start one with `docker compose -f docker-compose.test.yml up -d valkey_tests`"
71 )
72 return client
75@pytest.fixture
76def valkey_store(valkey_client):
77 """The real Valkey-backed store, so the Lua script itself is exercised rather than a stand-in."""
78 return ratelimit.ValkeyCounterStore(VALKEY_TEST_HOST, VALKEY_TEST_PORT)
81@pytest.fixture
82def key_prefix():
83 """A prefix unique to this test run, so counters can't collide with a previous run's leftovers."""
84 return f"test:{uuid4().hex}"
87def _limited(method: str, ip: str | None = None) -> bool:
88 """should_rate_limit for an unauthenticated call from this IP; the limiter reads no other header."""
89 headers = CouchersHeaders(
90 token=None,
91 is_api_key=False,
92 ip_address=ip,
93 user_agent=None,
94 client_platform=None,
95 ui_lang=None,
96 user_id_str=None,
97 sofa=None,
98 )
99 return ratelimit.should_rate_limit(method, headers, None)
102def test_ip_to_key_ipv4():
103 assert ratelimit.ip_to_key("1.2.3.4", 64) == "1.2.3.4/32"
106def test_ip_to_key_ipv6_masks_to_prefix():
107 assert ratelimit.ip_to_key("2001:db8::1", 64) == "2001:db8::/64"
108 assert ratelimit.ip_to_key("2001:0db8:0000:0000:dead:beef:0:1", 64) == "2001:db8::/64"
109 assert ratelimit.ip_to_key("2001:db8::1", 64) == ratelimit.ip_to_key("2001:db8::ffff", 64)
112def test_ip_to_key_ipv6_prefix_configurable():
113 assert ratelimit.ip_to_key("2001:db8:abcd:1234::1", 48) == "2001:db8:abcd::/48"
116def test_resolve_method_rate_limits_method_override():
117 limits = ratelimit.resolve_method_rate_limits(AUTHENTICATE)
118 assert limits.rpc == {"per_ip": 10, "per_user": 120, "global": 6000}
121def test_resolve_method_rate_limits_defaults():
122 limits = ratelimit.resolve_method_rate_limits(USERNAME_VALID)
123 assert limits.rpc == {"per_ip": 60, "per_user": 120, "global": 6000}
124 assert limits.svc == {"per_ip": 300, "per_user": 600, "global": 20000}
125 assert limits.api == {"per_ip": 600, "per_user": 1200, "global": 60000}
128def test_no_store_means_rate_limiting_is_off():
129 # this is the one test that goes through the real accessor, so it can't reuse another test's store
130 ratelimit._get_store.cache_clear()
131 assert ratelimit._get_store() is None
132 assert not _limited(AUTHENTICATE, ip="1.2.3.4")
135def test_trips_per_ip(feature_flags, store):
136 feature_flags.set("rate_limiting_enabled", True)
137 # Authenticate per_ip = 10 and every other limit is far higher, so the 11th call from an IP is the first
138 # to trip anything
139 for _ in range(10):
140 assert not _limited(AUTHENTICATE, ip="1.2.3.4")
141 assert _limited(AUTHENTICATE, ip="1.2.3.4")
144def test_per_ip_skipped_without_ip(feature_flags, store):
145 feature_flags.set("rate_limiting_enabled", True)
146 # no IP → the per_ip dimension is not counted, so the per_ip=10 limit can never trip
147 for _ in range(20):
148 assert not _limited(AUTHENTICATE)
151def test_separate_subnets_counted_separately(feature_flags, store):
152 feature_flags.set("rate_limiting_enabled", True)
153 for _ in range(11):
154 _limited(AUTHENTICATE, ip="2001:db8:1::1")
155 assert not _limited(AUTHENTICATE, ip="2001:db8:2::1")
158def test_store_error_fails_open(feature_flags, monkeypatch):
159 feature_flags.set("rate_limiting_enabled", True)
160 _use_store(monkeypatch, BrokenStore())
161 captured = []
162 monkeypatch.setattr("couchers.middleware.ratelimit.sentry_sdk.capture_exception", lambda e: captured.append(e))
163 monkeypatch.setattr("couchers.middleware.ratelimit.sentry_sdk.set_tag", lambda *a, **k: None)
165 assert not _limited(AUTHENTICATE, ip="1.2.3.4")
166 assert len(captured) == 1
169def test_interceptor_superuser_exempt_when_enforcing(db, feature_flags, monkeypatch):
170 feature_flags.set("rate_limiting_enabled", True)
171 _use_store(monkeypatch, AlwaysTripStore())
172 superuser, token = generate_user(is_superuser=True)
174 # real_api_session, not api_session: only the real server runs the interceptor the limiter lives in
175 with real_api_session(token) as api:
176 assert api.Ping(api_pb2.PingReq()).user.user_id == superuser.id
179def test_interceptor_non_superuser_still_blocked_when_enforcing(db, feature_flags, monkeypatch):
180 feature_flags.set("rate_limiting_enabled", True)
181 _use_store(monkeypatch, AlwaysTripStore())
182 _, token = generate_user()
184 with real_api_session(token) as api:
185 with pytest.raises(grpc.RpcError) as e:
186 api.Ping(api_pb2.PingReq())
187 assert e.value.code() == grpc.StatusCode.RESOURCE_EXHAUSTED
190def test_interceptor_no_store_allows(db, feature_flags, monkeypatch):
191 feature_flags.set("rate_limiting_enabled", True)
192 _use_store(monkeypatch, None)
193 with auth_api_session() as (auth_api, _):
194 assert auth_api.UsernameValid(auth_pb2.UsernameValidReq(username="test")).valid
197def test_interceptor_shadow_allows(db, feature_flags, monkeypatch):
198 feature_flags.set("rate_limiting_enabled", False)
199 _use_store(monkeypatch, AlwaysTripStore())
200 with auth_api_session() as (auth_api, _):
201 assert auth_api.UsernameValid(auth_pb2.UsernameValidReq(username="test")).valid
204def test_interceptor_enforce_rejects(db, feature_flags, monkeypatch):
205 feature_flags.set("rate_limiting_enabled", True)
206 _use_store(monkeypatch, AlwaysTripStore())
207 with auth_api_session() as (auth_api, _):
208 with pytest.raises(grpc.RpcError) as e:
209 auth_api.UsernameValid(auth_pb2.UsernameValidReq(username="test"))
210 assert e.value.code() == grpc.StatusCode.RESOURCE_EXHAUSTED
213def test_interceptor_fails_open_when_enforcing(db, feature_flags, monkeypatch):
214 feature_flags.set("rate_limiting_enabled", True)
215 _use_store(monkeypatch, BrokenStore())
216 with auth_api_session() as (auth_api, _):
217 assert auth_api.UsernameValid(auth_pb2.UsernameValidReq(username="test")).valid
220def _metric_value(counter, name: str, **labels: str) -> float:
221 return float(
222 sum(
223 s.value
224 for m in counter.collect()
225 for s in m.samples
226 if s.name == name and all(s.labels.get(k) == v for k, v in labels.items())
227 )
228 )
231def test_interceptor_emits_metrics_on_enforce(db, feature_flags, monkeypatch):
232 feature_flags.set("rate_limiting_enabled", True)
233 _use_store(monkeypatch, AlwaysTripStore())
235 blocked_before = _metric_value(
236 metrics.rate_limit_checks_counter,
237 "couchers_rate_limit_checks_total",
238 method=USERNAME_VALID,
239 decision="blocked",
240 )
241 # no IP/user on this call, so the global dimension trips at every scope
242 trip_before = _metric_value(
243 metrics.rate_limit_trips_counter,
244 "couchers_rate_limit_trips_total",
245 method=USERNAME_VALID,
246 scope="rpc",
247 dimension="global",
248 enforced="true",
249 )
251 with auth_api_session() as (auth_api, _):
252 with pytest.raises(grpc.RpcError):
253 auth_api.UsernameValid(auth_pb2.UsernameValidReq(username="test"))
255 assert (
256 _metric_value(
257 metrics.rate_limit_checks_counter,
258 "couchers_rate_limit_checks_total",
259 method=USERNAME_VALID,
260 decision="blocked",
261 )
262 == blocked_before + 1
263 )
264 assert (
265 _metric_value(
266 metrics.rate_limit_trips_counter,
267 "couchers_rate_limit_trips_total",
268 method=USERNAME_VALID,
269 scope="rpc",
270 dimension="global",
271 enforced="true",
272 )
273 == trip_before + 1
274 )
277# The tests below run the real Lua script against a real Valkey; everything above uses a stand-in store.
280def test_valkey_store_counts_and_trips(valkey_store, key_prefix):
281 key = f"{key_prefix}:counted"
282 for _ in range(3):
283 assert valkey_store.incr_and_check([(key, 3)], 120) == []
284 assert valkey_store.incr_and_check([(key, 3)], 120) == [0]
285 assert valkey_store.incr_and_check([(key, 3)], 120) == [0]
288def test_valkey_store_returns_indices_of_tripped_entries(valkey_store, key_prefix):
289 # a limit of 0 trips on the first increment, a high limit never does; this pins the Lua script's
290 # 1-based indices being translated back to the 0-based positions of the entries passed in
291 entries = [
292 (f"{key_prefix}:high:0", 100),
293 (f"{key_prefix}:zero:1", 0),
294 (f"{key_prefix}:high:2", 100),
295 (f"{key_prefix}:zero:3", 0),
296 ]
297 assert valkey_store.incr_and_check(entries, 120) == [1, 3]
300def test_valkey_store_counts_keys_independently(valkey_store, key_prefix):
301 a, b = f"{key_prefix}:a", f"{key_prefix}:b"
302 for _ in range(5):
303 valkey_store.incr_and_check([(a, 5)], 120)
304 assert valkey_store.incr_and_check([(a, 5), (b, 5)], 120) == [0]
307def test_valkey_store_sets_ttl_on_first_increment(valkey_store, valkey_client, key_prefix):
308 key = f"{key_prefix}:ttl"
309 valkey_store.incr_and_check([(key, 100)], 120)
310 # without a TTL the counter would never reset and the key would leak
311 assert 0 < valkey_client.ttl(key) <= 120
314def test_valkey_store_does_not_extend_ttl_on_later_increments(valkey_store, valkey_client, key_prefix):
315 key = f"{key_prefix}:ttl-once"
316 valkey_store.incr_and_check([(key, 100)], 120)
317 # pull the expiry in, then increment again: the window must not slide, or a sustained flood would keep
318 # renewing its own counter and the fixed window would never roll over
319 valkey_client.expire(key, 5)
320 valkey_store.incr_and_check([(key, 100)], 120)
321 assert 0 < valkey_client.ttl(key) <= 5
324def test_rate_limiting_end_to_end_against_valkey(feature_flags, valkey_store, monkeypatch):
325 feature_flags.set("rate_limiting_enabled", True)
326 _use_store(monkeypatch, valkey_store)
327 # a /64 unique to this run, so the per-IP counters start clean
328 ip = f"2001:db8:{uuid4().hex[:4]}:{uuid4().hex[:4]}::1"
330 # Authenticate annotates per_ip = 10
331 for _ in range(10):
332 assert not _limited(AUTHENTICATE, ip=ip)
333 assert _limited(AUTHENTICATE, ip=ip)
335 other_ip = f"2001:db8:{uuid4().hex[:4]}:{uuid4().hex[:4]}::1"
336 assert not _limited(AUTHENTICATE, ip=other_ip)