Coverage for app/backend/src/couchers/middleware/interceptors.py: 87%

288 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-10-01 13:53 +0000

1import logging 

2from collections.abc import Callable, Mapping 

3from dataclasses import dataclass, field 

4from datetime import datetime 

5from os import getpid 

6from threading import get_ident 

7from time import perf_counter_ns 

8from traceback import format_exception 

9from typing import Any, NoReturn, cast 

10from zoneinfo import ZoneInfo 

11 

12import grpc 

13import sentry_sdk 

14from google.protobuf.message import Message 

15from opentelemetry import trace 

16from sqlalchemy import Function, literal_column, select, update 

17from sqlalchemy.dialects.postgresql import insert as pg_insert 

18from sqlalchemy.orm import undefer 

19from sqlalchemy.sql import func 

20 

21from couchers.config import config 

22from couchers.constants import ( 

23 CALL_CANCELLED_ERROR_MESSAGE, 

24 COOKIES_AND_AUTH_HEADER_ERROR_MESSAGE, 

25 NONEXISTENT_API_CALL_ERROR_MESSAGE, 

26 PERMISSION_DENIED_ERROR_MESSAGE, 

27 RATE_LIMIT_ERROR_MESSAGE, 

28 UNAUTHORIZED_ERROR_MESSAGE, 

29 UNKNOWN_ERROR_MESSAGE, 

30) 

31from couchers.context import CouchersContext, make_interactive_context, make_media_context 

32from couchers.db import session_scope 

33from couchers.i18n import LocalizationContext 

34from couchers.metrics import ( 

35 observe_api_call, 

36 observe_in_servicer_duration_histogram, 

37 observe_in_servicer_perf_histograms, 

38 observe_in_servicer_pool_wait_histogram, 

39 observe_in_servicer_serde_histogram, 

40 observe_in_servicer_setup_errors_counter, 

41 observe_in_servicer_setup_histogram, 

42) 

43from couchers.middleware.errors import CallRejectedError 

44from couchers.middleware.perf import PerfResult, read_perf, start_perf 

45from couchers.middleware.proto_annotations import get_proto_annotations 

46from couchers.middleware.ratelimit import should_rate_limit 

47from couchers.middleware.sanitize import sanitized_bytes 

48from couchers.models import APICall, ClientPlatform, User, UserActivity, UserSession 

49from couchers.proto import annotations_pb2 

50from couchers.proto.annotations_pb2 import AuthLevel 

51from couchers.utils import ( 

52 create_lang_cookie, 

53 create_session_cookies, 

54 generate_sofa_cookie, 

55 parse_api_key, 

56 parse_session_cookie, 

57 parse_sofa_cookie, 

58 parse_ui_lang_cookie, 

59 parse_user_id_cookie, 

60) 

61 

62logger = logging.getLogger(__name__) 

63 

64# the prometheus label shared by calls to methods with no servicer registered, whose name is whatever the caller sent 

65NONEXISTENT_METHOD_LABEL = "<nonexistent>" 

66 

67 

68@dataclass(frozen=True, slots=True, kw_only=True) 

69class UserAuthInfo: 

70 """Information about an authenticated user session.""" 

71 

72 user_id: int 

73 is_jailed: bool 

74 is_editor: bool 

75 is_superuser: bool 

76 token_expiry: datetime 

77 ui_language_preference: str | None 

78 timezone: str | None 

79 token: str = field(repr=False) 

80 is_api_key: bool 

81 

82 

83@dataclass(frozen=True, slots=True, kw_only=True) 

84class CouchersHeaders: 

85 # the user id cookie: client-supplied and unauthenticated, only good for spotting a desynced cookie 

86 user_id_str: str | None 

87 token: str | None = field(repr=False) 

88 sofa: str | None 

89 # which mechanism the token came in on, not whether it authenticated: a bad key still reads True here, while 

90 # the context's is_api_key is False whenever there's no session at all 

91 is_api_key: bool 

92 ip_address: str | None 

93 user_agent: str | None 

94 client_platform: ClientPlatform | None 

95 ui_lang: str | None 

96 

97 

98def _binned_now() -> Function[Any]: 

