import copy import hashlib import os from pathlib import Path import socket import stat import sys import tempfile import threading from types import SimpleNamespace import unittest from unittest import mock APP_DIR = os.path.abspath(os.path.join(os.path.dirname(__file__), '..', 'app')) if APP_DIR not in sys.path: sys.path.insert(0, APP_DIR) import managed_files class ManagedFileConfigurationTests(unittest.TestCase): def limits(self): return { 'max_relative_path_bytes': 1024, 'max_component_bytes': 255, 'max_path_depth': 16, 'max_listing_entries': 500, 'max_listing_bytes': 262144, 'max_file_bytes': 64 * 1024 * 1024, } def root(self, path='/data/managed-files/exports', **permissions): return { 'path': path, 'permissions': { 'list': True, 'read': True, 'create_replace': True, 'delete': True, **permissions, }, 'limits': self.limits(), } def assert_error(self, value, category, field=None): with self.assertRaises(managed_files.ManagedFileConfigurationError) as raised: managed_files.normalize_managed_file_roots(value) self.assertEqual(raised.exception.category, category) if field is not None: self.assertEqual(raised.exception.field, field) self.assertEqual(str(raised.exception), 'managed file root configuration is invalid') def test_empty_registry_and_deterministic_lookup(self): normalized, registry = managed_files.normalize_managed_file_roots({}) self.assertEqual(normalized, {}) self.assertEqual(registry, managed_files.ManagedFileRootRegistry()) self.assertEqual(registry.root_ids(), ()) self.assertIsNone(registry.get('missing')) self.assertIsNone(registry.get(1)) value = { 'z-export': self.root('/data/managed-files/z-export'), 'a-export': self.root('/data/managed-files/a-export'), } normalized, registry = managed_files.normalize_managed_file_roots(value) self.assertEqual(tuple(normalized), ('a-export', 'z-export')) self.assertEqual(registry.root_ids(), ('a-export', 'z-export')) selected = registry.get('a-export') self.assertEqual(selected.absolute_path, '/data/managed-files/a-export') for operation in managed_files.ManagedFileOperation: self.assertTrue(selected.permissions.allows(operation)) self.assertFalse(selected.permissions.allows('read')) self.assert_error(None, 'type', ('root',)) def test_predefined_runtime_roots_are_exact_and_read_only(self): predefined = { 'runtime-keychecks': '/data/runtime-linux/keychecks', 'runtime-logs': '/data/runtime-linux/logs', 'runtime-results': '/data/runtime-linux/results', } value = { root_id: self.root(path, create_replace=False, delete=False) for root_id, path in predefined.items() } normalized, registry = managed_files.normalize_managed_file_roots(value) self.assertEqual(registry.root_ids(), tuple(sorted(predefined))) for root_id, path in predefined.items(): root = registry.get(root_id) self.assertEqual(normalized[root_id]['path'], path) self.assertTrue(root.permissions.allows(managed_files.ManagedFileOperation.LIST)) self.assertTrue(root.permissions.allows(managed_files.ManagedFileOperation.READ)) self.assertFalse(root.permissions.allows( managed_files.ManagedFileOperation.CREATE_REPLACE, )) self.assertFalse(root.permissions.allows(managed_files.ManagedFileOperation.DELETE)) self.assertNotIn(path, repr(root)) for bad in ( {root_id: self.root(path)}, {root_id.removeprefix('runtime-'): self.root( path, create_replace=False, delete=False, )}, {root_id: self.root(f'/data/managed-files/{root_id}')}, ): self.assert_error(bad, 'deployment_path', ('root', 'path')) result_root = self.root( predefined['runtime-results'], create_replace=False, delete=False, ) result_root['limits']['max_file_bytes'] = ( managed_files.MAX_RESULT_FILE_BYTES ) _, registry = managed_files.normalize_managed_file_roots({ 'runtime-results': result_root, }) self.assertEqual( registry.get('runtime-results').limits.max_file_bytes, managed_files.MAX_RESULT_FILE_BYTES, ) dynamic = self.root() dynamic['limits']['max_file_bytes'] = managed_files.MAX_RESULT_FILE_BYTES self.assert_error({'exports': dynamic}, 'bounds') def test_root_ids_shape_and_count_are_bounded(self): for root_id in ('', 'Upper', 'under_score', '1-start', 'a' * 65, 1): self.assert_error({root_id: self.root()}, 'bounds', ('root', 'id')) too_many = { f'root-{index}': self.root(f'/data/managed-files/root-{index}') for index in range(managed_files.MAX_MANAGED_ROOTS + 1) } self.assert_error(too_many, 'bounds', ('root',)) def test_root_permission_and_limit_shapes_are_exact(self): root = self.root() for section in ('path', 'permissions', 'limits'): missing = copy.deepcopy(root) del missing[section] self.assert_error({'exports': missing}, 'schema') extra = copy.deepcopy(root) extra['command'] = 'fixture' self.assert_error({'exports': extra}, 'unknown_key') for key in tuple(root['permissions']): missing = copy.deepcopy(root) del missing['permissions'][key] self.assert_error({'exports': missing}, 'schema') invalid = copy.deepcopy(root) invalid['permissions']['read'] = 1 self.assert_error({'exports': invalid}, 'type') extra = copy.deepcopy(root) extra['permissions']['execute'] = False self.assert_error({'exports': extra}, 'unknown_key') for key in tuple(root['limits']): missing = copy.deepcopy(root) del missing['limits'][key] self.assert_error({'exports': missing}, 'schema') for value in (True, 0, getattr(managed_files, key.upper()) + 1): invalid = copy.deepcopy(root) invalid['limits'][key] = value self.assert_error({'exports': invalid}, 'bounds') extra = copy.deepcopy(root) extra['limits']['max_downloads'] = 1 self.assert_error({'exports': extra}, 'unknown_key') invalid = copy.deepcopy(root) invalid['limits']['max_relative_path_bytes'] = 128 self.assert_error({'exports': invalid}, 'bounds') invalid = copy.deepcopy(root) invalid['limits']['max_listing_bytes'] = 128 self.assert_error({'exports': invalid}, 'bounds') def test_paths_are_literal_normalized_and_narrowly_allowlisted(self): invalid_paths = ( '', 'relative', 'C:\\data', '/C:/data', '/', '/data/managed-files/', '/data//managed-files', '/data/./managed-files', '/data/managed-files/../config', '/data/managed-files/bad\\name', '/data/managed-files/bad\x00name', '/data/managed-files/' + 'a' * 256, '/data/managed-files/nested/child', ) for path in invalid_paths: self.assert_error( {'exports': self.root(path)}, 'deployment_path', ('root', 'path'), ) forbidden = ( '/data', '/data/config', '/data/config/child', '/data/secrets', '/data/secrets/runtime.yaml', '/data/runtime-document-candidates', '/data/postgres-linux', '/data/runtime-document-candidates/config.yaml', '/data/runtime-linux', '/data/runtime-linux/postgres', '/data/runtime-linux/results/private', '/data/runtime-linux/result_spool', '/data/runtime-linux/state', '/data/runtime-linux/keychecks/private', '/data/runtime-linux/queues', '/data/runtime-linux/postman_cache', '/data/scanner-result-bundles', '/data/scanner-result-bundles/private/archive', '/data/scanner-work', '/data/host-agent', '/data/host-agent/operations', '/data/windows-archive', '/opt', '/opt/truf', '/opt/truf/app', '/run', '/run/truf/host-agent.sock', '/var/run', '/var/run/docker.sock', ) for path in forbidden: self.assert_error( {'exports': self.root(path)}, 'deployment_path', ('root', 'path'), ) duplicate = { 'first': self.root('/data/managed-files/shared'), 'second': self.root('/data/managed-files/shared'), } self.assert_error(duplicate, 'deployment_path', ('root', 'path')) def test_configuration_parser_performs_no_filesystem_work(self): config = {'supervisor': {'worker_api': {'admin': { 'managed_file_roots': {'exports': self.root()}, }}}} with mock.patch('builtins.open', side_effect=AssertionError('open')), \ mock.patch('os.stat', side_effect=AssertionError('stat')), \ mock.patch('os.lstat', side_effect=AssertionError('lstat')), \ mock.patch.object( Path, 'resolve', side_effect=AssertionError('resolve'), ): registry = managed_files.managed_file_root_registry_from_config(config) self.assertEqual(registry.root_ids(), ('exports',)) def test_configuration_parser_rejects_malformed_parent_sections(self): for config in (None, [], 'config'): with self.assertRaises(managed_files.ManagedFileConfigurationError): managed_files.managed_file_root_registry_from_config(config) for value in (None, [], '', 'invalid', 0, False): configs = [ {'supervisor': value}, {'supervisor': {'worker_api': value}}, ] if value is not None: configs.append({'supervisor': {'worker_api': {'admin': value}}}) for config in configs: with self.assertRaises(managed_files.ManagedFileConfigurationError) as raised: managed_files.managed_file_root_registry_from_config(config) self.assertEqual(raised.exception.category, 'type') self.assertEqual( managed_files.managed_file_root_registry_from_config({}), managed_files.ManagedFileRootRegistry(), ) self.assertEqual( managed_files.managed_file_root_registry_from_config({ 'supervisor': {'worker_api': {'admin': None}}, }), managed_files.ManagedFileRootRegistry(), ) def test_errors_do_not_echo_ids_or_paths(self): sentinel_id = 'private-root' sentinel_path = '/private/sentinel/path' with self.assertRaises(managed_files.ManagedFileConfigurationError) as raised: managed_files.normalize_managed_file_roots({ sentinel_id: self.root(sentinel_path), }) text = str(raised.exception) self.assertNotIn(sentinel_id, text) self.assertNotIn(sentinel_path, text) class _ManagedFileTraversalTestHelpers: def limits(self, **overrides): values = { 'max_relative_path_bytes': 1024, 'max_component_bytes': 255, 'max_path_depth': 16, 'max_listing_entries': 500, 'max_listing_bytes': 262144, 'max_file_bytes': 64 * 1024 * 1024, } values.update(overrides) return managed_files.ManagedFileLimits(**values) def root(self, path='/fixture/root', **permissions): return managed_files.ManagedFileRoot( root_id='exports', absolute_path=path, permissions=managed_files.ManagedFilePermissions( allow_list=permissions.get('allow_list', True), allow_read=permissions.get('allow_read', True), allow_create_replace=permissions.get('allow_create_replace', False), allow_delete=permissions.get('allow_delete', False), ), limits=self.limits(), ) def assert_access_error(self, category, callback): with self.assertRaises(managed_files.ManagedFileAccessError) as raised: callback() self.assertEqual(raised.exception.category, category) self.assertEqual(str(raised.exception), 'managed file access failed') class ManagedFileTraversalTests(_ManagedFileTraversalTestHelpers, unittest.TestCase): def test_relative_paths_are_canonical_and_utf8_bounded(self): limits = self.limits() self.assertEqual( managed_files.parse_managed_relative_path('nested/caf\N{LATIN SMALL LETTER E WITH ACUTE}.txt', limits), ('nested', 'caf\N{LATIN SMALL LETTER E WITH ACUTE}.txt'), ) invalid = ( '', '/absolute', 'C:drive', 'nested/C:drive', 'bad\\name', 'bad\x00name', '.', '..', './file', 'dir/../file', 'dir//file', 'dir/', 'a/b/c', '12345', 'a/\ud800', '.truf-managed-file-private.tmp', ) bounded = self.limits( max_relative_path_bytes=4, max_component_bytes=4, max_path_depth=2, ) for value in invalid: self.assert_access_error( 'invalid_path', lambda value=value: managed_files.parse_managed_relative_path( value, bounded, ), ) self.assert_access_error( 'invalid_path', lambda: managed_files.parse_managed_relative_path(1, limits), ) with self.assertRaises(managed_files.ManagedFileAccessError) as raised: managed_files.parse_managed_relative_path('../private-path', limits) self.assertNotIn('private-path', str(raised.exception)) def test_descriptor_walk_uses_exact_relative_nofollow_opens(self): flags = {'directory': 1, 'list': 2, 'inspect': 4, 'read': 8} opened = mock.Mock(side_effect=(10, 11, 12, 21, 22, 23)) def fstat(descriptor): regular = descriptor in (22, 23) return SimpleNamespace( st_mode=stat.S_IFREG if regular else stat.S_IFDIR, st_nlink=1, st_dev=100 if regular else descriptor, st_ino=200 if regular else descriptor, ) with mock.patch.object(managed_files, '_descriptor_flags', return_value=flags), \ mock.patch.object(managed_files.os, 'open', opened), \ mock.patch.object(managed_files.os, 'dup', return_value=20) as duplicated, \ mock.patch.object(managed_files.os, 'fstat', side_effect=fstat), \ mock.patch.object(managed_files.os, 'close') as closed: traversal = managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((self.root(),)), ) with traversal.open_read_file('exports', 'nested/file.txt') as target: self.assertEqual(target.fileno(), 23) self.assertEqual(repr(target), '') self.assertNotIn('23', repr(target)) self.assert_access_error('closed', target.fileno) traversal.close() traversal.close() self.assertEqual(opened.call_args_list, [ mock.call('/', 1), mock.call('fixture', 1, dir_fd=10), mock.call('root', 1, dir_fd=11), mock.call('nested', 1, dir_fd=20), mock.call('file.txt', 4, dir_fd=21), mock.call('/proc/self/fd/22', 8), ]) duplicated.assert_called_once_with(12) for descriptor in (10, 11, 20, 21, 22, 23, 12): self.assertIn(mock.call(descriptor), closed.call_args_list) def test_base_exception_closes_partial_root_and_operation_descriptors(self): class Cancelled(BaseException): pass flags = {'directory': 1, 'list': 2, 'inspect': 4, 'read': 8} with mock.patch.object(managed_files, '_descriptor_flags', return_value=flags), \ mock.patch.object( managed_files.os, 'open', side_effect=(10, KeyboardInterrupt()), ), \ mock.patch.object( managed_files.os, 'fstat', return_value=SimpleNamespace(st_mode=stat.S_IFDIR), ), \ mock.patch.object(managed_files.os, 'close') as closed: with self.assertRaises(KeyboardInterrupt): managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((self.root(),)), ) closed.assert_called_once_with(10) opened = mock.Mock(side_effect=(10, 11, 12, 21, 22)) def fstat(descriptor): if descriptor == 22: raise Cancelled() return SimpleNamespace(st_mode=stat.S_IFDIR, st_nlink=1) with mock.patch.object(managed_files, '_descriptor_flags', return_value=flags), \ mock.patch.object(managed_files.os, 'open', opened), \ mock.patch.object(managed_files.os, 'dup', return_value=20), \ mock.patch.object(managed_files.os, 'fstat', side_effect=fstat), \ mock.patch.object(managed_files.os, 'close') as closed: traversal = managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((self.root(),)), ) with self.assertRaises(Cancelled): with traversal.open_read_file('exports', 'nested/file.txt'): self.fail('unreachable') traversal.close() for descriptor in (20, 21, 22, 12): self.assertIn(mock.call(descriptor), closed.call_args_list) opened = mock.Mock(side_effect=(10, 11, 12, 21)) close_calls = [] interrupted = False def close(descriptor): nonlocal interrupted close_calls.append(descriptor) if descriptor == 20 and not interrupted: interrupted = True raise Cancelled() with mock.patch.object(managed_files, '_descriptor_flags', return_value=flags), \ mock.patch.object(managed_files.os, 'open', opened), \ mock.patch.object(managed_files.os, 'dup', return_value=20), \ mock.patch.object( managed_files.os, 'fstat', return_value=SimpleNamespace(st_mode=stat.S_IFDIR, st_nlink=1), ), \ mock.patch.object(managed_files.os, 'close', side_effect=close): traversal = managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((self.root(),)), ) with self.assertRaises(Cancelled): with traversal.open_read_file('exports', 'nested/file.txt'): self.fail('unreachable') traversal.close() self.assertEqual(close_calls.count(20), 1) self.assertIn(21, close_calls) self.assertIn(12, close_calls) def test_proc_descriptor_reopen_failure_is_not_target_absence(self): flags = {'directory': 1, 'list': 2, 'inspect': 4, 'read': 8} opened = mock.Mock(side_effect=(10, 11, 12, 22, FileNotFoundError())) def fstat(descriptor): regular = descriptor == 22 return SimpleNamespace( st_mode=stat.S_IFREG if regular else stat.S_IFDIR, st_nlink=1, st_dev=100, st_ino=200, ) with mock.patch.object(managed_files, '_descriptor_flags', return_value=flags), \ mock.patch.object(managed_files.os, 'open', opened), \ mock.patch.object(managed_files.os, 'dup', return_value=20), \ mock.patch.object(managed_files.os, 'fstat', side_effect=fstat), \ mock.patch.object(managed_files.os, 'close'): traversal = managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((self.root(),)), ) self.assert_access_error( 'filesystem_unavailable', lambda: self._enter_context( traversal.open_read_file('exports', 'file.txt'), ), ) traversal.close() def test_nonregular_device_modes_are_rejected_before_read_reopen(self): traversal = object.__new__(managed_files.ManagedFileTraversal) traversal._flags = {'inspect': 4, 'read': 8} for mode in ( stat.S_IFIFO, stat.S_IFSOCK, stat.S_IFCHR, stat.S_IFBLK, stat.S_IFDIR): opened = mock.Mock(return_value=17) with self.subTest(mode=mode), mock.patch.object( managed_files.os, 'open', opened, ), mock.patch.object( managed_files.os, 'fstat', return_value=SimpleNamespace( st_mode=mode | 0o600, st_nlink=1, ), ), mock.patch.object(managed_files.os, 'close'): self.assert_access_error( 'unsafe_target', lambda: traversal._open_file_at(9, 'unsafe-target'), ) opened.assert_called_once_with('unsafe-target', 4, dir_fd=9) @staticmethod def _enter_context(context): with context: return None def test_empty_registry_is_portable_and_nonempty_registry_fails_closed(self): with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry()) as traversal: self.assert_access_error( 'unknown_root', lambda: traversal.open_read_file('missing', 'file'), ) with mock.patch.object(managed_files.sys, 'platform', 'win32'): self.assert_access_error( 'filesystem_unavailable', lambda: managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((self.root(),)), ), ) @unittest.skipUnless(sys.platform.startswith('linux'), 'Linux descriptor semantics required') class ManagedFileTraversalLinuxTests( _ManagedFileTraversalTestHelpers, unittest.TestCase): def setUp(self): self.temporary = tempfile.TemporaryDirectory() self.root_path = os.path.join(self.temporary.name, 'root') os.mkdir(self.root_path) self.root_value = self.root(self.root_path) self.registry = managed_files.ManagedFileRootRegistry((self.root_value,)) def tearDown(self): self.temporary.cleanup() def test_list_read_and_retained_root_survive_path_replacement(self): nested = os.path.join(self.root_path, 'nested') os.mkdir(nested) with open(os.path.join(nested, 'file.txt'), 'wb') as handle: handle.write(b'original') with managed_files.ManagedFileTraversal(self.registry) as traversal: with traversal.open_list_directory('exports') as opened: self.assertEqual(os.listdir(opened.fileno()), ['nested']) with traversal.open_list_directory('exports', 'nested') as opened: self.assertEqual(os.listdir(opened.fileno()), ['file.txt']) with traversal.open_read_file('exports', 'nested/file.txt') as opened: self.assertEqual(os.read(opened.fileno(), 32), b'original') moved = os.path.join(self.temporary.name, 'moved-root') os.rename(self.root_path, moved) os.mkdir(self.root_path) with open(os.path.join(self.root_path, 'replacement.txt'), 'wb') as handle: handle.write(b'replacement') with traversal.open_read_file('exports', 'nested/file.txt') as opened: self.assertEqual(os.read(opened.fileno(), 32), b'original') self.assert_access_error( 'not_found', lambda: self._enter(traversal.open_read_file( 'exports', 'replacement.txt', )), ) def test_links_and_special_files_are_rejected(self): outside = os.path.join(self.temporary.name, 'outside') os.mkdir(outside) with open(os.path.join(outside, 'outside.txt'), 'wb') as handle: handle.write(b'outside') os.symlink(outside, os.path.join(self.root_path, 'link-dir')) os.symlink( os.path.join(outside, 'outside.txt'), os.path.join(self.root_path, 'link-file'), ) with open(os.path.join(self.root_path, 'linked.txt'), 'wb') as handle: handle.write(b'linked') os.link( os.path.join(self.root_path, 'linked.txt'), os.path.join(self.root_path, 'hardlink.txt'), ) os.mkdir(os.path.join(self.root_path, 'directory')) os.mkfifo(os.path.join(self.root_path, 'fifo')) unix_socket = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) special_paths = [ 'link-dir/outside.txt', 'link-file', 'linked.txt', 'hardlink.txt', 'directory', 'fifo', ] try: unix_socket.bind(os.path.join(self.root_path, 'socket')) special_paths.append('socket') except AssertionError as exc: if str(exc) != 'container-unit permits loopback sockets only': raise unix_socket.close() unix_socket = None try: with managed_files.ManagedFileTraversal(self.registry) as traversal: for path in special_paths: self.assert_access_error( 'unsafe_target', lambda path=path: self._enter( traversal.open_read_file('exports', path), ), ) self.assert_access_error( 'unsafe_target', lambda: self._enter(traversal.open_list_directory( 'exports', 'link-dir', )), ) finally: if unix_socket is not None: unix_socket.close() root_link = os.path.join(self.temporary.name, 'root-link') os.symlink(self.root_path, root_link) linked_registry = managed_files.ManagedFileRootRegistry(( self.root(root_link), )) self.assert_access_error( 'root_unavailable', lambda: managed_files.ManagedFileTraversal(linked_registry), ) def test_hardlinks_and_special_targets_are_rejected_by_all_file_operations(self): regular = os.path.join(self.root_path, 'regular.bin') hardlink = os.path.join(self.root_path, 'hardlink.bin') directory = os.path.join(self.root_path, 'safe-directory') fifo = os.path.join(self.root_path, 'named-pipe') symlink = os.path.join(self.root_path, 'linked.bin') Path(regular).write_bytes(b'linked-content') os.link(regular, hardlink) os.mkdir(directory) os.mkfifo(fifo) os.symlink(regular, symlink) writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) registry = managed_files.ManagedFileRootRegistry((writable,)) expected = hashlib.sha256(b'linked-content').hexdigest() with managed_files.ManagedFileTraversal(registry) as traversal: listing = traversal.list_directory('exports') self.assertEqual( tuple((entry.name, entry.kind) for entry in listing.entries), (('safe-directory', 'directory'),), ) for path in ('regular.bin', 'hardlink.bin', 'safe-directory', 'named-pipe', 'linked.bin'): with self.subTest(path=path, operation='download'): self.assert_access_error( 'unsafe_target', lambda path=path: traversal.download_file('exports', path), ) with self.subTest(path=path, operation='replace'): self.assert_access_error( 'unsafe_target', lambda path=path: traversal.create_replace_file( 'exports', path, b'replacement', expected_sha256=expected, ), ) with self.subTest(path=path, operation='delete'): self.assert_access_error( 'unsafe_target', lambda path=path: traversal.delete_file( 'exports', path, expected_sha256=expected, ), ) def test_component_swap_remains_anchored_to_open_descriptor(self): live = os.path.join(self.root_path, 'live') outside = os.path.join(self.temporary.name, 'outside') os.mkdir(live) os.mkdir(outside) with open(os.path.join(live, 'file.txt'), 'wb') as handle: handle.write(b'inside') with open(os.path.join(outside, 'file.txt'), 'wb') as handle: handle.write(b'outside') original_open = os.open swapped = False def racing_open(path, flags, *args, **kwargs): nonlocal swapped if path == 'file.txt' and not swapped: swapped = True os.rename(live, os.path.join(self.root_path, 'original')) os.symlink(outside, live) return original_open(path, flags, *args, **kwargs) with managed_files.ManagedFileTraversal(self.registry) as traversal, \ mock.patch.object(managed_files.os, 'open', side_effect=racing_open): with traversal.open_read_file('exports', 'live/file.txt') as opened: self.assertEqual(os.read(opened.fileno(), 32), b'inside') self.assertTrue(swapped) def test_root_and_final_symlink_swaps_never_follow_targets(self): outside = os.path.join(self.temporary.name, 'outside-race') os.mkdir(outside) Path(os.path.join(outside, 'file.bin')).write_bytes(b'outside') original_open = managed_files.os.open descriptor_flags = managed_files._descriptor_flags() moved_root = os.path.join(self.temporary.name, 'moved-root-race') swapped_root = False def swap_root_before_open(path, flags, *args, **kwargs): nonlocal swapped_root if path == 'root' and kwargs.get('dir_fd') is not None and not swapped_root: swapped_root = True os.rename(self.root_path, moved_root) os.symlink(outside, self.root_path) return original_open(path, flags, *args, **kwargs) with mock.patch.object( managed_files, '_descriptor_flags', return_value=descriptor_flags, ), mock.patch.object( managed_files.os, 'open', side_effect=swap_root_before_open, ): self.assert_access_error( 'root_unavailable', lambda: managed_files.ManagedFileTraversal(self.registry), ) self.assertTrue(swapped_root) os.unlink(self.root_path) os.rename(moved_root, self.root_path) target = os.path.join(self.root_path, 'file.bin') retained = os.path.join(self.root_path, 'retained.bin') Path(target).write_bytes(b'inside') swapped_final = False def swap_final_before_reopen(path, flags, *args, **kwargs): nonlocal swapped_final if str(path).startswith('/proc/self/fd/') and not swapped_final: swapped_final = True os.rename(target, retained) os.symlink(os.path.join(outside, 'file.bin'), target) return original_open(path, flags, *args, **kwargs) with managed_files.ManagedFileTraversal(self.registry) as traversal, \ mock.patch.object( managed_files.os, 'open', side_effect=swap_final_before_reopen, ): download = traversal.download_file('exports', 'file.bin') self.assertTrue(swapped_final) self.assertEqual(download.content, b'inside') self.assertNotEqual(download.content, b'outside') def test_permission_validation_and_close_precede_descriptor_duplication(self): with managed_files.ManagedFileTraversal(self.registry) as traversal, \ mock.patch.object( managed_files.os, 'dup', wraps=managed_files.os.dup, ) as duplicated: self.assert_access_error( 'unknown_root', lambda: traversal.open_read_file('missing', 'file'), ) self.assert_access_error( 'invalid_path', lambda: traversal.open_read_file('exports', '../file'), ) duplicated.assert_not_called() self.assert_access_error( 'closed', lambda: self._enter(traversal.open_read_file('exports', 'file')), ) denied = self.root(self.root_path, allow_list=False, allow_read=False) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((denied,))) as traversal: self.assert_access_error( 'operation_not_allowed', lambda: traversal.open_read_file('exports', 'file'), ) self.assert_access_error( 'operation_not_allowed', lambda: traversal.open_list_directory('exports'), ) def test_bounded_listing_and_download_return_safe_metadata(self): os.mkdir(os.path.join(self.root_path, 'directory')) with open(os.path.join(self.root_path, 'z.txt'), 'wb') as handle: handle.write(b'z-value') with open(os.path.join(self.root_path, 'a.txt'), 'wb') as handle: handle.write(b'a-value') with open( os.path.join(self.root_path, '.truf-managed-file-hidden.tmp'), 'wb') as handle: handle.write(b'incomplete') with open(os.path.join(self.root_path, 'linked.txt'), 'wb') as handle: handle.write(b'linked') os.link( os.path.join(self.root_path, 'linked.txt'), os.path.join(self.root_path, 'linked-again.txt'), ) with managed_files.ManagedFileTraversal(self.registry) as traversal: listing = traversal.list_directory('exports') self.assertEqual( tuple((entry.name, entry.kind, entry.byte_count) for entry in listing.entries), ( ('a.txt', 'file', 7), ('directory', 'directory', None), ('z.txt', 'file', 7), ), ) self.assertEqual( listing.name_bytes, sum(len(entry.name.encode('utf-8')) for entry in listing.entries), ) download = traversal.download_file('exports', 'a.txt') self.assertEqual(download.content, b'a-value') self.assertEqual( download.identity, managed_files.ManagedFileIdentity( hashlib.sha256(b'a-value').hexdigest(), 7, ), ) self.assertNotIn('a-value', repr(download)) bounded_root = managed_files.ManagedFileRoot( root_id='exports', absolute_path=self.root_path, permissions=self.root_value.permissions, limits=self.limits( max_listing_entries=2, max_listing_bytes=255, max_file_bytes=4, ), ) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((bounded_root,))) as traversal: self.assert_access_error( 'limit_exceeded', lambda: traversal.list_directory('exports'), ) self.assert_access_error( 'limit_exceeded', lambda: traversal.download_file('exports', 'a.txt'), ) def test_result_projection_root_is_allowlisted_streamed_and_serialized(self): allowed = { 'found_secrets.jsonl': b'found-active', 'found_secrets.g000014.jsonl': b'found-rotated', 'scan_results.jsonl': b'scan-active', 'scan_results.g000123.jsonl': b'scan-rotated', } forbidden = { '.jsonl-projector.lock': b'lock', 'scan_errors.log': b'errors', 'scanner_active.db': b'database', 'found_secrets.ledger': b'ledger', 'scan_results.g123.jsonl': b'bad-generation', 'scan_results.g000001.jsonl.recovery': b'recovery', } for name, payload in {**allowed, **forbidden}.items(): Path(os.path.join(self.root_path, name)).write_bytes(payload) os.mkdir(os.path.join(self.root_path, '.projection-tmp')) os.mkdir(os.path.join(self.root_path, '.projection-quarantine')) os.mkdir(os.path.join(self.root_path, 'scan_results.g999999.jsonl')) result_root = managed_files.ManagedFileRoot( managed_files.RUNTIME_RESULT_ROOT_ID, self.root_path, managed_files.ManagedFilePermissions(True, True, False, False), self.limits(max_file_bytes=managed_files.MAX_RESULT_FILE_BYTES), ) reads = [] original_read = managed_files.os.read def bounded_read(descriptor, size): reads.append(size) return original_read(descriptor, size) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((result_root,))) as traversal, \ mock.patch.object(managed_files.os, 'read', side_effect=bounded_read): listing = traversal.list_directory(managed_files.RUNTIME_RESULT_ROOT_ID) self.assertEqual( tuple(entry.name for entry in listing.entries), tuple(sorted(allowed)), ) for name in ( *forbidden, '.projection-tmp', '.projection-quarantine'): self.assert_access_error( 'not_found', lambda name=name: traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, name, ), ) self.assert_access_error( 'unsafe_target', lambda: traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'scan_results.g999999.jsonl', ), ) self.assert_access_error( 'not_found', lambda: traversal.list_directory( managed_files.RUNTIME_RESULT_ROOT_ID, '.projection-tmp', ), ) first = traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'scan_results.jsonl', ) self.assertIsNone(first.content) self.assertIsNotNone(first.snapshot) self.assertNotIn('scan-active', repr(first)) self.assert_access_error( 'download_busy', lambda: traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'found_secrets.jsonl', ), ) snapshot_details = os.fstat(first.snapshot._handle.fileno()) root_details = os.stat(self.root_path) self.assertEqual(snapshot_details.st_dev, root_details.st_dev) self.assertEqual(snapshot_details.st_nlink, 0) self.assertEqual(b''.join(first.snapshot.chunks()), b'scan-active') second = traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'found_secrets.g000014.jsonl', ) self.assertEqual( b''.join(second.snapshot.chunks()), b'found-rotated', ) self.assertTrue(reads) self.assertLessEqual(max(reads), managed_files._READ_CHUNK_BYTES) def test_result_snapshot_retry_failure_and_shutdown_release_resources(self): Path(os.path.join(self.root_path, 'scan_results.jsonl')).write_bytes( b'stable-result', ) result_root = managed_files.ManagedFileRoot( managed_files.RUNTIME_RESULT_ROOT_ID, self.root_path, managed_files.ManagedFilePermissions(True, True, False, False), self.limits(max_file_bytes=managed_files.MAX_RESULT_FILE_BYTES), ) traversal = managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((result_root,)), ) original_snapshot = traversal._snapshot_descriptor attempts = 0 def fail_once(root, descriptor): nonlocal attempts attempts += 1 if attempts == 1: raise managed_files.ManagedFileAccessError('concurrent_change') return original_snapshot(root, descriptor) with mock.patch.object( traversal, '_snapshot_descriptor', side_effect=fail_once): retried = traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'scan_results.jsonl', ) self.assertEqual(attempts, 2) self.assertEqual(b''.join(retried.snapshot.chunks()), b'stable-result') with mock.patch.object( traversal, '_snapshot_descriptor', side_effect=managed_files.ManagedFileAccessError( 'filesystem_unavailable', )): self.assert_access_error( 'filesystem_unavailable', lambda: traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'scan_results.jsonl', ), ) active = traversal.download_file( managed_files.RUNTIME_RESULT_ROOT_ID, 'scan_results.jsonl', ) traversal.close() self.assertEqual(b''.join(active.snapshot.chunks()), b'') def test_path_listing_and_file_limits_at_exact_boundaries(self): limits = self.limits( max_relative_path_bytes=7, max_component_bytes=4, max_path_depth=2, max_listing_entries=10, max_listing_bytes=20, max_file_bytes=4, ) self.assertEqual( managed_files.parse_managed_relative_path('abc/def', limits), ('abc', 'def'), ) self.assertEqual( managed_files.parse_managed_relative_path('abcd', limits), ('abcd',), ) for path in ('abc/defg', 'abcde', 'a/b/c'): self.assert_access_error( 'invalid_path', lambda path=path: managed_files.parse_managed_relative_path( path, limits, ), ) nested_root = os.path.join(self.temporary.name, 'nested-limits') os.mkdir(nested_root) os.mkdir(os.path.join(nested_root, 'abc')) Path(os.path.join(nested_root, 'abc', 'def')).write_bytes(b'1234') Path(os.path.join(nested_root, 'abc', 'defg')).write_bytes(b'1234') bounded = managed_files.ManagedFileRoot( 'exports', nested_root, self.root_value.permissions, limits, ) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((bounded,))) as traversal: listing = traversal.list_directory('exports', 'abc') self.assertEqual(tuple(entry.name for entry in listing.entries), ('def',)) listing_root = os.path.join(self.temporary.name, 'listing-limits') os.mkdir(listing_root) Path(os.path.join(listing_root, 'a')).write_bytes(b'a') Path(os.path.join(listing_root, 'b')).write_bytes(b'b') entry_limited = managed_files.ManagedFileRoot( 'exports', listing_root, self.root_value.permissions, self.limits(max_listing_entries=2, max_listing_bytes=10), ) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((entry_limited,))) as traversal: self.assertEqual(len(traversal.list_directory('exports').entries), 2) Path(os.path.join(listing_root, 'c')).write_bytes(b'c') self.assert_access_error( 'limit_exceeded', lambda: traversal.list_directory('exports'), ) os.unlink(os.path.join(listing_root, 'c')) name_limited = managed_files.ManagedFileRoot( 'exports', listing_root, self.root_value.permissions, self.limits(max_listing_entries=10, max_listing_bytes=2), ) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((name_limited,))) as traversal: self.assertEqual(traversal.list_directory('exports').name_bytes, 2) Path(os.path.join(listing_root, 'cc')).write_bytes(b'c') self.assert_access_error( 'limit_exceeded', lambda: traversal.list_directory('exports'), ) file_root = os.path.join(self.temporary.name, 'file-limits') os.mkdir(file_root) Path(os.path.join(file_root, 'exact')).write_bytes(b'1234') Path(os.path.join(file_root, 'over')).write_bytes(b'12345') writable = managed_files.ManagedFileRoot( 'exports', file_root, managed_files.ManagedFilePermissions(True, True, True, True), self.limits(max_file_bytes=4), ) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal: self.assertEqual( traversal.download_file('exports', 'exact').content, b'1234', ) self.assert_access_error( 'limit_exceeded', lambda: traversal.download_file('exports', 'over'), ) traversal.create_replace_file( 'exports', 'created', b'abcd', expected_sha256=None, ) self.assert_access_error( 'limit_exceeded', lambda: traversal.create_replace_file( 'exports', 'too-large', b'abcde', expected_sha256=None, ), ) self.assertEqual(Path(os.path.join(file_root, 'created')).read_bytes(), b'abcd') self.assertFalse(os.path.lexists(os.path.join(file_root, 'too-large'))) def test_create_replace_download_and_delete_are_hash_checked_and_durable(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) old = b'old payload' new = b'new payload with a different length' old_hash = hashlib.sha256(old).hexdigest() new_hash = hashlib.sha256(new).hexdigest() target = os.path.join(self.root_path, 'artifact.bin') with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal: created = traversal.create_replace_file( 'exports', 'artifact.bin', old, expected_sha256=None, ) self.assertEqual(created, managed_files.ManagedFileMutation( None, managed_files.ManagedFileIdentity(old_hash, len(old)), True, )) details = os.stat(target, follow_symlinks=False) self.assertEqual(stat.S_IMODE(details.st_mode), 0o600) self.assertEqual(details.st_uid, os.geteuid()) self.assertEqual(details.st_nlink, 1) self.assertEqual(traversal.download_file( 'exports', 'artifact.bin', ).content, old) self.assert_access_error( 'hash_conflict', lambda: traversal.create_replace_file( 'exports', 'artifact.bin', b'collision', expected_sha256=None, ), ) self.assert_access_error( 'hash_conflict', lambda: traversal.create_replace_file( 'exports', 'artifact.bin', new, expected_sha256='0' * 64, ), ) self.assertEqual(Path(target).read_bytes(), old) unchanged = traversal.create_replace_file( 'exports', 'artifact.bin', old, expected_sha256=old_hash, ) self.assertEqual(unchanged, managed_files.ManagedFileMutation( managed_files.ManagedFileIdentity(old_hash, len(old)), managed_files.ManagedFileIdentity(old_hash, len(old)), False, )) replaced = traversal.create_replace_file( 'exports', 'artifact.bin', new, expected_sha256=old_hash, ) self.assertEqual(replaced, managed_files.ManagedFileMutation( managed_files.ManagedFileIdentity(old_hash, len(old)), managed_files.ManagedFileIdentity(new_hash, len(new)), True, )) self.assertEqual(Path(target).read_bytes(), new) self.assert_access_error( 'hash_conflict', lambda: traversal.delete_file( 'exports', 'artifact.bin', expected_sha256=old_hash, ), ) deleted = traversal.delete_file( 'exports', 'artifact.bin', expected_sha256=new_hash, ) self.assertEqual(deleted, managed_files.ManagedFileMutation( managed_files.ManagedFileIdentity(new_hash, len(new)), None, True, )) self.assertFalse(os.path.lexists(target)) self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) def test_mutation_identity_uses_mutation_permission_and_fsyncs_observed_state(self): writable = self.root( self.root_path, allow_read=False, allow_create_replace=True, allow_delete=True, ) payload = b'mutation-only content' payload_hash = hashlib.sha256(payload).hexdigest() target = os.path.join(self.root_path, 'artifact.bin') with open(target, 'wb') as handle: handle.write(payload) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal, \ mock.patch.object( traversal, '_fsync_directory', wraps=traversal._fsync_directory, ) as fsynced: self.assert_access_error( 'operation_not_allowed', lambda: traversal.download_file('exports', 'artifact.bin'), ) self.assertEqual( traversal.mutation_file_identity( 'exports', 'artifact.bin', managed_files.ManagedFileOperation.CREATE_REPLACE, ), managed_files.ManagedFileIdentity(payload_hash, len(payload)), ) self.assertEqual(fsynced.call_count, 1) os.chmod(target, 0o644) self.assert_access_error( 'unsafe_target', lambda: traversal.mutation_file_identity( 'exports', 'artifact.bin', managed_files.ManagedFileOperation.CREATE_REPLACE, require_private_sha256=payload_hash, ), ) os.chmod(target, 0o600) self.assertEqual( traversal.mutation_file_identity( 'exports', 'artifact.bin', managed_files.ManagedFileOperation.CREATE_REPLACE, require_private_sha256=payload_hash, ).sha256, payload_hash, ) os.unlink(target) self.assert_access_error( 'not_found', lambda: traversal.mutation_file_identity( 'exports', 'artifact.bin', managed_files.ManagedFileOperation.DELETE, ), ) self.assertEqual(fsynced.call_count, 4) self.assert_access_error( 'operation_not_allowed', lambda: traversal.mutation_file_identity( 'exports', 'artifact.bin', managed_files.ManagedFileOperation.READ, ), ) def test_failed_mutation_drops_payload_from_managed_file_tracebacks(self): writable = self.root(self.root_path, allow_create_replace=True) payload = b'traceback-managed-files-payload-sentinel' try: with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal, \ mock.patch.object( managed_files.os, 'write', side_effect=OSError('fixture write failure'), ): traversal.create_replace_file( 'exports', 'failed.bin', payload, expected_sha256=None, ) except managed_files.ManagedFileAccessError as error: current = error.__traceback__ found_managed_frame = False while current is not None: if Path(current.tb_frame.f_code.co_filename).name == 'managed_files.py': found_managed_frame = True self.assertNotIn( payload.decode('ascii'), repr(current.tb_frame.f_locals), ) current = current.tb_next self.assertTrue(found_managed_frame) else: self.fail('managed file write failure was not raised') def test_mutation_limits_permissions_and_unsafe_targets_fail_closed(self): readonly = self.root(self.root_path) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((readonly,))) as traversal: self.assert_access_error( 'operation_not_allowed', lambda: traversal.create_replace_file( 'exports', 'file', b'value', expected_sha256=None, ), ) self.assert_access_error( 'operation_not_allowed', lambda: traversal.delete_file( 'exports', 'file', expected_sha256='0' * 64, ), ) writable = managed_files.ManagedFileRoot( root_id='exports', absolute_path=self.root_path, permissions=managed_files.ManagedFilePermissions(True, True, True, True), limits=self.limits(max_file_bytes=4), ) with open(os.path.join(self.root_path, 'target'), 'wb') as handle: handle.write(b'old') os.link( os.path.join(self.root_path, 'target'), os.path.join(self.root_path, 'hardlink'), ) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal: self.assert_access_error( 'limit_exceeded', lambda: traversal.create_replace_file( 'exports', 'new', b'12345', expected_sha256=None, ), ) for invalid_hash in (None, '', 'A' * 64, '0' * 63, 1): if invalid_hash is None: continue self.assert_access_error( 'invalid_hash', lambda value=invalid_hash: traversal.delete_file( 'exports', 'target', expected_sha256=value, ), ) self.assert_access_error( 'invalid_content', lambda: traversal.create_replace_file( 'exports', 'new', bytearray(b'x'), expected_sha256=None, ), ) for path in ('target', 'hardlink'): self.assert_access_error( 'unsafe_target', lambda path=path: traversal.create_replace_file( 'exports', path, b'new', expected_sha256=hashlib.sha256(b'old').hexdigest(), ), ) self.assertEqual(Path(os.path.join(self.root_path, 'target')).read_bytes(), b'old') self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) def test_concurrent_compare_and_swap_has_one_winner(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) target = os.path.join(self.root_path, 'race.bin') Path(target).write_bytes(b'before') expected = hashlib.sha256(b'before').hexdigest() barrier = threading.Barrier(3) outcomes = [] lock = threading.Lock() registry = managed_files.ManagedFileRootRegistry((writable,)) with managed_files.ManagedFileTraversal(registry) as first, \ managed_files.ManagedFileTraversal(registry) as second: def replace(payload): barrier.wait() try: traversal = first if payload == b'first' else second result = traversal.create_replace_file( 'exports', 'race.bin', payload, expected_sha256=expected, ) outcome = ('success', result.after.sha256) except managed_files.ManagedFileAccessError as exc: outcome = ('error', exc.category) with lock: outcomes.append(outcome) threads = ( threading.Thread(target=replace, args=(b'first',)), threading.Thread(target=replace, args=(b'second',)), ) for thread in threads: thread.start() barrier.wait() for thread in threads: thread.join(10) self.assertFalse(thread.is_alive()) self.assertEqual(sum(item[0] == 'success' for item in outcomes), 1) self.assertEqual(sum(item == ('error', 'hash_conflict') for item in outcomes), 1) self.assertIn(Path(target).read_bytes(), (b'first', b'second')) def test_cross_instance_replace_and_delete_cannot_both_win(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) registry = managed_files.ManagedFileRootRegistry((writable,)) target = os.path.join(self.root_path, 'replace-delete.bin') Path(target).write_bytes(b'before') expected = hashlib.sha256(b'before').hexdigest() barrier = threading.Barrier(3) outcomes = [] lock = threading.Lock() with managed_files.ManagedFileTraversal(registry) as first, \ managed_files.ManagedFileTraversal(registry) as second: def mutate(action): barrier.wait() try: if action == 'replace': first.create_replace_file( 'exports', 'replace-delete.bin', b'after', expected_sha256=expected, ) else: second.delete_file( 'exports', 'replace-delete.bin', expected_sha256=expected, ) outcome = ('success', action) except managed_files.ManagedFileAccessError as exc: outcome = ('error', exc.category) with lock: outcomes.append(outcome) threads = ( threading.Thread(target=mutate, args=('replace',)), threading.Thread(target=mutate, args=('delete',)), ) for thread in threads: thread.start() barrier.wait() for thread in threads: thread.join(10) self.assertFalse(thread.is_alive()) self.assertEqual(sum(item[0] == 'success' for item in outcomes), 1) self.assertEqual(sum(item[0] == 'error' for item in outcomes), 1) if os.path.exists(target): self.assertEqual(Path(target).read_bytes(), b'after') def test_concurrent_create_delete_and_reader_atomicity(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) registry = managed_files.ManagedFileRootRegistry((writable,)) with managed_files.ManagedFileTraversal(registry) as first, \ managed_files.ManagedFileTraversal(registry) as second: create_barrier = threading.Barrier(3) create_outcomes = [] outcome_lock = threading.Lock() def create(traversal, payload): create_barrier.wait() try: traversal.create_replace_file( 'exports', 'create-race.bin', payload, expected_sha256=None, ) outcome = ('success', payload) except managed_files.ManagedFileAccessError as exc: outcome = ('error', exc.category) with outcome_lock: create_outcomes.append(outcome) create_threads = ( threading.Thread(target=create, args=(first, b'first-create')), threading.Thread(target=create, args=(second, b'second-create')), ) for thread in create_threads: thread.start() create_barrier.wait() for thread in create_threads: thread.join(10) self.assertFalse(thread.is_alive()) self.assertEqual( sum(outcome[0] == 'success' for outcome in create_outcomes), 1, ) self.assertEqual( sum(outcome == ('error', 'hash_conflict') for outcome in create_outcomes), 1, ) self.assertIn( Path(os.path.join(self.root_path, 'create-race.bin')).read_bytes(), (b'first-create', b'second-create'), ) delete_target = os.path.join(self.root_path, 'delete-race.bin') Path(delete_target).write_bytes(b'delete-me') delete_hash = hashlib.sha256(b'delete-me').hexdigest() delete_barrier = threading.Barrier(3) delete_outcomes = [] def delete(traversal): delete_barrier.wait() try: traversal.delete_file( 'exports', 'delete-race.bin', expected_sha256=delete_hash, ) outcome = ('success', None) except managed_files.ManagedFileAccessError as exc: outcome = ('error', exc.category) with outcome_lock: delete_outcomes.append(outcome) delete_threads = ( threading.Thread(target=delete, args=(first,)), threading.Thread(target=delete, args=(second,)), ) for thread in delete_threads: thread.start() delete_barrier.wait() for thread in delete_threads: thread.join(10) self.assertFalse(thread.is_alive()) self.assertEqual( sum(outcome[0] == 'success' for outcome in delete_outcomes), 1, ) self.assertEqual( sum(outcome[0] == 'error' for outcome in delete_outcomes), 1, ) self.assertFalse(os.path.lexists(delete_target)) reader_target = os.path.join(self.root_path, 'reader-race.bin') old_content = b'old-reader-content' new_content = b'new-reader-content' Path(reader_target).write_bytes(old_content) old_hash = hashlib.sha256(old_content).hexdigest() descriptor_opened = threading.Event() replacement_done = threading.Event() reader_results = [] original_read_descriptor = first._read_descriptor def blocked_read(*args, **kwargs): descriptor_opened.set() self.assertTrue(replacement_done.wait(10)) return original_read_descriptor(*args, **kwargs) def read_old_descriptor(): try: reader_results.append( first.download_file('exports', 'reader-race.bin').content, ) except BaseException as exc: reader_results.append(exc) with mock.patch.object( first, '_read_descriptor', side_effect=blocked_read): reader = threading.Thread(target=read_old_descriptor) reader.start() self.assertTrue(descriptor_opened.wait(10)) second.create_replace_file( 'exports', 'reader-race.bin', new_content, expected_sha256=old_hash, ) replacement_done.set() reader.join(10) self.assertFalse(reader.is_alive()) self.assertEqual(reader_results, [new_content]) self.assertEqual( first.download_file('exports', 'reader-race.bin').content, new_content, ) def test_download_reopens_after_in_place_write_and_atomic_replacement(self): target = os.path.join(self.root_path, 'download-race.bin') replacement = os.path.join(self.root_path, 'download-race-new.bin') old_content = b'OLD-OLD!' mixed_content = b'OLD-NEW!' new_content = b'NEW-NEW!' Path(target).write_bytes(old_content) original_times = os.stat(target, follow_symlinks=False) original_read = managed_files.os.read raced = False def mutate_then_replace(descriptor, size): nonlocal raced if not raced: raced = True Path(target).write_bytes(mixed_content) os.utime(target, ns=( original_times.st_atime_ns, original_times.st_mtime_ns, )) Path(replacement).write_bytes(new_content) os.replace(replacement, target) return original_read(descriptor, size) with managed_files.ManagedFileTraversal(self.registry) as traversal, \ mock.patch.object( managed_files.os, 'read', side_effect=mutate_then_replace, ): downloaded = traversal.download_file('exports', 'download-race.bin') self.assertTrue(raced) self.assertEqual(downloaded.content, new_content) self.assertNotIn(downloaded.content, (old_content, mixed_content)) def test_replace_cleans_temporary_file_when_target_changes_after_staging(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) target = os.path.join(self.root_path, 'changed.bin') Path(target).write_bytes(b'original') expected = hashlib.sha256(b'original').hexdigest() with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal: original_stage = traversal._stage_temporary_file def stage_and_change(*args, **kwargs): staged = original_stage(*args, **kwargs) Path(target).write_bytes(b'concurrent') return staged with mock.patch.object( traversal, '_stage_temporary_file', side_effect=stage_and_change): self.assert_access_error( 'concurrent_change', lambda: traversal.create_replace_file( 'exports', 'changed.bin', b'proposed', expected_sha256=expected, ), ) self.assertEqual(Path(target).read_bytes(), b'concurrent') self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) def test_post_publish_directory_fsync_failure_is_durability_uncertain(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) original_fsync = managed_files.os.fsync def fail_directory_fsync(descriptor): if stat.S_ISDIR(os.fstat(descriptor).st_mode): raise OSError('fixture directory fsync failure') return original_fsync(descriptor) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal, \ mock.patch.object( managed_files.os, 'fsync', side_effect=fail_directory_fsync, ): self.assert_access_error( 'durability_uncertain', lambda: traversal.create_replace_file( 'exports', 'published.bin', b'published', expected_sha256=None, ), ) self.assertEqual( Path(os.path.join(self.root_path, 'published.bin')).read_bytes(), b'published', ) self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) def test_temporary_cleanup_and_mutation_durability_errors(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) registry = managed_files.ManagedFileRootRegistry((writable,)) original_fsync = managed_files.os.fsync def fail_temporary_fsync(descriptor): if stat.S_ISREG(os.fstat(descriptor).st_mode): raise OSError('fixture temporary fsync failure') return original_fsync(descriptor) with managed_files.ManagedFileTraversal(registry) as traversal, \ mock.patch.object( managed_files.os, 'fsync', side_effect=fail_temporary_fsync, ): self.assert_access_error( 'filesystem_unavailable', lambda: traversal.create_replace_file( 'exports', 'unstaged.bin', b'payload', expected_sha256=None, ), ) self.assertFalse(os.path.lexists(os.path.join(self.root_path, 'unstaged.bin'))) self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) directory_descriptor = os.open( self.root_path, os.O_RDONLY | os.O_DIRECTORY, ) try: with mock.patch.object( managed_files.os, 'unlink', side_effect=FileNotFoundError, ), mock.patch.object( managed_files.os, 'fsync', wraps=original_fsync, ) as fsynced: managed_files.ManagedFileTraversal._cleanup_temporary_file( directory_descriptor, '.truf-managed-file-absent.tmp', ) fsynced.assert_called_once_with(directory_descriptor) with mock.patch.object( managed_files.os, 'unlink', side_effect=OSError('unlink'), ): self.assert_access_error( 'durability_uncertain', lambda: managed_files.ManagedFileTraversal._cleanup_temporary_file( directory_descriptor, '.truf-managed-file-unlink.tmp', ), ) with mock.patch.object( managed_files.os, 'unlink', side_effect=FileNotFoundError, ), mock.patch.object( managed_files.os, 'fsync', side_effect=OSError('fsync'), ): self.assert_access_error( 'durability_uncertain', lambda: managed_files.ManagedFileTraversal._cleanup_temporary_file( directory_descriptor, '.truf-managed-file-sync.tmp', ), ) finally: os.close(directory_descriptor) def fail_directory_fsync(descriptor): if stat.S_ISDIR(os.fstat(descriptor).st_mode): raise OSError('fixture directory fsync failure') return original_fsync(descriptor) replace_target = os.path.join(self.root_path, 'replace-uncertain.bin') Path(replace_target).write_bytes(b'before') before_hash = hashlib.sha256(b'before').hexdigest() with managed_files.ManagedFileTraversal(registry) as traversal, \ mock.patch.object( managed_files.os, 'fsync', side_effect=fail_directory_fsync, ): self.assert_access_error( 'durability_uncertain', lambda: traversal.create_replace_file( 'exports', 'replace-uncertain.bin', b'after', expected_sha256=before_hash, ), ) self.assertEqual(Path(replace_target).read_bytes(), b'after') delete_target = os.path.join(self.root_path, 'delete-uncertain.bin') Path(delete_target).write_bytes(b'delete') delete_hash = hashlib.sha256(b'delete').hexdigest() with managed_files.ManagedFileTraversal(registry) as traversal, \ mock.patch.object( managed_files.os, 'fsync', side_effect=fail_directory_fsync, ): self.assert_access_error( 'durability_uncertain', lambda: traversal.delete_file( 'exports', 'delete-uncertain.bin', expected_sha256=delete_hash, ), ) self.assertFalse(os.path.lexists(delete_target)) def test_post_publish_verification_failure_is_durable_but_uncertain(self): writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) directory_fsyncs = 0 original_fsync = managed_files.os.fsync def count_fsync(descriptor): nonlocal directory_fsyncs if stat.S_ISDIR(os.fstat(descriptor).st_mode): directory_fsyncs += 1 return original_fsync(descriptor) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal, \ mock.patch.object( managed_files.os, 'fsync', side_effect=count_fsync, ), mock.patch.object( traversal, '_verify_private_file', side_effect=managed_files.ManagedFileAccessError( 'filesystem_unavailable', ), ): self.assert_access_error( 'durability_uncertain', lambda: traversal.create_replace_file( 'exports', 'verified.bin', b'published', expected_sha256=None, ), ) self.assertGreaterEqual(directory_fsyncs, 1) self.assertEqual( Path(os.path.join(self.root_path, 'verified.bin')).read_bytes(), b'published', ) def test_base_exception_during_write_removes_temporary_file(self): class WriteCancelled(BaseException): pass class CloseCancelled(BaseException): pass writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) temporary_descriptor = None with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal: original_new = traversal._new_temporary_file original_close = managed_files._close_descriptor def capture_new(*args, **kwargs): nonlocal temporary_descriptor result = original_new(*args, **kwargs) temporary_descriptor = result[1] return result def cancel_close(descriptor): if descriptor == temporary_descriptor: raise CloseCancelled() return original_close(descriptor) with mock.patch.object( traversal, '_new_temporary_file', side_effect=capture_new, ), mock.patch.object( managed_files.os, 'write', side_effect=WriteCancelled(), ), mock.patch.object( managed_files, '_close_descriptor', side_effect=cancel_close, ): with self.assertRaises(WriteCancelled): traversal.create_replace_file( 'exports', 'cancelled.bin', b'payload', expected_sha256=None, ) if temporary_descriptor is not None: original_close(temporary_descriptor) self.assertFalse(os.path.lexists( os.path.join(self.root_path, 'cancelled.bin'), )) self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) def test_interrupted_parent_close_does_not_retain_mutation_lock(self): class CloseCancelled(BaseException): pass writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) parent_descriptor = None original_close = managed_files._close_descriptor registry = managed_files.ManagedFileRootRegistry((writable,)) with managed_files.ManagedFileTraversal(registry) as traversal: original_lock = traversal._lock_mutation_directory def capture_lock(descriptor): nonlocal parent_descriptor parent_descriptor = descriptor return original_lock(descriptor) def cancel_parent_close(descriptor): if descriptor == parent_descriptor: raise CloseCancelled() return original_close(descriptor) with mock.patch.object( traversal, '_lock_mutation_directory', side_effect=capture_lock, ), mock.patch.object( managed_files, '_close_descriptor', side_effect=cancel_parent_close, ): with self.assertRaises(CloseCancelled): traversal.create_replace_file( 'exports', 'unlocked.bin', b'payload', expected_sha256=None, ) self.assertIsNotNone(parent_descriptor) managed_files.fcntl.flock( parent_descriptor, managed_files.fcntl.LOCK_EX | managed_files.fcntl.LOCK_NB, ) managed_files.fcntl.flock( parent_descriptor, managed_files.fcntl.LOCK_UN, ) original_close(parent_descriptor) with managed_files.ManagedFileTraversal(registry) as traversal: current = hashlib.sha256(b'payload').hexdigest() result = traversal.create_replace_file( 'exports', 'unlocked.bin', b'next', expected_sha256=current, ) self.assertTrue(result.written) def test_fsync_cancellation_is_not_hidden_by_ordinary_mutation_error(self): class FsyncCancelled(BaseException): pass writable = self.root( self.root_path, allow_create_replace=True, allow_delete=True, ) Path(os.path.join(self.root_path, 'existing.bin')).write_bytes(b'existing') original_fsync = managed_files.os.fsync def cancel_directory_fsync(descriptor): if stat.S_ISDIR(os.fstat(descriptor).st_mode): raise FsyncCancelled() return original_fsync(descriptor) with managed_files.ManagedFileTraversal( managed_files.ManagedFileRootRegistry((writable,))) as traversal, \ mock.patch.object( managed_files.os, 'fsync', side_effect=cancel_directory_fsync, ): with self.assertRaises(FsyncCancelled): traversal.create_replace_file( 'exports', 'existing.bin', b'collision', expected_sha256=None, ) self.assertEqual( Path(os.path.join(self.root_path, 'existing.bin')).read_bytes(), b'existing', ) self.assertFalse(any( name.startswith('.truf-managed-file-') for name in os.listdir(self.root_path) )) @staticmethod def _enter(context): with context: return None if __name__ == '__main__': unittest.main()