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
« 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>
4"""
5RPC server for DAK
6"""
8import logging
9import os
10import sys
11from concurrent import futures
13import apt_pkg
14import grpc
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
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
40def unix_socket_path(listen_addr: str) -> str | None:
41 """Return the filesystem path of a `unix:` listen address.
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("//")
52def parse_socket_mode(mode: str) -> int:
53 """Parse an octal file mode string like "0660".
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
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.
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))
85 server.add_insecure_port(listen_addr)
86 if socket_chmod is not None:
87 os.chmod(*socket_chmod)
90def main() -> None:
91 cnf = Config()
93 apt_pkg.parse_commandline( # type: ignore[attr-defined]
94 cnf.Cnf,
95 [
96 ("o", "option", "", "ArbItem"),
97 ],
98 sys.argv,
99 )
101 daklog.Logger("rpc-server")
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)
112 if "RPC::Authorization::TokenFile" not in cnf:
113 raise Exception("RPC::Authorization::TokenFile is required.")
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 ]
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()