316 lines
10 KiB
Python
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()
|