99 return func.date_bin( 

100 literal_column("interval '1 hour'"), 

101 func.now(), 

102 literal_column("'2000-01-01'::timestamptz"), 

103 ) 

104 

105 

106def _try_get_and_update_user_details( 

107 token: str | None, 

108 is_api_key: bool, 

109 ip_address: str | None, 

110 user_agent: str | None, 

111 sofa: str | None, 

112 client_platform: ClientPlatform | None, 

113) -> UserAuthInfo | None: 

114 """ 

115 Tries to get session and user info corresponding to this token. 

116 

117 Also updates the user's last active time, token last active time, and increments API call count. 

118 

119 Returns UserAuthInfo if a valid session is found, None otherwise. 

120 """ 

121 if not token: 

122 return None 

123 

124 with session_scope() as session: 

125 result = session.execute( 

126 select(User, UserSession, User.is_jailed) 

127 .select_from(UserSession) 

128 .join(User, User.id == UserSession.user_id) 

129 .where(User.is_visible) 

130 .where(UserSession.token == token) 

131 .where(UserSession.is_valid) 

132 .where(UserSession.is_api_key == is_api_key) 

133 # User.timezone is deferred and read below for every authenticated call, so load it here rather 

134 # than paying a second round trip for its ST_Contains against timezone_areas 

135 .options(undefer(User.timezone)) 

136 ).one_or_none() 

137 

138 if not result: 

139 return None 

140 

141 user, user_session, is_jailed = result._tuple() 

142 

143 # update user last active time if it's been a while; a non-matching UPDATE takes no row lock, so this 

144 # costs nothing on the calls that don't move it 

145 touch_user = ( 

146 update(User) 

147 .where(User.id == user.id) 

148 .where(User.last_active < func.now() - literal_column("interval '5 minutes'")) 

149 .values(last_active=func.now()) 

150 .cte("touch_user") 

151 ) 

152 

153 # let's update the token 

154 touch_session = ( 

155 update(UserSession) 

156 .where(UserSession.token == token) 

157 .values(last_seen=func.now(), api_calls=UserSession.api_calls + 1) 

158 .cte("touch_session") 

159 ) 

160 

161 # upsert so concurrent requests for the same activity tuple don't race to insert and violate the index 

162 insert_stmt = pg_insert(UserActivity).values( 

163 user_id=user.id, 

164 period=_binned_now(), 

165 ip_address=ip_address, 

166 user_agent=user_agent, 

167 sofa=sofa, 

168 client_platform=client_platform, 

169 api_calls=1, 

170 ) 

171 # one statement, so the sessions and user_activity row locks that every concurrent call from the same 

172 # session queues on are held for a single round trip. postgres leaves the order it applies the CTEs in 

173 # undefined, but every caller runs this same statement, so they all take those locks the same way round 

174 session.execute( 

175 insert_stmt.on_conflict_do_update( 

176 index_elements=[ 

177 UserActivity.user_id, 

178 UserActivity.period, 

179 UserActivity.ip_address, 

180 UserActivity.user_agent, 

181 UserActivity.sofa, 

182 ], 

183 set_={ 

184 "api_calls": UserActivity.api_calls + 1, 

185 "client_platform": func.coalesce( 

186 insert_stmt.excluded.client_platform, UserActivity.client_platform 

187 ), 

188 }, 

189 ).add_cte(touch_user, touch_session) 

190 ) 

191 

192 # build before committing to avoid expire_on_commit reloading these attributes 

193 auth_info = UserAuthInfo( 

194 user_id=user.id, 

195 is_jailed=is_jailed, 

196 is_editor=user.is_editor, 

197 is_superuser=user.is_superuser, 

198 token_expiry=user_session.expiry, 

199 ui_language_preference=user.ui_language_preference, 

200 timezone=user.timezone, 

201 token=token, 

202 is_api_key=is_api_key, 

203 ) 

204 

205 session.commit() 

206 

207 return auth_info 

208 

209 

210def abort_handler[T, R]( 

211 message: str, 

212 status_code: grpc.StatusCode, 

213) -> grpc.RpcMethodHandler[T, R]: 

