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

1""" 

2API rate limiting. See docs/rate-limit-design.md. 

3 

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. 

7 

8This is independent of couchers.rate_limits, which is the per-user 24h action limiter. 

9""" 

10 

11import logging 

12import time 

13from dataclasses import dataclass 

14from functools import cache 

15from ipaddress import ip_network 

16from typing import TYPE_CHECKING 

17 

18import sentry_sdk 

19import valkey 

20 

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 

32 

33if TYPE_CHECKING: 

34 from couchers.middleware.interceptors import CouchersHeaders, UserAuthInfo 

35 

36logger = logging.getLogger(__name__) 

37 

38 

39@dataclass(frozen=True, slots=True) 

40class ResolvedLimits: 

41 """The fully-resolved per-minute limits for one method, per scope and dimension.""" 

42 

43 service_name: str 

44 rpc: dict[str, int] 

45 svc: dict[str, int] 

46 api: dict[str, int] 

47 

48 

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. 

53 

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 } 

64 

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) 

68 

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) 

74 

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 ) 

84 

85 

86def ip_to_key(ip: str, ipv6_prefix: int) -> str: 

87 """ 

88 Mask an IP to its subnet and return a canonical string key. 

89 

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)) 

95 

96 

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""" 

111 

112 

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) 

122 

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] 

130 

131 

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) 

138 

139 

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 

163 

164 

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. 

168 

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 

174 

175 store = _get_store() 

176 if store is None: 

177 return False 

178 

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 ) 

188 

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 

202 

203 if not tripped_idx: 

204 observe_rate_limit_check(method, "allowed") 

205 return False 

206 

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)