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

508 lines
24 KiB
Python

import ast
import copy
import importlib.util
import io
import math
import os
from pathlib import Path
import socket
import stat
import subprocess
import sys
import tempfile
import threading
import time
from collections import deque
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from contextlib import contextmanager
from types import ModuleType, SimpleNamespace
import unittest
from unittest import mock
APP_DIR = Path(__file__).resolve().parents[1] / 'app'
FAKE_TOKEN = '-fake-token-$HOME-$(false)-`false`-\\-%-"-\''
FAKE_USERNAME = 'fake-user-$HOME-\\-%-"-\''
def load_parts(filename, names, **bindings):
# Execute the production functions, never application imports/config/entrypoints.
path = APP_DIR / filename
nodes = []
for node in ast.parse(path.read_text(encoding='utf-8'), filename=str(path)).body:
declared = {getattr(node, 'name', None)}
if isinstance(node, ast.Assign):
declared = {target.id for target in node.targets if isinstance(target, ast.Name)}
if declared & set(names):
nodes.append(node)
module = ModuleType(path.stem)
module.__dict__.update(bindings, __file__=str(path))
exec(compile(ast.Module(body=nodes, type_ignores=[]), str(path), 'exec'), module.__dict__)
for name in names:
assert hasattr(module, name), name
return module
AUTHORITY = load_parts('lifecycle_authority.py', (
'CHILD_INSTANCE_FILE_ENV', 'CHILD_INSTANCE_ID_ENV', 'CHILD_TOKEN_ENV',
'CHILD_CONFIG_HASH_ENV', 'CHILD_SCRIPT_HASH_ENV', 'CHILD_MANIFEST_HASH_ENV',
'CHILD_DSN_HASH_ENV', 'CHILD_KIND_ENV', 'PRIVATE_CHILD_ENV_KEYS',
'_LIBPQ_PRIVATE_ENV_KEYS', '_PRIVATE_EXTERNAL_ENV_KEYS',
'strip_supervisor_credentials', 'LifecycleAuthorityError',
))
class ImmediateReader:
def __init__(self, *, target, name, daemon):
self.target = target
self.joins = []
def start(self):
self.target()
def join(self, timeout=None):
assert timeout is not None and 0 <= timeout <= 5
self.joins.append(timeout)
def is_alive(self):
return False
class IsolatedCase(unittest.TestCase):
def setUp(self):
self.root = Path(self.enterContext(tempfile.TemporaryDirectory(prefix='truf-portability-')))
self.blocked = mock.Mock(side_effect=AssertionError('live provider/database/network access'))
self.enterContext(mock.patch.object(socket, 'socket', self.blocked))
self.enterContext(mock.patch.object(socket, 'create_connection', self.blocked))
self.output = io.StringIO()
self.clock = SimpleNamespace(now=100.0)
self.clock.monotonic = lambda: self.clock.now
self.environment = {
'PATH': '', 'TEMP': str(self.root), 'TMP': str(self.root), 'TMPDIR': str(self.root),
'KEYCHECK_CANDIDATE_LEASE_SEC': '300',
'TRUF_MANAGED_POSTGRES_DSN': 'postgresql://fake:fake@127.0.0.1:1/fake',
'SCANNER_DB_URL': 'fake-database', 'PGPASSWORD': 'fake-password',
'TRUF_SUPERVISOR_TOKEN': 'fake-supervisor-token',
}
self.os = SimpleNamespace(**vars(os))
self.os.environ = self.environment
self.os.getenv = self.environment.get
def private_directory(self, path, **_options):
path = Path(path)
self.assertTrue(path.is_relative_to(self.root))
path.mkdir(parents=True, exist_ok=True, mode=0o700)
os.chmod(path, 0o700)
return str(path)
@contextmanager
def private_writer(self, path, **_options):
path = Path(path)
self.assertTrue(path.is_relative_to(self.root))
temporary = path.with_suffix('.fixture-tmp')
with temporary.open('xb') as output:
os.chmod(temporary, 0o600)
yield output
output.flush()
os.fsync(output.fileno())
os.replace(temporary, path)
def scanner(self, platform='posix'):
self.os.name = platform
self.os.chmod = mock.Mock(wraps=os.chmod)
work = self.root / 'private work'
self.private_directory(work)
process = SimpleNamespace(
job_membership_verified=True, pid=12345, payload_identity={},
poll=lambda: 0, returncode=0,
)
process_api = SimpleNamespace(**vars(subprocess))
process_api.CREATE_NEW_PROCESS_GROUP = 0x200
process_api.CREATE_NO_WINDOW = 0x8000000
scanner = load_parts('scanner.py', ('prepend_client_git_environment', 'run_command_streamed'),
os=self.os, stat=stat, subprocess=process_api, math=math, time=time,
tempfile=tempfile, contextmanager=contextmanager,
_client_scan_manifest=SimpleNamespace(get=lambda: None),
strip_supervisor_credentials=AUTHORITY.strip_supervisor_credentials,
scan_config=SimpleNamespace(min_free_gb=0, trufflehog_job_memory_limit_bytes=64 * 1024 ** 2),
command_output_limits=lambda: (65536, 65536),
_scan_policy_value=lambda name, default: {
'trufflehog_job_memory_limit_bytes': 64 * 1024 ** 2,
}.get(name, default),
_raise_if_scan_slot_fatal=mock.Mock(), scoped_scan_slot_lease=lambda: (True, None),
acquire_scan_slot=self.blocked, create_command_work_dir=lambda: str(work),
_shared_staging_owners=lambda _roots: [],
harden_private_file=mock.Mock(side_effect=lambda path: os.chmod(path, 0o600)),
require_git_clone_launch_authority=mock.Mock(),
require_trufflehog_launch_authority=self.blocked,
_check_command_staging=mock.Mock(return_value=''),
OwnedProcess=mock.Mock(return_value=process), write_temp_owner=mock.Mock(),
StreamedCommandOutput=lambda *values: SimpleNamespace(returncode=values[2]),
ScanSlotFatalError=type('ScanSlotFatalError', (RuntimeError,), {}),
cleanup_command_work_dir=mock.Mock(),
)
scanner.command = [
str(self.root / 'selected-git'), 'clone', '--no-checkout', '--no-recurse-submodules',
'--', 'https://example.invalid/owner/repo.git', str(self.root / 'checkout'),
]
scanner.work = work
return scanner
def run_clone(self, scanner, token=FAKE_TOKEN):
env = dict(self.environment, TRUF_GIT_USERNAME=FAKE_USERNAME,
HTTPS_PROXY='fake-proxy', all_proxy='fake-proxy', No_Proxy='fake-exception')
if token:
env['TRUF_GIT_TOKEN'] = token
with scanner.run_command_streamed(
scanner.command, 10, env, native_git_clone=True, staging_roots=(str(self.root),),
) as output:
self.assertEqual(output.returncode, 0)
scanner.require_git_clone_launch_authority.assert_called_once_with(scanner.command)
scanner.cleanup_command_work_dir.assert_called_once_with(str(scanner.work))
self.assertEqual(scanner.OwnedProcess.call_args.args[0], scanner.command)
return scanner.OwnedProcess.call_args.kwargs['env']
def runner(self):
runner = load_parts('keycheck_runner.py', (
'SERVICES', 'SERVICE_CAPABILITIES', 'KEYCHECK_CAPACITY_BLOCKED_EXIT',
'_unconfirmed_provider_processes', 'maybe_add', 'list_value', 'service_extra_args',
'service_supports_proxy', 'service_supports_flag', '_provider_diagnostic_path',
'_record_provider_startup_failure', '_stop_failed_provider_process',
'run_service', 'run_provider_services',
), os=self.os, sys=SimpleNamespace(executable=sys.executable, stdout=self.output),
subprocess=subprocess, threading=SimpleNamespace(Thread=ImmediateReader),
math=math, copy=copy, time=self.clock, deque=deque,
FIRST_COMPLETED=FIRST_COMPLETED, ThreadPoolExecutor=ThreadPoolExecutor, wait=wait,
LifecycleAuthorityError=AUTHORITY.LifecycleAuthorityError,
require_active_supervisor_child=mock.Mock(return_value={}),
supervised_child_environment=mock.Mock(return_value={}),
strip_supervisor_credentials=AUTHORITY.strip_supervisor_credentials,
ensure_private_directory=self.private_directory, private_atomic_writer=self.private_writer,
redact_argv=lambda command: command, OwnedProcess=self.blocked, ScannerDB=self.blocked,
print=mock.Mock(),
)
runner.__file__ = str(self.root / 'keycheck_runner.py')
(self.root / 'fixture-provider.py').write_text('# Never launched by the runner.\n', encoding='ascii')
runner.SERVICES = {'github': 'fixture-provider.py', 'gcp': 'fixture-provider.py'}
runner.args = SimpleNamespace(
input=None, proxy_file=None, config='unused-fixture-config', input_mode='postgres',
max_keys=3, _provider_deadline=107.0,
)
runner.layout = {
'keycheck_dir': str(self.root / 'keychecks'), 'results_dir': str(self.root / 'results'),
'proxy_file': str(self.root / 'unused-proxy'), 'database_url': '',
}
runner.diagnostic = self.root / 'keychecks' / 'github' / '.provider-last-run.log'
return runner
def process(self, runner, *outcomes, output=b'partial provider output\n'):
stream = mock.Mock()
stream.read1.side_effect = [output, b'']
stream.read.side_effect = AssertionError('buffered read would hide partial output')
stream.close.side_effect = AssertionError('observer must not close a blocked reader')
process = SimpleNamespace(
stdout=stream, wait=mock.Mock(side_effect=outcomes),
terminate=mock.Mock(), kill=mock.Mock(),
)
runner.OwnedProcess = mock.Mock(return_value=process)
return process
def run_provider(self, runner, service='github'):
return runner.run_service(service, runner.SERVICES[service], runner.args, runner.layout)
def durable_fixtures(self, runner):
self.private_directory(runner.diagnostic.parent)
files = {
runner.diagnostic.parent / 'committed.jsonl': b'{"event":"fake-committed-result"}\n',
self.root / 'leased-candidate.json': b'{"token":"fake-fence","state":"leased"}\n',
}
for path, payload in files.items():
path.write_bytes(payload)
return files
class AskpassTests(IsolatedCase):
def test_posix_helper_is_lf_private_and_uses_filtered_environment(self):
scanner = self.scanner()
env = self.run_clone(scanner)
path = Path(env['GIT_ASKPASS'])
content = path.read_bytes()
self.assertEqual(path, scanner.work / 'git-askpass.sh')
self.assertTrue(content.startswith(b'#!/bin/sh\n'))
self.assertNotIn(b'\r', content)
self.assertNotIn(FAKE_TOKEN.encode(), content)
self.assertNotIn(FAKE_USERNAME.encode(), content)
self.assertNotIn(FAKE_TOKEN, ' '.join(scanner.command))
self.assertEqual(env['TRUF_GIT_TOKEN'], FAKE_TOKEN)
self.assertEqual(env['GIT_TERMINAL_PROMPT'], '0')
self.assertEqual(env['NO_PROXY'], '*')
for name in ('HTTPS_PROXY', 'all_proxy', 'No_Proxy', 'SCANNER_DB_URL',
'PGPASSWORD', 'TRUF_MANAGED_POSTGRES_DSN', 'TRUF_SUPERVISOR_TOKEN'):
self.assertNotIn(name, env)
scanner.harden_private_file.assert_called_once_with(str(path))
scanner.os.chmod.assert_called_once_with(str(path), 0o700)
if os.name == 'posix':
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o700)
self.assertEqual(stat.S_IMODE(path.parent.stat().st_mode), 0o700)
def test_windows_batch_branch_is_preserved(self):
scanner = self.scanner('nt')
env = self.run_clone(scanner)
path = Path(env['GIT_ASKPASS'])
self.assertEqual(path.suffix, '.cmd')
content = path.read_text(encoding='ascii')
self.assertIn('@echo off', content)
self.assertIn('findstr /I "username"', content)
self.assertIn('(echo %TRUF_GIT_USERNAME%) else (echo %TRUF_GIT_TOKEN%)', content)
self.assertNotIn(FAKE_TOKEN, content)
scanner.harden_private_file.assert_called_once_with(str(path))
scanner.os.chmod.assert_not_called()
def test_anonymous_clone_does_not_create_an_askpass_file(self):
scanner = self.scanner()
env = self.run_clone(scanner, token=None)
self.assertEqual(env['GIT_ASKPASS'], 'true')
self.assertFalse(list(scanner.work.glob('git-askpass*')))
scanner.harden_private_file.assert_not_called()
def test_posix_helper_refuses_to_overwrite_an_existing_path(self):
scanner = self.scanner()
path = scanner.work / 'git-askpass.sh'
path.write_bytes(b'untouched fixture\n')
with self.assertRaises(FileExistsError):
self.run_clone(scanner)
self.assertEqual(path.read_bytes(), b'untouched fixture\n')
scanner.OwnedProcess.assert_not_called()
@unittest.skipUnless(os.name == 'posix', 'requires a native POSIX shell and permissions')
def test_posix_helper_executes_both_prompts_without_path_or_interpolation(self):
scanner = self.scanner()
env = self.run_clone(scanner)
env['PATH'] = ''
for prompt, expected in (
("Username for 'https://example.invalid': ", FAKE_USERNAME),
("uSeRnAmE for 'https://example.invalid': ", FAKE_USERNAME),
("Password for 'https://example.invalid': ", FAKE_TOKEN),
):
with self.subTest(prompt=prompt):
result = subprocess.run(
[env['GIT_ASKPASS'], prompt], env=env, cwd=self.root,
stdin=subprocess.DEVNULL, capture_output=True, timeout=5, check=True,
)
self.assertEqual(result.stdout, (expected + '\n').encode())
self.assertEqual(result.stderr, b'')
class ProviderDeadlineTests(IsolatedCase):
def test_scheduler_shares_one_deadline_across_later_slices(self):
runner = self.runner()
runner.args.max_keys = 0
remaining = {'github': 2, 'gcp': 1}
seen = []
def run(service, _script, args, _layout, _extra):
seen.append((service, args.max_keys, args._provider_deadline))
remaining[service] -= 1
self.clock.now += 2
return 0
runner.run_service = run
result = runner.run_provider_services(
list(remaining), runner.args, runner.layout,
{'scheduler_workers': 1, 'scheduler_batch_keys': 2, 'scheduler_deadline_sec': 7},
work_probe=lambda service: remaining[service] > 0,
)
self.assertEqual(result, (0, []))
self.assertEqual(seen, [('github', 2, 107.0), ('gcp', 2, 107.0), ('github', 2, 107.0)])
self.assertEqual(runner.args.max_keys, 0)
self.blocked.assert_not_called()
def test_scheduler_timeout_keeps_partial_durable_results_and_leases(self):
runner = self.runner()
runner.args.max_keys = 0
files = self.durable_fixtures(runner)
process = self.process(runner, subprocess.TimeoutExpired(['fake-sensitive-command'], 7), 0)
result = runner.run_provider_services(
['github'], runner.args, runner.layout, {'scheduler_deadline_sec': 7},
work_probe=lambda _service: True,
)
self.assertEqual(result, (124, [('github', 124)]))
self.assertEqual(process.wait.call_args_list, [mock.call(timeout=7.0), mock.call(timeout=5)])
process.terminate.assert_called_once_with()
process.kill.assert_not_called()
runner.OwnedProcess.assert_called_once()
self.assertIn(b'partial provider output\n', runner.diagnostic.read_bytes())
self.assertIn(b'TimeoutExpired', runner.diagnostic.read_bytes())
self.assertNotIn(b'fake-sensitive-command', runner.diagnostic.read_bytes())
self.assertTrue(any('outcome=deadline_exceeded' in str(call) for call in runner.print.call_args_list))
self.assertEqual(runner._unconfirmed_provider_processes, [])
for path, payload in files.items():
self.assertEqual(path.read_bytes(), payload)
self.blocked.assert_not_called()
def test_startup_time_is_charged_to_the_remaining_budget(self):
runner = self.runner()
process = self.process(runner, 7)
def launch(*_args, **_kwargs):
self.clock.now += 3
return process
runner.OwnedProcess.side_effect = launch
self.assertEqual(self.run_provider(runner), 7)
process.wait.assert_called_once_with(timeout=4.0)
command = runner.OwnedProcess.call_args.args[0]
env = runner.OwnedProcess.call_args.kwargs['env']
self.assertEqual(command[1:4], ['-I', '-S', '-B'])
self.assertNotIn('--input', command)
self.assertNotIn('--max-keys', command)
self.assertNotIn(env['SCANNER_DB_URL'], ' '.join(command))
self.assertEqual(env['KEYCHECK_PROVIDER_SLICE_KEYS'], '3')
self.assertEqual(env['KEYCHECK_CANDIDATE_LEASE_SEC'], '300')
self.assertEqual(runner.diagnostic.read_bytes(), b'partial provider output\n')
def test_expired_deadline_never_launches_a_provider(self):
runner = self.runner()
runner.args._provider_deadline = self.clock.now
self.assertEqual(self.run_provider(runner), 124)
self.blocked.assert_not_called()
def test_direct_call_uses_the_existing_scheduler_default(self):
runner = self.runner()
del runner.args._provider_deadline
process = self.process(runner, 0)
self.assertEqual(self.run_provider(runner), 0)
process.wait.assert_called_once_with(timeout=1800.0)
def test_unconfirmed_stop_retains_owner_and_refuses_further_launches(self):
runner = self.runner()
files = self.durable_fixtures(runner)
process = self.process(runner, *[
subprocess.TimeoutExpired(['fake-sensitive-command'], limit) for limit in (7, 5, 5)
])
with self.assertRaisesRegex(RuntimeError, 'cleanup was not confirmed'):
self.run_provider(runner)
self.assertEqual(process.wait.call_args_list, [
mock.call(timeout=7.0), mock.call(timeout=5), mock.call(timeout=5),
])
process.terminate.assert_called_once_with()
process.kill.assert_called_once_with()
self.assertIs(runner._unconfirmed_provider_processes[0], process)
self.assertIn(b'partial provider output', runner.diagnostic.read_bytes())
process.stdout.close.assert_not_called()
with self.assertRaisesRegex(RuntimeError, 'cleanup was not confirmed'):
self.run_provider(runner, 'gcp')
runner.OwnedProcess.assert_called_once()
for path, payload in files.items():
self.assertEqual(path.read_bytes(), payload)
self.blocked.assert_not_called()
def test_stop_does_not_treat_missing_exit_status_as_confirmation(self):
runner = self.runner()
process = self.process(runner, None, None)
self.assertFalse(runner._stop_failed_provider_process(process))
self.assertIs(runner._unconfirmed_provider_processes[0], process)
def test_kill_fallback_and_wait_error_preserve_partial_diagnostics(self):
runner = self.runner()
process = self.process(
runner, OSError('fake-sensitive-error'), subprocess.TimeoutExpired(['fake'], 5), -9,
)
self.assertEqual(self.run_provider(runner), 1)
process.kill.assert_called_once_with()
self.assertEqual(runner._unconfirmed_provider_processes, [])
diagnostic = runner.diagnostic.read_bytes()
self.assertIn(b'partial provider output', diagnostic)
self.assertIn(b'OSError', diagnostic)
self.assertNotIn(b'fake-sensitive-error', diagnostic)
def test_blocked_reader_is_not_closed_by_the_waiting_thread(self):
runner = self.runner()
reader = None
def blocked_reader(**options):
nonlocal reader
reader = ImmediateReader(**options)
reader.is_alive = lambda: True
return reader
runner.threading.Thread = blocked_reader
process = self.process(runner, 0, output=b'x' * 65536 + b'final partial output\n')
self.assertEqual(self.run_provider(runner), 1)
self.assertEqual(reader.joins, [5])
process.stdout.close.assert_not_called()
diagnostic = runner.diagnostic.read_bytes()
self.assertIn(b'final partial output\n', diagnostic)
self.assertIn(b'OutputDrainTimeout', diagnostic)
self.assertLessEqual(len(diagnostic), 65536 + 100)
def test_capacity_blocked_exit_keeps_existing_scheduler_semantics(self):
runner = self.runner()
self.process(runner, runner.KEYCHECK_CAPACITY_BLOCKED_EXIT)
runner.args.max_keys = 0
result = runner.run_provider_services(
['github'], runner.args, runner.layout, {}, work_probe=lambda _service: True,
)
self.assertEqual(result, (0, []))
runner.OwnedProcess.assert_called_once()
def test_nonfinite_deadlines_are_rejected_before_launch_or_database_access(self):
runner = self.runner()
runner.args._provider_deadline = float('nan')
with self.assertRaisesRegex(ValueError, 'finite'):
self.run_provider(runner)
with self.assertRaisesRegex(ValueError, 'finite'):
runner.run_provider_services(['github'], runner.args, runner.layout, {'scheduler_deadline_sec': float('inf')})
self.blocked.assert_not_called()
@unittest.skipUnless(sys.platform == 'linux', 'requires the real Linux OwnedProcess boundary')
def test_local_owned_provider_times_out_and_keeps_fsynced_partial_work(self):
runner = self.runner()
runner.time = time
runner.threading = threading
runner.args._provider_deadline = time.monotonic() + 3
spec = importlib.util.spec_from_file_location('portability_owned_process', APP_DIR / 'owned_process.py')
owned = importlib.util.module_from_spec(spec)
spec.loader.exec_module(owned)
processes = []
payload = (
"import os,time\n"
"root=os.environ['KEYCHECK_OUTPUT_DIR']\n"
"for name,data in [('committed.jsonl',b'fake-committed\\n'),('lease.json',b'fake-leased\\n')]:\n"
" with open(os.path.join(root,name),'xb') as f:\n"
" f.write(data);f.flush();os.fsync(f.fileno())\n"
"os.write(1,b'partial local provider output\\n')\n"
"time.sleep(60)\n"
)
def launch(_command, **options):
process = owned.OwnedProcess([sys.executable, '-I', '-S', '-B', '-c', payload], **options)
processes.append(process)
return process
runner.OwnedProcess = launch
started = time.monotonic()
try:
self.assertEqual(self.run_provider(runner), 124)
self.assertLess(time.monotonic() - started, 15)
self.assertEqual(len(processes), 1)
self.assertIsNotNone(processes[0].poll())
self.assertIn(b'partial local provider output\n', runner.diagnostic.read_bytes())
self.assertEqual((runner.diagnostic.parent / 'committed.jsonl').read_bytes(), b'fake-committed\n')
self.assertEqual((runner.diagnostic.parent / 'lease.json').read_bytes(), b'fake-leased\n')
self.assertEqual(runner._unconfirmed_provider_processes, [])
finally:
for process in processes:
process.terminate()
try:
process.wait(timeout=5)
except subprocess.TimeoutExpired:
process.kill()
process.wait(timeout=5)
process.stdout.close()
self.blocked.assert_not_called()
if __name__ == '__main__':
unittest.main()