214 def f(request: Any, context: CouchersContext) -> NoReturn: 

215 context.abort(status_code, message) 

216 

217 return grpc.unary_unary_rpc_method_handler(f) 

218 

219 

220def unauthenticated_handler[T, R]( 

221 message: str = UNAUTHORIZED_ERROR_MESSAGE, 

222 status_code: grpc.StatusCode = grpc.StatusCode.UNAUTHENTICATED, 

223) -> grpc.RpcMethodHandler[T, R]: 

224 return abort_handler(message, status_code) 

225 

226 

227def _log_call( 

228 *, 

229 method: str, 

230 status_code: str | None, 

231 user_id: int | None, 

232 is_api_key: bool, 

233 sofa: str | None, 

234 headers: CouchersHeaders | None, 

235 start: int, 

236 perf: PerfResult | None, 

237 request: Message | None = None, 

238 response: Message | None = None, 

239 exception: Exception | None = None, 

240 nonexistent_method: bool = False, 

241) -> None: 

242 """Record a finished call: one api_calls row, plus the per-call Prometheus observations.""" 

243 duration = (perf_counter_ns() - start) / 1e6 # ms 

244 metric_method = NONEXISTENT_METHOD_LABEL if nonexistent_method else method 

245 

246 req_bytes = sanitized_bytes(request) 

247 res_bytes = sanitized_bytes(response) 

248 response_truncated = False 

249 truncate_res_bytes_length = 16 * 1024 # 16 kB 

250 if res_bytes and len(res_bytes) > truncate_res_bytes_length: 250 ↛ 251line 250 didn't jump to line 251 because the condition on line 250 was never true

251 res_bytes = res_bytes[:truncate_res_bytes_length] 

252 response_truncated = True 

253 

254 traceback = "".join(format_exception(type(exception), exception, exception.__traceback__)) if exception else None 

255 

256 with session_scope() as session: 

257 session.add( 

258 APICall( 

259 is_api_key=is_api_key, 

260 method=method, 

261 status_code=status_code, 

262 duration=duration, 

263 user_id=user_id, 

264 request=req_bytes, 

265 response=res_bytes, 

266 response_truncated=response_truncated, 

267 traceback=traceback, 

268 db_query_count=perf.db_query_count if perf else None, 

269 db_write_query_count=perf.db_write_query_count if perf else None, 

270 db_time_ms=perf.db_time_ms if perf else None, 

271 cpu_ms=perf.cpu_ms if perf else None, 

272 client_platform=headers.client_platform if headers else None, 

273 ip_address=headers.ip_address if headers else None, 

274 user_agent=headers.user_agent if headers else None, 

275 sofa=sofa, 

276 ) 

277 ) 

278 

279 observe_in_servicer_duration_histogram( 

280 metric_method, user_id, status_code or "", type(exception).__name__ if exception else "", duration / 1000 

281 ) 

282 observe_api_call(metric_method, headers.client_platform if headers else None) 

283 logger.debug(f"{user_id=}, {method=}, {duration=} ms") 

284 

285 

286def _log_rejected_call( 

287 *, 

288 method: str, 

289 code: grpc.StatusCode, 

290 start: int, 

291 handler_call_details: grpc.HandlerCallDetails, 

292 user_id: int | None = None, 

293 exception: Exception | None = None, 

294 nonexistent_method: bool = False, 

295) -> None: 

296 """Log a call rejected during auth/setup.""" 

297 try: 

298 headers: CouchersHeaders | None = parse_headers(dict(handler_call_details.invocation_metadata)) 

299 except BadHeaders: 

300 headers = None 

301 

302 perf = read_perf() 

303 _log_call( 

304 method=method, 

305 status_code=code.name, 

306 user_id=user_id, 

307 is_api_key=headers.is_api_key if headers else False, 

308 sofa=headers.sofa if headers else None, 

309 headers=headers, 

310 start=start, 

311 perf=perf, 

312 exception=exception, 

313 nonexistent_method=nonexistent_method, 

314 ) 

315 observe_in_servicer_setup_histogram(NONEXISTENT_METHOD_LABEL if nonexistent_method else method, perf) 

316 

317 

