import sys from pathlib import Path import unittest import uuid from unittest import mock ROOT = Path(__file__).resolve().parents[1] APP = ROOT / 'app' sys.path.insert(0, str(APP)) import host_agent_runtime import scanner_db from host_agent_protocol import HostAgentAction, HostAgentRequest, HostAgentStatus def request_for(action='apply-config'): return HostAgentRequest( operation_id=str(uuid.uuid4()), action=HostAgentAction(action), active_config_sha256='a' * 64, active_secrets_sha256='b' * 64, candidate_config_sha256='c' * 64 if action in ('apply-config', 'apply-both') else None, candidate_secrets_sha256='d' * 64 if action in ('apply-secrets', 'apply-both') else None, ) class _Database: def __init__(self): self.closed = False def close(self): self.closed = True class _Session: publication_state = 'original' def __init__(self, *_args, **_kwargs): self.closed = False self.entered = False def __enter__(self): self.entered = True return self def close(self): self.closed = True class _State: terminal = None def __init__(self, request): self.request = request self.initialized = [] def terminal_result(self): return self.terminal def initialize(self, publication_state): self.initialized.append(publication_state) return {'phase': 'prepared'} class _ValidationFailureSession(_Session): instances = [] def __init__(self, *_args, **_kwargs): super().__init__() self.claim = None self.__class__.instances.append(self) def __enter__(self): self.entered = True self.claim = {'replayed': False} raise host_agent_runtime.HostApplyError('validation') class _ValidationState(_State): instances = [] def __init__(self, request): super().__init__(request) self.phase = 'prepared' self.publication_state = None self.forward_category = None self.safe_detail = None self.result = None self.__class__.instances.append(self) def initialize(self, publication_state): if self.publication_state is None: self.publication_state = publication_state return self._phase() def _phase(self): return { 'phase': self.phase, 'publication_state': self.publication_state, 'forward_category': self.forward_category, 'safe_detail': self.safe_detail, } def advance( self, expected_phase, next_phase, publication_state, **evidence, ): if self.phase != expected_phase: raise AssertionError((self.phase, expected_phase, next_phase)) self.phase = next_phase self.publication_state = publication_state self.forward_category = evidence.get('forward_category') self.safe_detail = evidence.get('safe_detail') return self._phase() def publish_result(self, result, **evidence): self.result = {'result': result, **evidence} return self.result class _ReplayValidationState: record = None fail_next_result = True results = [] def __init__(self, request): self.request = request def terminal_result(self): return None def initialize(self, publication_state): if self.__class__.record is None: self.__class__.record = { 'phase': 'prepared', 'publication_state': publication_state, 'forward_category': None, 'safe_detail': None, } return dict(self.__class__.record) def advance( self, expected_phase, next_phase, publication_state, **evidence, ): if self.__class__.record['phase'] != expected_phase: raise AssertionError((self.__class__.record['phase'], expected_phase)) self.__class__.record = { 'phase': next_phase, 'publication_state': publication_state, 'forward_category': evidence.get('forward_category'), 'safe_detail': evidence.get('safe_detail'), } return dict(self.__class__.record) def publish_result(self, result, **evidence): if self.__class__.fail_next_result: self.__class__.fail_next_result = False raise OSError('simulated result publication failure') value = {'result': result, **evidence} self.__class__.results.append(value) return value class _Thread: instances = [] def __init__(self, *, target, args, name, daemon): self.target = target self.args = args self.name = name self.daemon = daemon self.started = False self.joined = False self.__class__.instances.append(self) def start(self): self.started = True def join(self): self.joined = True def run(self): self.target(*self.args) class HostAgentRuntimeTests(unittest.TestCase): def setUp(self): _Thread.instances.clear() _ValidationFailureSession.instances.clear() _ValidationState.instances.clear() _ReplayValidationState.record = None _ReplayValidationState.fail_next_result = True _ReplayValidationState.results.clear() def test_package_capabilities_map_only_fixed_roots(self): config = { 'supervisor': {'worker_api': {'compatibility_profiles': { 'linux': {'package_manifest': '/data/worker-packages/linux.json'}, }}}, } payloads = { Path('/etc/truf/worker-packages/linux.json'): b'linux', } with mock.patch.object(host_agent_runtime, 'load_yaml_document', return_value=config), \ mock.patch.object( host_agent_runtime, '_resolve_package_manifest_path', side_effect=lambda _config, value: value, ), \ mock.patch.object( host_agent_runtime, '_stable_root_file', side_effect=lambda path, _maximum: payloads[path], ) as stable, \ mock.patch.object( host_agent_runtime, 'load_worker_package_manifest_bytes', side_effect=lambda payload: {'capabilities': [payload.decode('ascii')]}, ): evidence = host_agent_runtime.load_fixed_package_capabilities(b'config') self.assertEqual(evidence, { 'linux': { 'package_manifest': '/data/worker-packages/linux.json', 'capabilities': ['linux'], }, }) self.assertEqual(stable.call_count, 1) config['supervisor']['worker_api']['compatibility_profiles']['linux'][ 'package_manifest' ] = '/data/runtime-linux/foreign.json' with mock.patch.object(host_agent_runtime, 'load_yaml_document', return_value=config), \ mock.patch.object( host_agent_runtime, '_resolve_package_manifest_path', side_effect=lambda _config, value: value, ): with self.assertRaisesRegex( host_agent_runtime.HostRuntimeError, 'host operation runtime failed', ): host_agent_runtime.load_fixed_package_capabilities(b'secret detail') config['supervisor']['worker_api']['compatibility_profiles']['linux'][ 'package_manifest' ] = '/opt/truf/app/packages/image.json' with mock.patch.object(host_agent_runtime, 'load_yaml_document', return_value=config), \ mock.patch.object( host_agent_runtime, '_resolve_package_manifest_path', side_effect=lambda _config, value: value, ): with self.assertRaises(host_agent_runtime.HostRuntimeError): host_agent_runtime.load_fixed_package_capabilities(b'config') def test_dispatch_accepts_only_after_claim_state_and_worker_start(self): request = request_for() database = _Database() with mock.patch.object( host_agent_runtime.ScannerDB, 'host_agent_authority', return_value=database, ), mock.patch.object(host_agent_runtime, 'HostApplySession', _Session), \ mock.patch.object(host_agent_runtime, 'HostOperationState', _State), \ mock.patch.object(host_agent_runtime.threading, 'Thread', _Thread), \ mock.patch.object(host_agent_runtime, 'execute_fixed_operation') as execute: dispatcher = host_agent_runtime.FixedHostOperationDispatcher() self.assertEqual(dispatcher.handle(request), HostAgentStatus.ACCEPTED) worker = _Thread.instances[0] self.assertTrue(worker.started) self.assertFalse(worker.daemon) self.assertEqual(dispatcher.handle(request), HostAgentStatus.ACCEPTED) self.assertEqual(dispatcher.handle(request_for()), HostAgentStatus.REJECTED) worker.run() execute.assert_called_once() self.assertTrue(database.closed) self.assertTrue(worker.args[0].closed) dispatcher.close() def test_terminal_replay_never_opens_database_and_start_failure_cleans_up(self): request = request_for('restart') terminal = {'result': 'succeeded'} with mock.patch.object(_State, 'terminal', terminal), \ mock.patch.object(host_agent_runtime, 'HostOperationState', _State), \ mock.patch.object( host_agent_runtime.ScannerDB, 'host_agent_authority', ) as authority: dispatcher = host_agent_runtime.FixedHostOperationDispatcher() self.assertEqual(dispatcher.handle(request), HostAgentStatus.ACCEPTED) authority.assert_not_called() database = _Database() with mock.patch.object(_State, 'terminal', None), \ mock.patch.object(host_agent_runtime, 'HostOperationState', _State), \ mock.patch.object( host_agent_runtime.ScannerDB, 'host_agent_authority', return_value=database, ), mock.patch.object(host_agent_runtime, 'HostApplySession', _Session), \ mock.patch.object(host_agent_runtime.threading, 'Thread', _Thread): dispatcher = host_agent_runtime.FixedHostOperationDispatcher() with mock.patch.object(_Thread, 'start', side_effect=RuntimeError('thread detail')): self.assertEqual(dispatcher.handle(request), HostAgentStatus.UNAVAILABLE) self.assertTrue(database.closed) def test_claimed_validation_failure_is_durably_terminal_before_acceptance(self): request = request_for('apply-secrets') database = _Database() with mock.patch.object( host_agent_runtime.ScannerDB, 'host_agent_authority', return_value=database, ), mock.patch.object( host_agent_runtime, 'HostApplySession', _ValidationFailureSession, ), mock.patch.object( host_agent_runtime, 'HostOperationState', _ValidationState, ), mock.patch.object(host_agent_runtime.threading, 'Thread') as thread: status = host_agent_runtime.FixedHostOperationDispatcher().handle(request) self.assertEqual(status, HostAgentStatus.ACCEPTED) thread.assert_not_called() self.assertTrue(database.closed) self.assertTrue(_ValidationFailureSession.instances[0].closed) state = _ValidationState.instances[0] self.assertEqual(state.phase, 'failed') self.assertEqual(state.result, { 'result': 'failed', 'safe_category': 'validation_failed', 'safe_detail': 'validation_failed', 'resulting_identity': { 'active_config_sha256': request.active_config_sha256, 'active_secrets_sha256': request.active_secrets_sha256, }, }) def test_validation_result_publication_failure_replays_idempotently(self): request = request_for('apply-config') databases = [_Database(), _Database()] with mock.patch.object( host_agent_runtime.ScannerDB, 'host_agent_authority', side_effect=databases, ), mock.patch.object( host_agent_runtime, 'HostApplySession', _ValidationFailureSession, ), mock.patch.object( host_agent_runtime, 'HostOperationState', _ReplayValidationState, ), mock.patch.object(host_agent_runtime.threading, 'Thread') as thread: dispatcher = host_agent_runtime.FixedHostOperationDispatcher() self.assertEqual( dispatcher.handle(request), HostAgentStatus.UNAVAILABLE, ) self.assertEqual(_ReplayValidationState.record, { 'phase': 'failed', 'publication_state': 'original', 'forward_category': 'validation_failed', 'safe_detail': 'validation_failed', }) self.assertEqual( dispatcher.handle(request), HostAgentStatus.ACCEPTED, ) thread.assert_not_called() self.assertTrue(all(database.closed for database in databases)) self.assertEqual(len(_ReplayValidationState.results), 1) self.assertEqual( _ReplayValidationState.results[0]['resulting_identity'], { 'active_config_sha256': request.active_config_sha256, 'active_secrets_sha256': request.active_secrets_sha256, }, ) def test_validation_failure_never_rewrites_nonoriginal_prepared_state(self): request = request_for() state = _ValidationState(request) state.phase = 'prepared' state.publication_state = 'candidate' self.assertFalse( host_agent_runtime.FixedHostOperationDispatcher. _record_validation_failure(state, request), ) self.assertEqual(state.phase, 'prepared') self.assertEqual(state.publication_state, 'candidate') self.assertIsNone(state.result) def test_cleanup_failures_do_not_strand_dispatcher_or_skip_database_close(self): request = request_for() database = _Database() session = _Session() state = _State(request) dispatcher = host_agent_runtime.FixedHostOperationDispatcher() dispatcher._active_request = b'active' dispatcher._worker = object() with mock.patch.object( session, 'close', side_effect=RuntimeError('session detail'), ), mock.patch.object( database, 'close', side_effect=RuntimeError('database detail'), ) as close_database, mock.patch.object( host_agent_runtime, 'execute_fixed_operation', ): dispatcher._execute(session, database, state) close_database.assert_called_once_with() self.assertIsNone(dispatcher._active_request) self.assertIsNone(dispatcher._worker) def test_database_authority_closes_connection_when_session_setup_fails(self): connection = mock.Mock() connection.execute.side_effect = RuntimeError('setup detail') with mock.patch.object( scanner_db, 'connect_host_agent_postgres', return_value=connection, ): with self.assertRaisesRegex(RuntimeError, 'setup detail'): scanner_db.ScannerDB.host_agent_authority() connection.close.assert_called_once_with() if __name__ == '__main__': unittest.main()