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

1""" 

2Reading our custom proto annotations (see proto/annotations.proto) off the descriptor pool. 

3 

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. 

6 

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

12 

13from functools import cache 

14from typing import Any, cast 

15 

16import grpc 

17from google.protobuf.descriptor import MethodDescriptor, ServiceDescriptor 

18from google.protobuf.descriptor_pool import DescriptorPool 

19from google.protobuf.message import Message 

20 

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 

29 

30 

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 

35 

36 

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 

40 

41 

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) 

46 

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) 

55 

56 

57class ProtoAnnotations: 

58 """The annotations on our API, read off one descriptor pool.""" 

59 

60 def __init__(self, pool: DescriptorPool) -> None: 

61 self._pool = pool 

62 self._auth_levels: dict[str, AuthLevel.ValueType] = {} 

63 

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 

69 

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] 

73 

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

77 

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

81 

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] 

89 

90 

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