318def _rejected_call_handler[T, R]( 

319 *, 

320 method: str, 

321 message: str, 

322 code: grpc.StatusCode, 

323 start: int, 

324 handler_call_details: grpc.HandlerCallDetails, 

325) -> grpc.RpcMethodHandler[T, R]: 

326 """Terminate a call that has no handler to run, logging it from the pool thread rather than the serving one.""" 

327 

328 def f(request: Any, context: grpc.ServicerContext) -> NoReturn: 

329 start_perf() 

330 _log_rejected_call( 

331 method=method, 

332 code=code, 

333 start=start, 

334 handler_call_details=handler_call_details, 

335 nonexistent_method=True, 

336 ) 

337 context.abort(code, message) 

338 

339 return grpc.unary_unary_rpc_method_handler(f) 

340 

341 

342type Cont[T, R] = Callable[[grpc.HandlerCallDetails], grpc.RpcMethodHandler[T, R] | None] 

343 

344 

345@dataclass(frozen=True, slots=True, kw_only=True) 

346class AdmittedCall: 

347 """What a call that's cleared to run carries into the handler body.""" 

348 

349 headers: CouchersHeaders 

350 auth_info: UserAuthInfo | None 

351 sofa: str 

352 new_sofa_cookie: str | None 

353 localization: LocalizationContext 

354 

355 

356@dataclass(frozen=True, slots=True, kw_only=True) 

357class RejectedCall: 

358 """What a call that didn't clear setup gets terminated and logged with.""" 

359 

360 code: grpc.StatusCode 

361 message: str 

362 user_id: int | None 

363 # set when setup broke rather than turned the call away, so the caller can report it 

364 exception: Exception | None 

365 

366 

367def admit_call(handler_call_details: grpc.HandlerCallDetails) -> AdmittedCall | RejectedCall: 

368 """ 

369 Pre-RPC setup handling. 

370 

371 Never raises: a call that doesn't make it through comes back as a RejectedCall carrying whatever setup had 

372 resolved before it stopped, so the caller can log the call it never ran. 

373 """ 

374 auth_info = None 

375 try: 

376 headers = parse_headers(dict(handler_call_details.invocation_metadata)) 

377 

378 # if this is not present in prod, it's a Big Bug in config 

379 assert config.DEV or headers.ip_address is not None 

380 

381 auth_level = get_proto_annotations().auth_level(handler_call_details.method) 

382 

383 auth_info = _try_get_and_update_user_details( 

384 headers.token, 

385 headers.is_api_key, 

386 headers.ip_address, 

387 headers.user_agent, 

388 headers.sofa, 

389 headers.client_platform, 

390 ) 

391 

392 check_permissions(auth_info, auth_level) 

393 

394 if should_rate_limit(handler_call_details.method, headers, auth_info): 

395 raise CallRejectedError(RATE_LIMIT_ERROR_MESSAGE, grpc.StatusCode.RESOURCE_EXHAUSTED) 

396 

397 if headers.sofa: 

398 sofa = headers.sofa 

399 new_sofa_cookie = None 

400 else: 

401 sofa, new_sofa_cookie = generate_sofa_cookie() 

402 

403 loc_context = LocalizationContext( 

404 locale=(auth_info.ui_language_preference if auth_info else headers.ui_lang) or "", 

405 timezone=ZoneInfo((auth_info and auth_info.timezone) or "Etc/UTC"), 

406 ) 

407 except BadHeaders: 

408 return RejectedCall( 

409 code=grpc.StatusCode.UNAUTHENTICATED, 

410 message=COOKIES_AND_AUTH_HEADER_ERROR_MESSAGE, 

411 user_id=None, 

412 exception=None, 

413 ) 

414 except CallRejectedError as e: 

415 return RejectedCall( 

416 code=e.code, message=e.msg, user_id=auth_info.user_id if auth_info else None, exception=None 

417 ) 

418 except Exception as e: 

419 return RejectedCall( 

420 code=grpc.StatusCode.INTERNAL, 

421 message=UNKNOWN_ERROR_MESSAGE, 

422 user_id=auth_info.user_id if auth_info else None, 

423 exception=e, 

424 ) 

425 

