import hmac import json import re import struct import time import uuid from dataclasses import dataclass from enum import Enum HOST_AGENT_SOCKET_PATH = '/run/truf/host-agent.sock' HOST_AGENT_RUNTIME_UID = 10001 MAX_REQUEST_PAYLOAD_BYTES = 1024 MAX_RESPONSE_PAYLOAD_BYTES = 256 CLIENT_CONNECT_TIMEOUT_SECONDS = 1.0 SERVER_READ_TIMEOUT_SECONDS = 2.0 EXCHANGE_TIMEOUT_SECONDS = 5.0 _FRAME_HEADER_BYTES = 4 _SHA256_RE = re.compile(r'^[0-9a-f]{64}$') _REQUEST_FIELDS = frozenset(( 'operation_id', 'action', 'active_config_sha256', 'active_secrets_sha256', 'candidate_config_sha256', 'candidate_secrets_sha256', )) _RESPONSE_FIELDS = frozenset(('operation_id', 'status')) class HostAgentProtocolError(ValueError): def __init__(self, category): self.category = category super().__init__('host agent protocol message is invalid') class HostAgentAction(str, Enum): APPLY_CONFIG = 'apply-config' APPLY_SECRETS = 'apply-secrets' APPLY_BOTH = 'apply-both' RESTART = 'restart' class HostAgentStatus(str, Enum): ACCEPTED = 'accepted' UNAVAILABLE = 'unavailable' REJECTED = 'rejected' INVALID = 'invalid' @dataclass(frozen=True, slots=True) class HostAgentRequest: operation_id: str action: HostAgentAction active_config_sha256: str active_secrets_sha256: str candidate_config_sha256: str | None candidate_secrets_sha256: str | None @dataclass(frozen=True, slots=True) class HostAgentResponse: operation_id: str | None status: HostAgentStatus def _canonical_uuid(value): if not isinstance(value, str) or not value: raise HostAgentProtocolError('operation_id') try: parsed = uuid.UUID(value) except (ValueError, AttributeError) as exc: raise HostAgentProtocolError('operation_id') from exc if parsed.int == 0 or str(parsed) != value: raise HostAgentProtocolError('operation_id') return value def _sha256(value, field, *, optional=False): if optional and value is None: return None if not isinstance(value, str) or _SHA256_RE.fullmatch(value) is None: raise HostAgentProtocolError(field) return value def _canonical_json(value): try: return json.dumps( value, ensure_ascii=True, sort_keys=True, separators=(',', ':'), allow_nan=False, ).encode('ascii') except (TypeError, ValueError, UnicodeError) as exc: raise HostAgentProtocolError('json') from exc def _strict_json(payload, *, maximum): if type(payload) is not bytes or not 1 <= len(payload) <= maximum: raise HostAgentProtocolError('bounds') def reject_duplicate(pairs): result = {} for key, value in pairs: if key in result: raise HostAgentProtocolError('duplicate_field') result[key] = value return result try: text = payload.decode('utf-8', errors='strict') value = json.loads( text, object_pairs_hook=reject_duplicate, parse_constant=lambda _value: (_ for _ in ()).throw( HostAgentProtocolError('constant') ), ) except HostAgentProtocolError: raise except (UnicodeDecodeError, json.JSONDecodeError, TypeError, ValueError) as exc: raise HostAgentProtocolError('json') from exc finally: text = None if not isinstance(value, dict): raise HostAgentProtocolError('shape') if not hmac.compare_digest(_canonical_json(value), payload): raise HostAgentProtocolError('canonical') return value def _normalize_request(value): if not isinstance(value, dict) or set(value) != _REQUEST_FIELDS: raise HostAgentProtocolError('shape') try: action = HostAgentAction(value.get('action')) except (TypeError, ValueError) as exc: raise HostAgentProtocolError('action') from exc active_config = _sha256(value.get('active_config_sha256'), 'active_config_sha256') active_secrets = _sha256(value.get('active_secrets_sha256'), 'active_secrets_sha256') candidate_config = _sha256( value.get('candidate_config_sha256'), 'candidate_config_sha256', optional=True, ) candidate_secrets = _sha256( value.get('candidate_secrets_sha256'), 'candidate_secrets_sha256', optional=True, ) required = { HostAgentAction.APPLY_CONFIG: (True, False), HostAgentAction.APPLY_SECRETS: (False, True), HostAgentAction.APPLY_BOTH: (True, True), HostAgentAction.RESTART: (False, False), }[action] if (candidate_config is not None, candidate_secrets is not None) != required: raise HostAgentProtocolError('candidate_identity') return HostAgentRequest( operation_id=_canonical_uuid(value.get('operation_id')), action=action, active_config_sha256=active_config, active_secrets_sha256=active_secrets, candidate_config_sha256=candidate_config, candidate_secrets_sha256=candidate_secrets, ) def _request_value(request): if not isinstance(request, HostAgentRequest): raise HostAgentProtocolError('request_type') return { 'operation_id': request.operation_id, 'action': request.action.value if isinstance(request.action, HostAgentAction) else request.action, 'active_config_sha256': request.active_config_sha256, 'active_secrets_sha256': request.active_secrets_sha256, 'candidate_config_sha256': request.candidate_config_sha256, 'candidate_secrets_sha256': request.candidate_secrets_sha256, } def encode_request_payload(request): normalized = _normalize_request(_request_value(request)) payload = _canonical_json(_request_value(normalized)) if len(payload) > MAX_REQUEST_PAYLOAD_BYTES: raise HostAgentProtocolError('bounds') return payload def decode_request_payload(payload): return _normalize_request(_strict_json(payload, maximum=MAX_REQUEST_PAYLOAD_BYTES)) def _normalize_response(value): if not isinstance(value, dict) or set(value) != _RESPONSE_FIELDS: raise HostAgentProtocolError('shape') try: status = HostAgentStatus(value.get('status')) except (TypeError, ValueError) as exc: raise HostAgentProtocolError('status') from exc operation_id = value.get('operation_id') if status is HostAgentStatus.INVALID: if operation_id is not None: raise HostAgentProtocolError('operation_id') else: operation_id = _canonical_uuid(operation_id) return HostAgentResponse(operation_id=operation_id, status=status) def _response_value(response): if not isinstance(response, HostAgentResponse): raise HostAgentProtocolError('response_type') return { 'operation_id': response.operation_id, 'status': response.status.value if isinstance(response.status, HostAgentStatus) else response.status, } def encode_response_payload(response): normalized = _normalize_response(_response_value(response)) payload = _canonical_json(_response_value(normalized)) if len(payload) > MAX_RESPONSE_PAYLOAD_BYTES: raise HostAgentProtocolError('bounds') return payload def decode_response_payload(payload): return _normalize_response(_strict_json(payload, maximum=MAX_RESPONSE_PAYLOAD_BYTES)) def _encode_frame(payload, maximum): if type(payload) is not bytes or not 1 <= len(payload) <= maximum: raise HostAgentProtocolError('bounds') return struct.pack('!I', len(payload)) + payload def _decode_frame(frame, maximum): if type(frame) is not bytes or len(frame) < _FRAME_HEADER_BYTES: raise HostAgentProtocolError('frame') length = struct.unpack('!I', frame[:_FRAME_HEADER_BYTES])[0] if not 1 <= length <= maximum or len(frame) != _FRAME_HEADER_BYTES + length: raise HostAgentProtocolError('frame') return frame[_FRAME_HEADER_BYTES:] def encode_request_frame(request): return _encode_frame(encode_request_payload(request), MAX_REQUEST_PAYLOAD_BYTES) def decode_request_frame(frame): return decode_request_payload(_decode_frame(frame, MAX_REQUEST_PAYLOAD_BYTES)) def encode_response_frame(response): return _encode_frame(encode_response_payload(response), MAX_RESPONSE_PAYLOAD_BYTES) def decode_response_frame(frame): return decode_response_payload(_decode_frame(frame, MAX_RESPONSE_PAYLOAD_BYTES)) def _remaining(deadline): remaining = deadline - time.monotonic() if remaining <= 0: raise HostAgentProtocolError('timeout') return remaining def receive_frame(sock, *, maximum, deadline): header = _receive_exact(sock, _FRAME_HEADER_BYTES, deadline) length = struct.unpack('!I', header)[0] if not 1 <= length <= maximum: raise HostAgentProtocolError('bounds') return header + _receive_exact(sock, length, deadline) def _receive_exact(sock, length, deadline): chunks = bytearray() try: while len(chunks) < length: sock.settimeout(_remaining(deadline)) chunk = sock.recv(length - len(chunks)) if not chunk: raise HostAgentProtocolError('truncated') chunks.extend(chunk) return bytes(chunks) except HostAgentProtocolError: raise except (OSError, TimeoutError) as exc: raise HostAgentProtocolError('transport') from exc finally: chunks.clear() chunk = None def require_eof(sock, *, deadline): try: sock.settimeout(_remaining(deadline)) if sock.recv(1): raise HostAgentProtocolError('trailing_data') except HostAgentProtocolError: raise except (OSError, TimeoutError) as exc: raise HostAgentProtocolError('transport') from exc def send_frame(sock, frame, *, deadline): if type(frame) is not bytes: raise HostAgentProtocolError('frame') view = memoryview(frame) try: while view: sock.settimeout(_remaining(deadline)) sent = sock.send(view) if sent <= 0: raise HostAgentProtocolError('transport') view = view[sent:] except HostAgentProtocolError: raise except (OSError, TimeoutError) as exc: raise HostAgentProtocolError('transport') from exc finally: view.release()