Coverage for app/backend/src/couchers/middleware/proto_annotations.py: 96%
47 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"""
2Reading our custom proto annotations (see proto/annotations.proto) off the descriptor pool.
4Everything that knows how descriptors, service/method options and extensions fit together lives here, so
5callers work in terms of "the auth level for this method" rather than in terms of protobuf machinery.
7The pool is an invariant of ProtoAnnotations rather than an argument threaded through every lookup, and
8get_proto_annotations() is the process-wide instance. Lookups memoize on the instance, which is bounded
9because the descriptors fix the set of methods; failures aren't memoized, so a request naming a method that
10doesn't exist raises every time rather than accumulating entries.
11"""
13from functools import cache
14from typing import Any, cast
16import grpc
17from google.protobuf.descriptor import MethodDescriptor, ServiceDescriptor
18from google.protobuf.descriptor_pool import DescriptorPool
19from google.protobuf.message import Message
21from couchers.constants import (
22 MISSING_AUTH_LEVEL_ERROR_MESSAGE,
23 NONEXISTENT_API_CALL_ERROR_MESSAGE,
24)
25from couchers.middleware.descriptor_pool import build_descriptor_pool
26from couchers.middleware.errors import CallRejectedError
27from couchers.proto import annotations_pb2
28from couchers.proto.annotations_pb2 import AuthLevel
31def split_method(method: str) -> tuple[str, str]:
32 """Split a gRPC method path, e.g. "/org.couchers.api.core.API/GetUser", into service and method names."""
33 _, service_name, method_name = method.split("/")
34 return service_name, method_name
37def optional_field(message: Message, field: str) -> int | None:
38 """Read an optional scalar field, honouring proto field presence, so an explicit 0 differs from unset."""
39 return getattr(message, field) if message.HasField(field) else None
42def validate_auth_level(auth_level: AuthLevel.ValueType) -> None:
43 # if unknown auth level, then it wasn't set and something's wrong
44 if auth_level == annotations_pb2.AUTH_LEVEL_UNKNOWN:
45 raise CallRejectedError(MISSING_AUTH_LEVEL_ERROR_MESSAGE, grpc.StatusCode.INTERNAL)
47 if auth_level not in { 47 ↛ 54line 47 didn't jump to line 54 because the condition on line 47 was never true
48 annotations_pb2.AUTH_LEVEL_OPEN,
49 annotations_pb2.AUTH_LEVEL_JAILED,
50 annotations_pb2.AUTH_LEVEL_SECURE,
51 annotations_pb2.AUTH_LEVEL_EDITOR,
52 annotations_pb2.AUTH_LEVEL_ADMIN,
53 }:
54 raise CallRejectedError(MISSING_AUTH_LEVEL_ERROR_MESSAGE, grpc.StatusCode.INTERNAL)
57class ProtoAnnotations:
58 """The annotations on our API, read off one descriptor pool."""
60 def __init__(self, pool: DescriptorPool) -> None:
61 self._pool = pool
62 self._auth_levels: dict[str, AuthLevel.ValueType] = {}
64 def _find_service(self, service_name: str) -> ServiceDescriptor:
65 try:
66 return cast(ServiceDescriptor, self._pool.FindServiceByName(service_name)) # type: ignore[no-untyped-call]
67 except KeyError:
68 raise CallRejectedError(NONEXISTENT_API_CALL_ERROR_MESSAGE, grpc.StatusCode.UNIMPLEMENTED) from None
70 def _find_method(self, method: str) -> MethodDescriptor:
71 service_name, method_name = split_method(method)
72 return cast(MethodDescriptor, self._find_service(service_name).FindMethodByName(method_name)) # type: ignore[no-untyped-call]
74 def service_extension(self, service_name: str, extension: Any) -> Message:
75 """The value of a service-level extension; protobuf returns the default instance when it isn't set."""
76 return cast(Message, self._find_service(service_name).GetOptions().Extensions[extension])
78 def method_extension(self, method: str, extension: Any) -> Message:
79 """The value of a method-level extension; protobuf returns the default instance when it isn't set."""
80 return cast(Message, self._find_method(method).GetOptions().Extensions[extension])
82 def auth_level(self, method: str) -> AuthLevel.ValueType:
83 if method not in self._auth_levels:
84 service_name, _ = split_method(method)
85 level = self._find_service(service_name).GetOptions().Extensions[annotations_pb2.auth_level]
86 validate_auth_level(level)
87 self._auth_levels[method] = level
88 return self._auth_levels[method]
91@cache
92def get_proto_annotations() -> ProtoAnnotations:
93 """The process-wide annotations, built off the descriptor set shipped with the backend."""
94 return ProtoAnnotations(build_descriptor_pool())