426 return AdmittedCall( 

427 headers=headers, 

428 auth_info=auth_info, 

429 sofa=sofa, 

430 new_sofa_cookie=new_sofa_cookie, 

431 localization=loc_context, 

432 ) 

433 

434 

435class CouchersMiddlewareInterceptor(grpc.ServerInterceptor): 

436 """ 

437 1. Does auth: extracts a session token from a cookie, and authenticates a user with that. 

438 

439 Sets context.user_id and context.token if authenticated, otherwise 

440 terminates the call with an UNAUTHENTICATED error code. 

441 

442 2. Makes sure cookies are in sync. 

443 

444 3. Injects a session to get a database transaction. 

445 

446 4. Measures and logs the time it takes to service each incoming call. 

447 

448 All of that happens in the returned handler, on a thread pool thread. gRPC runs intercept_service inline on the 

449 server's single completion-queue thread while holding the server-wide lock, and only submits to the pool once 

450 the interceptor chain has returned a handler, so blocking in the interceptor body (the auth query above all) 

451 serializes call dispatch for the entire worker process. 

452 """ 

453 

454 def __init__(self) -> None: 

455 # builds the descriptor pool at startup rather than on the first call 

456 get_proto_annotations() 

457 

458 def intercept_service[T = Message, R = Message]( 

459 self, 

460 continuation: Cont[T, R], 

461 handler_call_details: grpc.HandlerCallDetails, 

462 ) -> grpc.RpcMethodHandler[T, R]: 

463 start = perf_counter_ns() 

464 

465 method = handler_call_details.method 

466 

467 # only the handler lookup happens here, the rest waits for the handler thread; see the class docstring 

468 handler = continuation(handler_call_details) 

469 if not handler or not (prev_function := handler.unary_unary): 

470 return _rejected_call_handler( 

471 method=method, 

472 message=NONEXISTENT_API_CALL_ERROR_MESSAGE, 

473 code=grpc.StatusCode.UNIMPLEMENTED, 

474 start=start, 

475 handler_call_details=handler_call_details, 

476 ) 

477 

478 def function_without_couchers_stuff(req: Message, grpc_context: grpc.ServicerContext) -> Message | None: 

479 # accounting for the auth/setup phase; the handler re-arms its own below 

480 start_perf() 

481 

482 call = admit_call(handler_call_details) 

483 

484 if isinstance(call, RejectedCall): 

485 # anything unexpected goes to Sentry before the row below: if the DB is what's broken, that fails too 

486 if call.exception: 

487 observe_in_servicer_setup_errors_counter(method, type(call.exception).__name__) 

488 sentry_sdk.set_tag("context", "servicer_setup") 

489 sentry_sdk.set_tag("method", method) 

490 sentry_sdk.capture_exception(call.exception) 

491 _log_rejected_call( 

492 method=method, 

493 code=call.code, 

494 start=start, 

495 handler_call_details=handler_call_details, 

496 user_id=call.user_id, 

497 exception=call.exception, 

498 ) 

499 grpc_context.abort(call.code, call.message) 

500 

501 observe_in_servicer_setup_histogram(method, read_perf()) 

502 

503 headers = call.headers 

504 auth_info = call.auth_info 

505 sofa = call.sofa 

506 loc_context = call.localization 

507 

508 couchers_context = make_interactive_context( 

509 grpc_context=grpc_context, 

510 user_id=auth_info.user_id if auth_info else None, 

511 is_api_key=auth_info.is_api_key if auth_info else False, 

512 token=auth_info.token if auth_info else None, 

513 localization=loc_context, 

514 sofa=sofa, 

515 ) 

516 

517 with session_scope() as session: 

518 # force the checkout now so its wait is timed here rather than hiding in the handler's first query 

519 pool_wait_start = perf_counter_ns() 

520 session.connection() 

521 observe_in_servicer_pool_wait_histogram(method, (perf_counter_ns() - pool_wait_start) / 1e9) 

522 start_perf() 

523 

524 res: Message | None = None 

525 exception: Exception | None = None 

526 try: 

527 _res = prev_function(req, couchers_context, session) # type: ignore[call-arg, arg-type] 

