Initial server source import
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user