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

316 lines
10 KiB
Python

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()