528 # flush so pending ORM writes execute (and are counted) before we snapshot; a handler that only 

529 # session.add(...)s and returns would otherwise flush at commit, after read_perf() 

530 session.flush() 

531 res = cast(Message, _res) 

532 except Exception as e: 

533 exception = e 

534 

535 perf = read_perf() 

536 

537 if exception and couchers_context._grpc_context: 

538 context_code = couchers_context._grpc_context.code() # type: ignore[attr-defined] 

539 code = getattr(context_code, "name", None) 

540 else: 

541 code = None 

542 

543 _log_call( 

544 method=method, 

545 status_code=code, 

546 user_id=couchers_context._user_id, 

547 is_api_key=cast(bool, couchers_context._is_api_key), 

548 sofa=sofa, 

549 headers=headers, 

550 start=start, 

551 perf=perf, 

552 request=req, 

553 response=res, 

554 exception=exception, 

555 ) 

556 observe_in_servicer_perf_histograms(method, perf) 

557 

558 if exception: 

559 if not code: 

560 sentry_sdk.set_tag("context", "servicer") 

561 sentry_sdk.set_tag("method", method) 

562 sentry_sdk.set_tag("user_agent", headers.user_agent) 

563 sentry_sdk.set_tag("ui_lang", loc_context.preferred_locale) 

564 sentry_sdk.set_user( 

565 { 

566 "id": couchers_context._user_id, 

567 "ip_address": headers.ip_address, 

568 "sofa": sofa[:12], 

569 } 

570 ) 

571 sentry_sdk.capture_exception(exception) 

572 

573 raise exception 

574 

575 if auth_info and not auth_info.is_api_key: 

576 # check the two cookies are in sync & that language preference cookie is correct 

577 if headers.user_id_str != str(auth_info.user_id): 577 ↛ 581line 577 didn't jump to line 581 because the condition on line 577 was always true

578 couchers_context.set_cookies( 

579 create_session_cookies(auth_info.token, auth_info.user_id, auth_info.token_expiry) 

580 ) 

581 if auth_info.ui_language_preference and auth_info.ui_language_preference != headers.ui_lang: 

582 couchers_context.set_cookies(create_lang_cookie(auth_info.ui_language_preference)) 

583 

584 if call.new_sofa_cookie: 

585 couchers_context.set_cookies([call.new_sofa_cookie]) 

586 

587 if not grpc_context.is_active(): 587 ↛ 588line 587 didn't jump to line 588 because the condition on line 587 was never true

588 grpc_context.abort(grpc.StatusCode.INTERNAL, CALL_CANCELLED_ERROR_MESSAGE) 

589 

590 couchers_context._send_cookies() 

591 

592 return res 

593 

594 def timed_serde[A, B](fn: Callable[[A], B], direction: str) -> Callable[[A], B]: 

595 def wrapped(arg: A) -> B: 

596 t0 = perf_counter_ns() 

597 result = fn(arg) 

598 observe_in_servicer_serde_histogram(method, direction, (perf_counter_ns() - t0) / 1e9) 

599 return result 

600 

601 return wrapped 

602 

603 # always set for our generated-proto methods, but grpc types them as optional 

604 assert handler.request_deserializer is not None and handler.response_serializer is not None 

605 return grpc.unary_unary_rpc_method_handler( 

606 function_without_couchers_stuff, 

607 request_deserializer=timed_serde(handler.request_deserializer, "deserialize"), 

608 response_serializer=timed_serde(handler.response_serializer, "serialize"), 

609 ) 

610 

611 

612def parse_headers(headers: Mapping[str, str | bytes]) -> CouchersHeaders: 

613 if "cookie" in headers and "authorization" in headers: 

614 # for security reasons, only one of "cookie" or "authorization" can be present 

615 raise BadHeaders("Both cookies and authorization are present in headers") 

616 elif "cookie" in headers: 

617 # the session token is passed in cookies, i.e., in the `cookie` header 

618 token, is_api_key = parse_session_cookie(headers), False 

619 elif "authorization" in headers: 

620 # the session token is passed in the `authorization` header 

621 token, is_api_key = parse_api_key(headers), True 

622 else: 

