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