Files
truf-server/tests/test_host_agent_protocol.py
T
2026-09-30 20:30:56 +03:00

530 lines
22 KiB
Python

import json
from pathlib import Path
import socket
import struct
import sys
import threading
import time
import unittest
import uuid
from unittest import mock
ROOT = Path(__file__).resolve().parents[1]
APP = ROOT / 'app'
if str(APP) not in sys.path:
sys.path.insert(0, str(APP))
import host_agent_client
import host_agent_protocol
import host_agent_server
from host_agent_client import (
HostAgentClient,
HostAgentRejectedError,
HostAgentUnavailableError,
)
from host_agent_protocol import (
MAX_REQUEST_PAYLOAD_BYTES,
HostAgentAction,
HostAgentProtocolError,
HostAgentRequest,
HostAgentResponse,
HostAgentStatus,
decode_request_frame,
decode_request_payload,
decode_response_frame,
encode_request_frame,
encode_request_payload,
encode_response_frame,
)
def request_for(action='apply-config'):
candidates = {
'apply-config': ('c' * 64, None),
'apply-secrets': (None, 'd' * 64),
'apply-both': ('c' * 64, 'd' * 64),
'restart': (None, None),
}[action]
return HostAgentRequest(
operation_id=str(uuid.uuid4()),
action=HostAgentAction(action),
active_config_sha256='a' * 64,
active_secrets_sha256='b' * 64,
candidate_config_sha256=candidates[0],
candidate_secrets_sha256=candidates[1],
)
class HostAgentProtocolTests(unittest.TestCase):
def test_all_actions_round_trip_as_canonical_exact_schema(self):
for action in ('apply-config', 'apply-secrets', 'apply-both', 'restart'):
with self.subTest(action=action):
request = request_for(action)
payload = encode_request_payload(request)
self.assertEqual(decode_request_payload(payload), request)
self.assertEqual(decode_request_frame(encode_request_frame(request)), request)
parsed = json.loads(payload)
self.assertEqual(set(parsed), {
'operation_id', 'action', 'active_config_sha256',
'active_secrets_sha256', 'candidate_config_sha256',
'candidate_secrets_sha256',
})
self.assertEqual(
payload,
json.dumps(
parsed, ensure_ascii=True, sort_keys=True,
separators=(',', ':'), allow_nan=False,
).encode('ascii'),
)
def test_unknown_duplicate_missing_and_noncanonical_fields_are_rejected(self):
request = request_for()
canonical = encode_request_payload(request)
parsed = json.loads(canonical)
malformed = []
for field in tuple(parsed):
changed = dict(parsed)
changed.pop(field)
malformed.append(json.dumps(changed, sort_keys=True, separators=(',', ':')).encode())
for field in ('path', 'service', 'command', 'argv', 'environment', 'docker', 'options'):
changed = dict(parsed, **{field: 'forbidden'})
malformed.append(json.dumps(changed, sort_keys=True, separators=(',', ':')).encode())
malformed.extend((
canonical.replace(b'"action":', b' "action":', 1),
canonical + b'\n',
b'{' + canonical[1:canonical.find(b',') + 1] + canonical[1:],
b'[]',
b'{"action":NaN}',
b'\xff',
))
for payload in malformed:
with self.subTest(payload=payload[:60]):
with self.assertRaises(HostAgentProtocolError):
decode_request_payload(payload)
def test_identity_types_hashes_uuid_and_action_matrix_are_exact(self):
parsed = json.loads(encode_request_payload(request_for('apply-both')))
cases = (
('operation_id', str(uuid.UUID(int=0))),
('operation_id', parsed['operation_id'].upper()),
('operation_id', 1),
('action', 'shell'),
('active_config_sha256', 'A' * 64),
('active_secrets_sha256', 'b' * 63),
('candidate_config_sha256', None),
('candidate_secrets_sha256', True),
)
for field, value in cases:
with self.subTest(field=field, value=value):
changed = dict(parsed)
changed[field] = value
payload = json.dumps(changed, sort_keys=True, separators=(',', ':')).encode()
with self.assertRaises(HostAgentProtocolError):
decode_request_payload(payload)
def test_frames_reject_zero_oversize_truncation_and_trailing_data(self):
frame = encode_request_frame(request_for())
for changed in (
b'', b'\x00\x00\x00', struct.pack('!I', 0),
struct.pack('!I', MAX_REQUEST_PAYLOAD_BYTES + 1),
frame[:-1], frame + b'x',
):
with self.subTest(frame=changed[:20]):
with self.assertRaises(HostAgentProtocolError):
decode_request_frame(changed)
def test_response_schema_is_closed_and_identity_bound(self):
operation_id = str(uuid.uuid4())
for status in (
HostAgentStatus.ACCEPTED, HostAgentStatus.UNAVAILABLE,
HostAgentStatus.REJECTED,
):
response = HostAgentResponse(operation_id, status)
self.assertEqual(decode_response_frame(encode_response_frame(response)), response)
invalid = HostAgentResponse(None, HostAgentStatus.INVALID)
self.assertEqual(decode_response_frame(encode_response_frame(invalid)), invalid)
for response in (
HostAgentResponse(None, HostAgentStatus.ACCEPTED),
HostAgentResponse(operation_id, HostAgentStatus.INVALID),
):
with self.assertRaises(HostAgentProtocolError):
encode_response_frame(response)
def test_absolute_deadline_is_not_extended_by_trickled_bytes(self):
class TrickleSocket:
def __init__(self):
self.timeouts = []
def settimeout(self, value):
self.timeouts.append(value)
def recv(self, _length):
return b'\x00'
connection = TrickleSocket()
with mock.patch.object(
host_agent_protocol.time, 'monotonic',
side_effect=(1.0, 4.0, 7.0, 10.1),
):
with self.assertRaisesRegex(HostAgentProtocolError, 'invalid') as raised:
host_agent_protocol.receive_frame(
connection, maximum=MAX_REQUEST_PAYLOAD_BYTES, deadline=10.0,
)
self.assertEqual(raised.exception.category, 'timeout')
self.assertEqual(len(connection.timeouts), 3)
class _ClientSocket:
def __init__(self, response, *, peer_uid=0):
self.response = bytearray(response)
self.peer_uid = peer_uid
self.sent = bytearray()
self.closed = False
self.shutdown_how = None
self.connected = None
def settimeout(self, _value):
pass
def connect(self, path):
self.connected = path
def getsockopt(self, _level, _option, _length):
return struct.pack('3i', 123, self.peer_uid, 0)
def send(self, value):
self.sent.extend(value)
return len(value)
def recv(self, length):
if not self.response:
return b''
result = bytes(self.response[:length])
del self.response[:length]
return result
def shutdown(self, how):
self.shutdown_how = how
def close(self):
self.closed = True
class HostAgentClientTests(unittest.TestCase):
def test_client_sends_exact_request_and_accepts_only_root_response(self):
request = request_for('apply-both')
response = encode_response_frame(HostAgentResponse(
request.operation_id, HostAgentStatus.ACCEPTED,
))
connection = _ClientSocket(response)
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
mock.patch.object(
host_agent_client, '_peer_credentials',
side_effect=lambda value: (0, value.peer_uid, 0),
), \
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
mock.patch.object(host_agent_client.socket, 'socket', return_value=connection):
result = HostAgentClient().dispatch(
operation_id=request.operation_id, action=request.action.value,
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,
)
self.assertEqual(result.status, HostAgentStatus.ACCEPTED)
self.assertEqual(decode_request_frame(bytes(connection.sent)), request)
self.assertEqual(connection.shutdown_how, socket.SHUT_WR)
self.assertTrue(connection.closed)
def test_client_rejects_nonroot_mismatch_rejection_and_unavailable_socket(self):
request = request_for()
cases = (
(_ClientSocket(encode_response_frame(HostAgentResponse(
request.operation_id, HostAgentStatus.ACCEPTED,
)), peer_uid=1000), HostAgentUnavailableError),
(_ClientSocket(encode_response_frame(HostAgentResponse(
str(uuid.uuid4()), HostAgentStatus.ACCEPTED,
))), HostAgentUnavailableError),
(_ClientSocket(encode_response_frame(HostAgentResponse(
request.operation_id, HostAgentStatus.REJECTED,
))), HostAgentRejectedError),
(_ClientSocket(encode_response_frame(HostAgentResponse(
request.operation_id, HostAgentStatus.UNAVAILABLE,
))), HostAgentUnavailableError),
)
for connection, error in cases:
with self.subTest(error=error.__name__), \
mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
mock.patch.object(
host_agent_client, '_peer_credentials',
side_effect=lambda value: (123, value.peer_uid, 0),
), \
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
mock.patch.object(host_agent_client.socket, 'socket', return_value=connection):
with self.assertRaises(error):
HostAgentClient().dispatch(
operation_id=request.operation_id, action=request.action.value,
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,
)
self.assertTrue(connection.closed)
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=False), \
mock.patch.object(host_agent_client.socket, 'socket') as socket_factory:
with self.assertRaises(HostAgentUnavailableError):
HostAgentClient().dispatch(
operation_id=request.operation_id, action=request.action.value,
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,
)
socket_factory.assert_not_called()
def test_client_rejects_response_pipelining_and_closes_on_baseexception(self):
request = request_for()
response = encode_response_frame(HostAgentResponse(
request.operation_id, HostAgentStatus.ACCEPTED,
))
pipelined = _ClientSocket(response + b'x')
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
mock.patch.object(host_agent_client, '_peer_credentials', return_value=(123, 0, 0)), \
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
mock.patch.object(host_agent_client.socket, 'socket', return_value=pipelined):
with self.assertRaises(HostAgentUnavailableError):
HostAgentClient().dispatch(
operation_id=request.operation_id, action=request.action.value,
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,
)
self.assertTrue(pipelined.closed)
interrupted = _ClientSocket(response)
interrupted.send = mock.Mock(side_effect=KeyboardInterrupt())
with mock.patch.object(host_agent_client, '_fixed_socket_is_safe', return_value=True), \
mock.patch.object(host_agent_client, '_peer_credentials', return_value=(123, 0, 0)), \
mock.patch.object(host_agent_client.socket, 'AF_UNIX', 1, create=True), \
mock.patch.object(host_agent_client.socket, 'socket', return_value=interrupted):
with self.assertRaises(KeyboardInterrupt):
HostAgentClient().dispatch(
operation_id=request.operation_id, action=request.action.value,
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,
)
self.assertTrue(interrupted.closed)
class HostAgentServerTests(unittest.TestCase):
def _exchange(self, request_bytes, *, credentials=(123, 10001, 10001), handler=None):
server, client = socket.socketpair()
result = []
def run():
with server, mock.patch.object(
host_agent_server, '_peer_credentials', return_value=credentials,
):
kwargs = {} if handler is None else {'handler': handler}
result.append(host_agent_server.serve_connection(server, **kwargs))
thread = threading.Thread(target=run)
thread.start()
try:
client.sendall(request_bytes)
client.shutdown(socket.SHUT_WR)
except (BrokenPipeError, ConnectionAbortedError, ConnectionResetError):
pass
response = bytearray()
while True:
try:
chunk = client.recv(4096)
except (ConnectionAbortedError, ConnectionResetError):
chunk = b''
if not chunk:
break
response.extend(chunk)
client.close()
thread.join(2)
self.assertFalse(thread.is_alive())
return result, bytes(response)
def test_authorized_valid_request_is_validated_but_unavailable_until_task_9_2(self):
request = request_for()
result, response = self._exchange(encode_request_frame(request))
self.assertEqual(result, [True])
self.assertEqual(
decode_response_frame(response),
HostAgentResponse(request.operation_id, HostAgentStatus.UNAVAILABLE),
)
def test_unauthorized_peer_is_rejected_before_body_read_and_gets_no_response(self):
request = request_for()
result, response = self._exchange(
encode_request_frame(request), credentials=(123, 1000, 10001),
)
self.assertEqual(result, [False])
self.assertEqual(response, b'')
server, client = socket.socketpair()
try:
with mock.patch.object(
host_agent_server, '_peer_credentials', return_value=(123, 1000, 10001),
), mock.patch.object(
host_agent_server, 'receive_frame',
side_effect=AssertionError('body must not be read'),
):
self.assertFalse(host_agent_server.serve_connection(server))
finally:
server.close()
client.close()
def test_invalid_and_trailing_requests_receive_only_bounded_invalid_response(self):
for frame in (
struct.pack('!I', 2) + b'{}',
encode_request_frame(request_for()) + b'x',
):
with self.subTest(frame=frame[:30]):
result, response = self._exchange(frame)
self.assertEqual(result, [True])
self.assertEqual(
decode_response_frame(response),
HostAgentResponse(None, HostAgentStatus.INVALID),
)
def test_deep_unknown_request_field_is_rejected_before_handler(self):
parsed = json.loads(encode_request_payload(request_for()))
nested = 'forbidden'
for _ in range(128):
nested = [nested]
parsed['options'] = nested
payload = json.dumps(
parsed, ensure_ascii=True, sort_keys=True,
separators=(',', ':'), allow_nan=False,
).encode('ascii')
self.assertLess(len(payload), MAX_REQUEST_PAYLOAD_BYTES)
handler = mock.Mock()
result, response = self._exchange(
struct.pack('!I', len(payload)) + payload,
handler=handler,
)
self.assertEqual(result, [True])
self.assertEqual(
decode_response_frame(response),
HostAgentResponse(None, HostAgentStatus.INVALID),
)
handler.assert_not_called()
def test_handler_cannot_change_operation_identity(self):
request = request_for()
result, response = self._exchange(
encode_request_frame(request),
handler=lambda _request: HostAgentResponse(
str(uuid.uuid4()), HostAgentStatus.ACCEPTED,
),
)
self.assertEqual(result, [True])
self.assertEqual(
decode_response_frame(response),
HostAgentResponse(request.operation_id, HostAgentStatus.UNAVAILABLE),
)
def test_listener_requires_exact_stream_type_and_accepting_state(self):
class FakeListener:
def __init__(self, *, socket_type=socket.SOCK_STREAM, accepting=1, family=None):
self.family = getattr(socket, 'AF_UNIX', 1) if family is None else family
self.socket_type = socket_type
self.accepting = accepting
def getsockopt(self, _level, option):
if option == socket.SO_TYPE:
return self.socket_type
if option == socket.SO_ACCEPTCONN:
return self.accepting
raise OSError('unexpected option')
def getsockname(self):
return host_agent_server.HOST_AGENT_SOCKET_PATH
details = mock.Mock(st_mode=__import__('stat').S_IFSOCK, st_uid=0)
cases = (
FakeListener(socket_type=getattr(socket, 'SOCK_SEQPACKET', 5)),
FakeListener(accepting=0),
FakeListener(family=socket.AF_INET),
)
for listener in cases:
with self.subTest(listener=listener.__dict__), \
mock.patch.object(host_agent_server.socket, 'socket', FakeListener), \
mock.patch.object(host_agent_server.socket, 'AF_UNIX', 1, create=True), \
mock.patch.object(host_agent_server.os, 'lstat', return_value=details):
with self.assertRaises(host_agent_server.HostAgentServerError):
host_agent_server._validate_listener(listener)
listener = FakeListener()
with mock.patch.object(host_agent_server.socket, 'socket', FakeListener), \
mock.patch.object(host_agent_server.socket, 'AF_UNIX', 1, create=True), \
mock.patch.object(host_agent_server.os, 'lstat', return_value=details):
host_agent_server._validate_listener(listener)
def test_invalid_inherited_listener_is_closed(self):
created = []
class ActivatedListener:
def __init__(self, *, fileno):
self.fileno = fileno
self.family = socket.AF_INET
self.closed = False
created.append(self)
def getsockopt(self, _level, option):
if option == socket.SO_TYPE:
return socket.SOCK_STREAM
if option == socket.SO_ACCEPTCONN:
return 1
raise OSError('unexpected option')
def close(self):
self.closed = True
environment = {
'LISTEN_PID': str(__import__('os').getpid()),
'LISTEN_FDS': '1',
}
with mock.patch.object(host_agent_server.os, 'geteuid', return_value=0, create=True), \
mock.patch.dict(host_agent_server.os.environ, environment, clear=True), \
mock.patch.object(host_agent_server.socket, 'socket', ActivatedListener), \
mock.patch.object(host_agent_server.socket, 'AF_UNIX', 1, create=True):
with self.assertRaises(host_agent_server.HostAgentServerError):
host_agent_server.inherited_systemd_listener()
self.assertEqual(len(created), 1)
self.assertEqual(created[0].fileno, host_agent_server.SYSTEMD_LISTEN_FD)
self.assertTrue(created[0].closed)
def test_server_stop_event_bounds_accept_loop(self):
stop_event = threading.Event()
class PollListener:
timeout = None
def settimeout(self, value):
self.timeout = value
def accept(self):
stop_event.set()
raise socket.timeout()
listener = PollListener()
with mock.patch.object(host_agent_server, '_validate_listener'):
host_agent_server.serve_forever(listener, stop_event=stop_event)
self.assertEqual(listener.timeout, host_agent_server.ACCEPT_POLL_SECONDS)
if __name__ == '__main__':
unittest.main()