623 # no session found 

624 token, is_api_key = None, False 

625 

626 ip_address = headers.get("x-couchers-real-ip") 

627 user_agent = headers.get("user-agent") 

628 

629 # the client (web app or native app) declares its platform via this header 

630 client_platform_raw = headers.get("x-couchers-client-platform") 

631 client_platform = ( 

632 ClientPlatform[client_platform_raw] 

633 if isinstance(client_platform_raw, str) and client_platform_raw in ClientPlatform.__members__ 

634 else None 

635 ) 

636 

637 ui_lang = parse_ui_lang_cookie(headers) 

638 user_id_str = parse_user_id_cookie(headers) 

639 sofa = parse_sofa_cookie(headers) 

640 

641 return CouchersHeaders( 

642 user_id_str=user_id_str, 

643 token=token, 

644 sofa=sofa, 

645 is_api_key=is_api_key, 

646 ip_address=ip_address if isinstance(ip_address, str) else None, 

647 user_agent=user_agent if isinstance(user_agent, str) else None, 

648 client_platform=client_platform, 

649 ui_lang=ui_lang, 

650 ) 

651 

652 

653class BadHeaders(Exception): 

654 pass 

655 

656 

657def check_permissions(auth_info: UserAuthInfo | None, auth_level: AuthLevel.ValueType) -> None: 

658 if not auth_info: 

659 # if this isn't an open service, fail 

660 if auth_level != annotations_pb2.AUTH_LEVEL_OPEN: 

661 raise CallRejectedError(UNAUTHORIZED_ERROR_MESSAGE, grpc.StatusCode.UNAUTHENTICATED) 

662 else: 

663 # a valid user session was found - check permissions 

664 if auth_level == annotations_pb2.AUTH_LEVEL_ADMIN and not auth_info.is_superuser: 

665 raise CallRejectedError(PERMISSION_DENIED_ERROR_MESSAGE, grpc.StatusCode.PERMISSION_DENIED) 

666 

667 if auth_level == annotations_pb2.AUTH_LEVEL_EDITOR and not auth_info.is_editor: 

668 raise CallRejectedError(PERMISSION_DENIED_ERROR_MESSAGE, grpc.StatusCode.PERMISSION_DENIED) 

669 

670 # if the user is jailed and this isn't an open or jailed service, fail 

671 if auth_info.is_jailed and auth_level not in [ 

672 annotations_pb2.AUTH_LEVEL_OPEN, 

673 annotations_pb2.AUTH_LEVEL_JAILED, 

674 ]: 

675 raise CallRejectedError(PERMISSION_DENIED_ERROR_MESSAGE, grpc.StatusCode.UNAUTHENTICATED) 

676 

677 

678class MediaInterceptor(grpc.ServerInterceptor): 

679 """ 

680 Extracts an "Authorization: Bearer <hex>" header and calls the 

681 is_authorized function. Terminates the call with an HTTP error 

682 code if not authorized. 

683 

684 Also adds a session to called APIs. 

685 """ 

686 

687 def __init__(self, is_authorized: Callable[[str], bool]): 

688 self._is_authorized = is_authorized 

689 

690 def intercept_service[T, R]( 

691 self, 

692 continuation: Cont[T, R], 

693 handler_call_details: grpc.HandlerCallDetails, 

694 ) -> grpc.RpcMethodHandler[T, R]: 

695 handler = continuation(handler_call_details) 

696 if not handler: 696 ↛ 697line 696 didn't jump to line 697 because the condition on line 696 was never true

697 raise RuntimeError("No handler") 

698 

699 prev_func = handler.unary_unary 

700 if not prev_func: 700 ↛ 701line 700 didn't jump to line 701 because the condition on line 700 was never true

701 raise RuntimeError(f"No prev_function, {handler}") 

702 

703 metadata = dict(handler_call_details.invocation_metadata) 

704 

705 token = parse_api_key(metadata) 

706 

707 if not token or not self._is_authorized(token): 707 ↛ 708line 707 didn't jump to line 708 because the condition on line 707 was never true

708 return unauthenticated_handler() 

709 

710 def function_without_session(request: T, grpc_context: grpc.ServicerContext) -> R: 

