import os import socket import stat import struct import time from host_agent_protocol import ( EXCHANGE_TIMEOUT_SECONDS, HOST_AGENT_RUNTIME_UID, HOST_AGENT_SOCKET_PATH, MAX_REQUEST_PAYLOAD_BYTES, SERVER_READ_TIMEOUT_SECONDS, HostAgentProtocolError, HostAgentResponse, HostAgentStatus, decode_request_frame, encode_response_frame, receive_frame, require_eof, send_frame, ) SYSTEMD_LISTEN_FD = 3 ACCEPT_POLL_SECONDS = 1.0 class HostAgentServerError(RuntimeError): def __init__(self, category): self.category = category super().__init__('host operations agent server failed') def _peer_credentials(connection): if not hasattr(socket, 'SO_PEERCRED'): raise HostAgentServerError('peer_credentials_unavailable') try: raw = connection.getsockopt( socket.SOL_SOCKET, socket.SO_PEERCRED, struct.calcsize('3i'), ) pid, uid, gid = struct.unpack('3i', raw) except (OSError, struct.error) as exc: raise HostAgentServerError('peer_credentials_unavailable') from exc return pid, uid, gid def unavailable_handler(_request): return HostAgentStatus.UNAVAILABLE def serve_connection(connection, *, handler=unavailable_handler, accepted_at=None): if not isinstance(connection, socket.socket): raise HostAgentServerError('connection_invalid') started = time.monotonic() if accepted_at is None else accepted_at try: peer_pid, peer_uid, _peer_gid = _peer_credentials(connection) except HostAgentServerError: return False if peer_pid <= 0 or peer_uid != HOST_AGENT_RUNTIME_UID: return False request = None response = None try: read_deadline = min( started + SERVER_READ_TIMEOUT_SECONDS, started + EXCHANGE_TIMEOUT_SECONDS, ) frame = receive_frame( connection, maximum=MAX_REQUEST_PAYLOAD_BYTES, deadline=read_deadline, ) require_eof(connection, deadline=read_deadline) request = decode_request_frame(frame) except HostAgentProtocolError: response = HostAgentResponse(None, HostAgentStatus.INVALID) else: try: outcome = handler(request) if isinstance(outcome, HostAgentResponse): response = outcome else: response = HostAgentResponse( request.operation_id, HostAgentStatus(outcome), ) if response.operation_id != request.operation_id: raise HostAgentServerError('handler_identity_invalid') except BaseException as exc: if not isinstance(exc, Exception): raise response = HostAgentResponse( request.operation_id, HostAgentStatus.UNAVAILABLE, ) try: send_frame( connection, encode_response_frame(response), deadline=started + EXCHANGE_TIMEOUT_SECONDS, ) except HostAgentProtocolError: return False finally: frame = request = response = outcome = None return True def _validate_listener(listener): if not isinstance(listener, socket.socket): raise HostAgentServerError('listener_invalid') unix_family = getattr(socket, 'AF_UNIX', None) if unix_family is None: raise HostAgentServerError('listener_invalid') try: socket_type = listener.getsockopt(socket.SOL_SOCKET, socket.SO_TYPE) accepting = listener.getsockopt(socket.SOL_SOCKET, socket.SO_ACCEPTCONN) except OSError as exc: raise HostAgentServerError('listener_invalid') from exc if ( listener.family != unix_family or socket_type != socket.SOCK_STREAM or accepting != 1 ): raise HostAgentServerError('listener_invalid') try: if listener.getsockname() != HOST_AGENT_SOCKET_PATH: raise HostAgentServerError('listener_path_invalid') details = os.lstat(HOST_AGENT_SOCKET_PATH) except HostAgentServerError: raise except OSError as exc: raise HostAgentServerError('listener_unavailable') from exc if not stat.S_ISSOCK(details.st_mode) or details.st_uid != 0: raise HostAgentServerError('listener_owner_invalid') def inherited_systemd_listener(): geteuid = getattr(os, 'geteuid', None) if geteuid is None or geteuid() != 0: raise HostAgentServerError('root_required') if os.environ.get('LISTEN_PID') != str(os.getpid()): raise HostAgentServerError('socket_activation_invalid') if os.environ.get('LISTEN_FDS') != '1': raise HostAgentServerError('socket_activation_invalid') try: listener = socket.socket(fileno=SYSTEMD_LISTEN_FD) _validate_listener(listener) except BaseException: try: listener.close() except (OSError, UnboundLocalError): pass raise return listener def serve_forever(listener, *, handler=unavailable_handler, stop_event=None): _validate_listener(listener) if stop_event is not None: listener.settimeout(ACCEPT_POLL_SECONDS) while stop_event is None or not stop_event.is_set(): try: connection, _address = listener.accept() except InterruptedError: continue except TimeoutError: continue with connection: serve_connection(connection, handler=handler)