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

1905 lines
81 KiB
Python

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), '<ManagedFileOpenedTarget>')
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()