711 with session_scope() as session: 

712 return prev_func(request, make_media_context(grpc_context), session) # type: ignore[call-arg, arg-type] 

713 

714 return grpc.unary_unary_rpc_method_handler( 

715 function_without_session, 

716 request_deserializer=handler.request_deserializer, 

717 response_serializer=handler.response_serializer, 

718 ) 

719 

720 

721class OTelInterceptor(grpc.ServerInterceptor): 

722 """ 

723 OpenTelemetry tracing 

724 """ 

725 

726 def __init__(self) -> None: 

727 self.tracer = trace.get_tracer(__name__) 

728 

729 def intercept_service[T, R]( 

730 self, 

731 continuation: Cont[T, R], 

732 handler_call_details: grpc.HandlerCallDetails, 

733 ) -> grpc.RpcMethodHandler[T, R]: 

734 handler = continuation(handler_call_details) 

735 if not handler: 

736 raise RuntimeError("No handler") 

737 

738 prev_func = handler.unary_unary 

739 if not prev_func: 

740 raise RuntimeError(f"No prev_function, {handler}") 

741 

742 method = handler_call_details.method 

743 

744 def tracing_function(request: T, context: grpc.ServicerContext) -> R: 

745 # method is of the form "/org.couchers.api.core.API/GetUser" 

746 _, service_name, method_name = method.split("/") 

747 

748 headers = dict(handler_call_details.invocation_metadata) 

749 

750 with self.tracer.start_as_current_span("handler") as rollspan: 

751 rollspan.set_attribute("rpc.method_full", method) 

752 rollspan.set_attribute("rpc.service", service_name) 

753 rollspan.set_attribute("rpc.method", method_name) 

754 

755 rollspan.set_attribute("rpc.thread", get_ident()) 

756 rollspan.set_attribute("rpc.pid", getpid()) 

757 

758 res = prev_func(request, context) 

759 

760 rollspan.set_attribute("web.user_agent", headers.get("user-agent") or "") 

761 rollspan.set_attribute("web.ip_address", headers.get("x-couchers-real-ip") or "") 

762 

763 return res 

764 

765 return grpc.unary_unary_rpc_method_handler( 

766 tracing_function, 

767 request_deserializer=handler.request_deserializer, 

768 response_serializer=handler.response_serializer, 

769 ) 

770 

771 

772class ErrorSanitizationInterceptor(grpc.ServerInterceptor): 

773 """ 

774 If the call resulted in a non-gRPC error, this strips away the error details. 

775 

776 It's important to put this first, so that it does not interfere with other interceptors. 

777 """ 

778 

779 def intercept_service[T, R]( 

780 self, 

781 continuation: Cont[T, R], 

782 handler_call_details: grpc.HandlerCallDetails, 

783 ) -> grpc.RpcMethodHandler[T, R]: 

784 handler = continuation(handler_call_details) 

785 if not handler: 785 ↛ 786line 785 didn't jump to line 786 because the condition on line 785 was never true

786 raise RuntimeError("No handler") 

787 

788 prev_func = handler.unary_unary 

789 if not prev_func: 789 ↛ 790line 789 didn't jump to line 790 because the condition on line 789 was never true

790 raise RuntimeError(f"No prev_function, {handler}") 

791 

792 def sanitizing_function(req: T, context: grpc.ServicerContext) -> R: 

793 try: 

794 res = prev_func(req, context) 

795 except Exception as e: 

796 code = context.code() # type: ignore[attr-defined] 

797 # the code is one of the RPC error codes if this was failed through abort(), otherwise it's None 

798 if not code: 

799 logger.exception(e) 

800 logger.info("Probably an unknown error! Sanitizing...") 

801 context.abort(grpc.StatusCode.INTERNAL, UNKNOWN_ERROR_MESSAGE) 

802 else: 

803 logger.warning(f"RPC error: {code} in method {handler_call_details.method}") 

804 raise e 

805 return res 

806 

807 return grpc.unary_unary_rpc_method_handler( 

808 sanitizing_function, 

809 request_deserializer=handler.request_deserializer, 

810 response_serializer=handler.response_serializer, 

811 )