Initial server source import
This commit is contained in:
@@ -0,0 +1,507 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user