168 lines
5.4 KiB
Python
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)
|