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