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

407 lines
15 KiB
Python

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