Coverage for app/backend/src/couchers/middleware/ratelimit.py: 98%
87 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
1"""
2API rate limiting. See docs/rate-limit-design.md.
4Limits are a pair of (scope, dimension). Scopes nest: per-RPC ⊂ per-servicer ⊂ all-API; a single request
5increments a counter at every level. Dimensions are per-IP (keyed by subnet), per-user, and global. A
6request is rejected if any applicable limit is exceeded.
8This is independent of couchers.rate_limits, which is the per-user 24h action limiter.
9"""
11import logging
12import time
13from dataclasses import dataclass
14from functools import cache
15from ipaddress import ip_network
16from typing import TYPE_CHECKING
18import sentry_sdk
19import valkey
21from couchers.config import config
22from couchers.constants import RATE_LIMIT_WINDOW_SECONDS
23from couchers.experimentation import get_global_boolean_value
24from couchers.metrics import (
25 observe_rate_limit_check,
26 observe_rate_limit_duration,
27 observe_rate_limit_store_error,
28 observe_rate_limit_trip,
29)
30from couchers.middleware.proto_annotations import get_proto_annotations, optional_field, split_method
31from couchers.proto import annotations_pb2
33if TYPE_CHECKING:
34 from couchers.middleware.interceptors import CouchersHeaders, UserAuthInfo
36logger = logging.getLogger(__name__)
39@dataclass(frozen=True, slots=True)
40class ResolvedLimits:
41 """The fully-resolved per-minute limits for one method, per scope and dimension."""
43 service_name: str
44 rpc: dict[str, int]
45 svc: dict[str, int]
46 api: dict[str, int]
49@cache
50def resolve_method_rate_limits(method: str) -> ResolvedLimits:
51 """
52 Resolve the limits for a method from its proto annotations, falling back to global defaults.
54 per-RPC: method rate_limit.<dim> → service rate_limit_default.<dim> → global rpc default
55 per-servicer: service rate_limit_aggregate.<dim> → global svc default
56 all-API: global api default
57 """
58 dimensions = ("per_ip", "per_user", "global")
59 defaults = {
60 "rpc": {"per_ip": 60, "per_user": 120, "global": 6000},
61 "svc": {"per_ip": 300, "per_user": 600, "global": 20000},
62 "api": {"per_ip": 600, "per_user": 1200, "global": 60000},
63 }
65 def resolve(*values: int | None) -> int:
66 # the last value is always a global default, so there is always one to find
67 return next(value for value in values if value is not None)
69 annotations = get_proto_annotations()
70 service_name, _ = split_method(method)
71 method_rl = annotations.method_extension(method, annotations_pb2.rate_limit)
72 service_default = annotations.service_extension(service_name, annotations_pb2.rate_limit_default)
73 service_aggregate = annotations.service_extension(service_name, annotations_pb2.rate_limit_aggregate)
75 return ResolvedLimits(
76 service_name=service_name,
77 rpc={
78 dim: resolve(optional_field(method_rl, dim), optional_field(service_default, dim), defaults["rpc"][dim])
79 for dim in dimensions
80 },
81 svc={dim: resolve(optional_field(service_aggregate, dim), defaults["svc"][dim]) for dim in dimensions},
82 api=dict(defaults["api"]),
83 )
86def ip_to_key(ip: str, ipv6_prefix: int) -> str:
87 """
88 Mask an IP to its subnet and return a canonical string key.
90 IPv4 is keyed at /32 (the exact address); IPv6 is masked to ipv6_prefix bits (default /64).
91 """
92 network = ip_network(ip, strict=False)
93 prefix = 32 if network.version == 4 else ipv6_prefix
94 return str(ip_network(f"{network.network_address}/{prefix}", strict=False))
97_LUA_SCRIPT = """
98local tripped = {}
99local ttl = tonumber(ARGV[1])
100for i, key in ipairs(KEYS) do
101 local count = redis.call('INCR', key)
102 if count == 1 then
103 redis.call('EXPIRE', key, ttl)
104 end
105 if count > tonumber(ARGV[i + 1]) then
106 tripped[#tripped + 1] = i
107 end
108end
109return tripped
110"""
113class ValkeyCounterStore:
114 def __init__(self, host: str, port: int) -> None:
115 self._client = valkey.Valkey(
116 host=host,
117 port=port,
118 socket_connect_timeout=0.1,
119 socket_timeout=0.1,
120 )
121 self._script = self._client.register_script(_LUA_SCRIPT)
123 def incr_and_check(self, entries: list[tuple[str, int]], ttl: int) -> list[int]:
124 """Increment each key's counter; return the indices of entries whose count now exceeds its limit."""
125 keys = [key for key, _ in entries]
126 limits = [str(limit) for _, limit in entries]
127 result = self._script(keys=keys, args=[str(ttl), *limits])
128 # Lua returns 1-based indices.
129 return [i - 1 for i in result]
132@cache
133def _get_store() -> ValkeyCounterStore | None:
134 """The process-wide counter store, or None when no store is configured and rate limiting is off."""
135 if not config.VALKEY_HOST: 135 ↛ 137line 135 didn't jump to line 137 because the condition on line 135 was always true
136 return None
137 return ValkeyCounterStore(config.VALKEY_HOST, config.VALKEY_PORT)
140def _build_entries(
141 limits: ResolvedLimits, method: str, ip_address: str | None, user_id: int | None, bucket: int
142) -> list[tuple[str, str, str, int]]:
143 """The applicable (scope, dimension, key, limit) tuples for one request, across all scopes and dimensions."""
144 dim_ids: dict[str, str | None] = {
145 "per_ip": ip_to_key(ip_address, config.RATE_LIMIT_IPV6_PREFIX) if ip_address else None,
146 "per_user": str(user_id) if user_id is not None else None,
147 "global": "*",
148 }
149 scopes = (
150 ("rpc", method, limits.rpc),
151 ("svc", limits.service_name, limits.svc),
152 ("api", "*", limits.api),
153 )
154 entries = []
155 for scope, scope_id, scope_limits in scopes:
156 for dim, limit in scope_limits.items():
157 dim_id = dim_ids[dim]
158 if dim_id is None:
159 continue
160 key = f"rl:{scope}:{scope_id}:{dim}:{dim_id}:{bucket}"
161 entries.append((scope, dim, key, limit))
162 return entries
165def should_rate_limit(method: str, headers: CouchersHeaders, auth_info: UserAuthInfo | None) -> bool:
166 """
167 Count this request against every applicable limit and decide whether it should be rejected.
169 True only when a limit tripped and enforcement is on; shadow mode (the default) allows the request, as
170 does having no counter store configured at all, which turns rate limiting off entirely.
171 """
172 if auth_info and auth_info.is_superuser:
173 return False
175 store = _get_store()
176 if store is None:
177 return False
179 start = time.perf_counter_ns()
180 try:
181 entries = _build_entries(
182 resolve_method_rate_limits(method),
183 method,
184 headers.ip_address,
185 auth_info.user_id if auth_info else None,
186 int(time.time() // RATE_LIMIT_WINDOW_SECONDS),
187 )
189 try:
190 # keys carry their window bucket, so a stale key is only ever read within its own window; we let
191 # Valkey expire them a window later as housekeeping
192 tripped_idx = store.incr_and_check(
193 [(key, limit) for _, _, key, limit in entries], 2 * RATE_LIMIT_WINDOW_SECONDS
194 )
195 except Exception as e:
196 # nothing could be counted, so fail open: going dark on the counters shouldn't take the API down
197 observe_rate_limit_store_error(type(e).__name__)
198 observe_rate_limit_check(method, "failed_open")
199 sentry_sdk.set_tag("context", "rate_limiting")
200 sentry_sdk.capture_exception(e)
201 return False
203 if not tripped_idx:
204 observe_rate_limit_check(method, "allowed")
205 return False
207 enforced = get_global_boolean_value("rate_limiting_enabled", False)
208 for i in tripped_idx:
209 scope, dimension, _, _ = entries[i]
210 observe_rate_limit_trip(method, scope, dimension, enforced)
211 observe_rate_limit_check(method, "blocked" if enforced else "shadowed")
212 return enforced
213 finally:
214 observe_rate_limit_duration((time.perf_counter_ns() - start) / 1e9)