Coverage for dak/rpc_server.py: 66%

61 statements  

« prev     ^ index     » next       coverage.py v7.6.0, created at 2026-08-03 16:46 +0000

1# SPDX-License-Identifier: GPL-2.0-or-later 

2# © 2026, Ansgar 🙀 <ansgar@debian.org> 

3 

4""" 

5RPC server for DAK 

6""" 

7 

8import logging 

9import os 

10import sys 

11from concurrent import futures 

12 

13import apt_pkg 

14import grpc 

15 

16import daklib.policy_rpc 

17from dak.policyqueue.v1 import policyqueue_pb2_grpc 

18from daklib import daklog 

19from daklib.config import Config 

20from daklib.daklog import DakLogHandler 

21from daklib.dbconn import DBConn 

22from daklib.rpc_auth import AuthenticationInterceptor, TokenAuth, load_tokens_from_file 

23from daklib.rpc_log import LoggingInterceptor, RequestContextFilter 

24from daklib.rpc_peer import PeerAddressInterceptor 

25 

26 

27def init_server( 

28 *, interceptors: list[grpc.ServerInterceptor] | None = None 

29) -> grpc.Server: 

30 server = grpc.server( 

31 thread_pool=futures.ThreadPoolExecutor(max_workers=10), 

32 interceptors=interceptors or [], 

33 ) 

34 policyqueue_pb2_grpc.add_PolicyQueueServiceServicer_to_server( 

35 daklib.policy_rpc.PolicyQueueServiceServicer(conn=DBConn()), server 

36 ) 

37 return server 

38 

39 

40def unix_socket_path(listen_addr: str) -> str | None: 

41 """Return the filesystem path of a `unix:` listen address. 

42 

43 Returns None for TCP addresses and for `unix-abstract:` sockets, 

44 which have no filesystem path. 

45 """ 

46 if not listen_addr.startswith("unix:"): 

47 return None 

48 # handle both gRPC forms: `unix:path` and `unix://absolute_path` 

49 return listen_addr.removeprefix("unix:").removeprefix("//") 

50 

51 

52def parse_socket_mode(mode: str) -> int: 

53 """Parse an octal file mode string like "0660". 

54 

55 Modes above 0o777 (setuid/setgid/sticky) are rejected; execute bits 

56 are always cleared, as they have no meaning for a socket. 

57 """ 

58 try: 

59 parsed = int(mode, 8) 

60 except ValueError: 

61 raise ValueError(f"invalid octal file mode: {mode!r}") from None 

62 if not 0 <= parsed <= 0o777: 

63 raise ValueError(f"file mode out of range: {mode!r}") 

64 return parsed & 0o666 

65 

66 

67def setup_listen_socket( 

68 server: grpc.Server, listen_addr: str, mode: str | None = None 

69) -> None: 

70 """Bind `server` to `listen_addr` and apply socket permissions. 

71 

72 A `mode` (octal file mode string) requires a filesystem `unix:` 

73 address; it is validated before binding and applied after. 

74 """ 

75 socket_chmod: tuple[str, int] | None = None 

76 if mode: 

77 socket_path = unix_socket_path(listen_addr) 

78 if socket_path is None: 

79 raise Exception( 

80 "RPC::ListenAddressMode requires RPC::ListenAddress to be a" 

81 " filesystem unix: socket." 

82 ) 

83 socket_chmod = (socket_path, parse_socket_mode(mode)) 

84 

85 server.add_insecure_port(listen_addr) 

86 if socket_chmod is not None: 

87 os.chmod(*socket_chmod) 

88 

89 

90def main() -> None: 

91 cnf = Config() 

92 

93 apt_pkg.parse_commandline( # type: ignore[attr-defined] 

94 cnf.Cnf, 

95 [ 

96 ("o", "option", "", "ArbItem"), 

97 ], 

98 sys.argv, 

99 ) 

100 

101 daklog.Logger("rpc-server") 

102 

103 context_filter = RequestContextFilter() 

104 daklog_formatter = logging.Formatter( 

105 "[%(request_id)s] [%(peer)s] [%(auth_sub)s] %(name)s: %(message)s" 

106 ) 

107 for h in logging.getLogger().handlers: 

108 h.addFilter(context_filter) 

109 if isinstance(h, DakLogHandler): 

110 h.setFormatter(daklog_formatter) 

111 

112 if "RPC::Authorization::TokenFile" not in cnf: 

113 raise Exception("RPC::Authorization::TokenFile is required.") 

114 

115 interceptors: list[grpc.ServerInterceptor] = [ 

116 # PeerAddressInterceptor must come first: the header-derived peer 

117 # has to be set before LoggingInterceptor logs "request started", 

118 # and its transport-address wrapper must run outside the logging 

119 # wrapper so "request finished" sees the resolved peer. 

120 PeerAddressInterceptor(peer_header=cnf.get("RPC::PeerAddressHeader")), 

121 LoggingInterceptor(), 

122 AuthenticationInterceptor( 

123 TokenAuth(load_tokens_from_file(cnf["RPC::Authorization::TokenFile"])) 

124 ), 

125 ] 

126 

127 server = init_server(interceptors=interceptors) 

128 listen_addr = cnf.get("RPC::ListenAddress") 

129 if not listen_addr: 

130 raise Exception("RPC::ListenAddress is required.") 

131 setup_listen_socket(server, listen_addr, cnf.get("RPC::ListenAddressMode")) 

132 server.start() 

133 server.wait_for_termination()