Files
truf-server/app/host_agent_server.py
2026-09-30 20:30:56 +03:00

168 lines
5.4 KiB
Python

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)