import copy import ctypes import io import os from pathlib import Path import posixpath import queue import signal import subprocess import sys import tempfile import time from types import SimpleNamespace import unittest from unittest import mock ROOT = Path(__file__).resolve().parents[1] HELPER = Path(__file__).with_name("owned_process_helper.py") sys.path.insert(0, str(ROOT / "app")) import owned_process def identity(pid, linux=False): prefix = "proc-start-ticks:" if linux or os.name != "nt" else "windows-filetime:" return { "pid": pid, "creation_time": prefix + "4567", "creation_time_unix": 1700000045.67, "executable": "/usr/bin/python3.12" if linux else os.path.abspath(sys.executable), } class LinuxIdentityTests(unittest.TestCase): def setUp(self): fields = [b"S"] + [b"0"] * 18 + [b"4567"] self.stat = b"321 (name with ) parentheses (and nonascii \xff)) " + b" ".join(fields) self.boot_stat = b"cpu 1 2 3 4\nbtime 1700000000\n" self.platform = SimpleNamespace( name="posix", path=posixpath, readlink=mock.Mock(return_value="/usr/bin/python3.12"), sysconf=mock.Mock(return_value=100), ) for patch in ( mock.patch.object(owned_process, "os", self.platform), mock.patch("builtins.open", side_effect=self.open_proc), ): patch.start() self.addCleanup(patch.stop) def open_proc(self, path, mode): self.assertEqual(mode, "rb") if path == "/proc/321/stat": return io.BytesIO(self.stat) self.assertEqual(path, "/proc/stat") return io.BytesIO(self.boot_stat) def test_kernel_identity_handles_parentheses_and_nonascii_process_names(self): self.assertEqual(owned_process._linux_process_identity(321), identity(321, linux=True)) self.platform.readlink.assert_called_once_with("/proc/321/exe") self.platform.sysconf.assert_called_once_with("SC_CLK_TCK") def test_invalid_proc_data_fails_closed(self): valid_stat, valid_boot = self.stat, self.boot_stat for stat, boot in ( (b"garbage", valid_boot), (valid_stat.replace(b"321 (", b"322 ("), valid_boot), (valid_stat.rsplit(b" ", 1)[0] + b" -1", valid_boot), (valid_stat, b"cpu 1 2 3\n"), (valid_stat, b"btime invalid\n"), ): with self.subTest(stat=stat, boot=boot): self.stat, self.boot_stat = stat, boot with self.assertRaises(owned_process.OwnedProcessError) as error: owned_process._linux_process_identity(321) self.assertEqual(error.exception.stage, "identity") def test_exited_or_unreadable_executable_has_no_command_text_fallback(self): for error in (FileNotFoundError(), PermissionError()): with self.subTest(error=type(error).__name__): self.platform.readlink.side_effect = error with self.assertRaises(owned_process.OwnedProcessError): owned_process._linux_process_identity(321) def test_invalid_clock_rate_fails_closed(self): self.platform.sysconf.return_value = 0 with self.assertRaises(owned_process.OwnedProcessError): owned_process._linux_process_identity(321) def test_invalid_identity_fields_are_rejected(self): for key, value in ( ("pid", True), ("pid", 322), ("creation_time", ""), ("creation_time", "windows-filetime:4567"), ("creation_time", "proc-start-ticks:-1"), ("creation_time_unix", None), ("creation_time_unix", True), ("creation_time_unix", float("nan")), ("creation_time_unix", float("inf")), ("creation_time_unix", 10 ** 1000), ("executable", "python"), ("executable", "/bad\0path"), ): with self.subTest(key=key, value=value): invalid = {**identity(321, linux=True), key: value} with self.assertRaises(owned_process.OwnedProcessError): owned_process._validated_identity(invalid, 321) class StartupHandshakeTests(unittest.TestCase): def test_rejected_handshakes_always_request_and_reap_startup(self): valid = { "ok": True, "pid": 702, "job_membership_verified": True, "payload_identity": identity(702), "host_identity": identity(701), "job_members": [identity(701), identity(702)], "job_accounting": {"child_subreaper_verified": True}, } for key, value in ( ("ok", 1), ("pid", 701), ("pid", "702"), ("payload_identity", {"pid": 702}), ("host_identity", identity(703)), ("job_members", []), ("job_members", [{"pid": 702}]), ("job_membership_verified", False), ("job_accounting", None), ): with self.subTest(key=key): status = {**copy.deepcopy(valid), key: value} with mock.patch.object(owned_process.subprocess, "Popen", return_value=mock.Mock(pid=701)), \ mock.patch.object(owned_process, "_write_packet"), \ mock.patch.object(owned_process.OwnedProcess, "_wait_for_startup", return_value=status), \ mock.patch.object(owned_process.OwnedProcess, "_reap_failed_start") as reap, \ mock.patch.object(owned_process.OwnedProcess, "_request_stop", autospec=True) as stop: with self.assertRaises(owned_process.OwnedProcessError): owned_process.OwnedProcess(["fixture"]) reap.assert_called_once_with() process = stop.call_args.args[0] owned_process.OwnedProcess._close_fd(process._control_fd) process._control_fd = None def test_linux_failed_start_never_hard_kills_observer(self): process = object.__new__(owned_process.OwnedProcess) host = mock.Mock() host.wait.side_effect = [KeyboardInterrupt(), SystemExit(4), OSError("unconfirmed"), 127] process._host_process = host with mock.patch.object(owned_process, "os", SimpleNamespace(name="posix")), \ mock.patch.object(owned_process.time, "sleep"): process._reap_failed_start() self.assertEqual(host.wait.call_args_list, [mock.call()] * 4) host.terminate.assert_not_called() host.kill.assert_not_called() def test_linux_kill_uses_only_the_owned_control_channel(self): process = object.__new__(owned_process.OwnedProcess) process._host_process = mock.Mock() with mock.patch.object(process, "_request_stop") as stop, \ mock.patch.object(owned_process, "os", SimpleNamespace(name="posix")): process.kill() stop.assert_called_once_with() self.assertEqual(process._host_process.mock_calls, []) def test_forked_proxy_detaches_without_stopping_or_taking_inherited_lock(self): for method in ("terminate", "kill", "__del__"): with self.subTest(method=method): process = object.__new__(owned_process.OwnedProcess) process._owner_pid = 700 process._host_process = mock.Mock() process._control_fd = 42 process._configuration_sent = True process._control_lock = mock.MagicMock() process._control_lock.__enter__.side_effect = AssertionError("inherited lock") platform = SimpleNamespace( name="posix", getpid=lambda: 701, close=mock.Mock(), write=mock.Mock(), set_blocking=mock.Mock(), ) try: with mock.patch.object(owned_process, "os", platform): getattr(process, method)() self.assertIsNone(process._control_fd) platform.close.assert_called_once_with(42) platform.write.assert_not_called() platform.set_blocking.assert_not_called() process._control_lock.__enter__.assert_not_called() self.assertEqual(process._host_process.mock_calls, []) finally: process._control_fd = None def test_status_reader_is_the_only_owner_that_closes_its_fd(self): for response in ({"ok": True}, EOFError("fixture")): with self.subTest(response=response): read_fd, write_fd = os.pipe() result = queue.Queue(maxsize=1) owner = queue.SimpleQueue() owner.put(read_fd) try: with mock.patch.object(owned_process, "_read_packet", side_effect=[response]): owned_process._status_reader(owner, result) with self.assertRaises(OSError): os.fstat(read_fd) success, value = result.get_nowait() self.assertEqual(success, isinstance(response, dict)) self.assertIs(value, response) finally: os.close(write_fd) def test_interrupted_reader_start_never_double_closes_or_reads_a_cancelled_fd(self): for started, error in ((True, TimeoutError("fixture")), (False, KeyboardInterrupt())): with self.subTest(started=started): pending = [] descriptors = [] def thread(target, args, daemon): def start(): fd = args[0].get_nowait() args[0].put(fd) descriptors.append(fd) if started: target(*args) else: pending.append((target, args)) raise error return SimpleNamespace(start=start) def reap(): # An unread status pipe must not hold up the host's exit. with self.assertRaises(OSError): os.fstat(descriptors[0]) with mock.patch.object(owned_process.subprocess, "Popen", return_value=mock.Mock(pid=701)), \ mock.patch.object(owned_process, "_write_packet"), \ mock.patch.object(owned_process, "_read_packet", return_value={}) as read, \ mock.patch.object(owned_process.threading, "Thread", side_effect=thread), \ mock.patch.object(owned_process.OwnedProcess, "_reap_failed_start", side_effect=reap), \ mock.patch.object(owned_process.os, "close", wraps=os.close) as close: with self.assertRaises(type(error)): owned_process.OwnedProcess(["fixture"]) for target, args in pending: target(*args) self.assertEqual(read.call_count, int(started)) self.assertEqual(close.call_args_list.count(mock.call(descriptors[0])), 1) class LinuxContainmentTests(unittest.TestCase): def platform(self): return SimpleNamespace( name="posix", P_ALL=0, WEXITED=4, WNOHANG=1, WNOWAIT=0x1000000, CLD_EXITED=1, CLD_KILLED=2, CLD_DUMPED=3, waitid=mock.Mock(), getpgid=mock.Mock(return_value=701), getsid=mock.Mock(return_value=701), getpid=mock.Mock(return_value=701), ) def test_subreaper_is_enabled_verified_and_session_checked(self): for enabled, result, group in ((1, 0, 701), (0, 0, 701), (1, -1, 701), (1, 0, 700)): with self.subTest(enabled=enabled, result=result, group=group): def prctl_call(operation, argument, *unused): if operation == 37: ctypes.cast(argument, ctypes.POINTER(ctypes.c_int)).contents.value = enabled return result prctl = mock.Mock(side_effect=prctl_call) platform = self.platform() platform.getpgid.return_value = group signals = SimpleNamespace(SIGCHLD=17, SIG_DFL=0, signal=mock.Mock()) with mock.patch.object(owned_process, "os", platform), \ mock.patch.object(owned_process, "sys", SimpleNamespace(platform="linux")), \ mock.patch.object(owned_process, "signal", signals), \ mock.patch.object(owned_process.ctypes, "CDLL", return_value=SimpleNamespace(prctl=prctl)): if enabled == 1 and result == 0 and group == 701: owned_process._linux_enable_subreaper() self.assertEqual([call.args[0] for call in prctl.call_args_list], [36, 37]) signals.signal.assert_called_once_with(17, 0) else: with self.assertRaises(OSError): owned_process._linux_enable_subreaper() def test_unsupported_posix_platform_fails_before_native_call(self): with mock.patch.object(owned_process.sys, "platform", "darwin"), \ mock.patch.object(owned_process.ctypes, "CDLL") as library: with self.assertRaises(OSError): owned_process._linux_enable_subreaper() library.assert_not_called() def test_wait_observes_without_reaping_and_preserves_signal_status(self): platform = self.platform() payload = mock.Mock(pid=702) for status, expected in ( (None, None), (SimpleNamespace(si_pid=702, si_code=1, si_status=37), 37), (SimpleNamespace(si_pid=702, si_code=2, si_status=9), -9), (SimpleNamespace(si_pid=702, si_code=3, si_status=11), -11), ): with self.subTest(expected=expected), mock.patch.object(owned_process, "os", platform): platform.waitid.return_value = status self.assertEqual(owned_process._linux_payload_returncode(payload), expected) platform.waitid.assert_called_with(0, 0, 4 | 1 | 0x1000000) payload.poll.assert_not_called() payload.wait.assert_not_called() def test_live_root_reaps_adopted_zombies_without_reaping_itself(self): platform = self.platform() platform.waitpid = mock.Mock() platform.waitid.side_effect = [ SimpleNamespace(si_pid=703, si_code=1, si_status=0), SimpleNamespace(si_pid=704, si_code=2, si_status=9), SimpleNamespace(si_pid=702, si_code=1, si_status=23), ] with mock.patch.object(owned_process, "os", platform): self.assertEqual(owned_process._linux_payload_returncode(mock.Mock(pid=702)), 23) self.assertEqual(platform.waitpid.call_args_list, [mock.call(703, 0), mock.call(704, 0)]) def test_live_reaping_is_bounded_to_keep_stop_handling_responsive(self): platform = self.platform() platform.waitpid = mock.Mock() platform.waitid.return_value = SimpleNamespace(si_pid=703, si_code=1, si_status=0) with mock.patch.object(owned_process, "os", platform): self.assertIsNone(owned_process._linux_payload_returncode(mock.Mock(pid=702))) self.assertEqual(platform.waitid.call_count, owned_process._MAX_REAP_BATCH) def test_group_is_killed_before_reaping_then_every_adopted_child_is_drained(self): calls = [] payload = mock.Mock(pid=702, returncode=None) def wait(): calls.append("payload-reap") payload.returncode = 23 payload.wait.side_effect = wait with mock.patch.object( owned_process, "_terminate_posix_group", side_effect=lambda pid: calls.append("group-kill"), ) as kill, mock.patch.object( owned_process, "_linux_reap_adopted_batch", return_value=0, ), mock.patch.object( owned_process, "_linux_direct_children", side_effect=([703], [], []), ), mock.patch.object( owned_process, "_linux_kill_exact_descendant", side_effect=lambda pid: calls.append(f"child-kill-{pid}"), ), mock.patch.object( owned_process.os, "waitpid", side_effect=ChildProcessError(), ), mock.patch.object(owned_process.time, "sleep"): owned_process._linux_stop_and_reap(payload) self.assertEqual(calls, ["group-kill", "payload-reap", "child-kill-703"]) kill.assert_called_once_with(702) def test_unconfirmed_cleanup_is_not_success(self): payload = mock.Mock(pid=702, returncode=None) with mock.patch.object(owned_process, "_terminate_posix_group", side_effect=PermissionError()): with self.assertRaises(PermissionError): owned_process._linux_stop_and_reap(payload) payload.wait.assert_not_called() payload.returncode = 23 with mock.patch.object( owned_process, "_linux_reap_adopted_batch", side_effect=PermissionError(), ): with self.assertRaises(PermissionError): owned_process._linux_stop_and_reap(payload) def test_signal_callback_only_latches_request_before_launch(self): platform = SimpleNamespace(name="posix", close=mock.Mock()) signals = SimpleNamespace(SIGTERM=15, SIGINT=2, signal=mock.Mock()) event = mock.Mock() def config(*args): callback = signals.signal.call_args_list[0].args[1] callback(15, None) callback(15, None) return {} with mock.patch.object(owned_process, "os", platform), \ mock.patch.object(owned_process, "signal", signals), \ mock.patch.object(owned_process.threading, "Event", return_value=event), \ mock.patch.object(owned_process, "_open_inherited_fd", side_effect=(40, 41)), \ mock.patch.object(owned_process, "_require_isolated_host"), \ mock.patch.object(owned_process, "_linux_enable_subreaper"), \ mock.patch.object(owned_process, "_read_packet", side_effect=config), \ mock.patch.object(owned_process, "_write_packet"), \ mock.patch.object(owned_process.subprocess, "Popen") as popen: self.assertEqual(owned_process._host_main("control", "status"), 127) event.set.assert_not_called() popen.assert_not_called() def test_negative_exit_resignals_without_setting_a_sigkill_handler(self): for signum in (9, 15): with self.subTest(signum=signum): calls = [] platform = SimpleNamespace( name="posix", getpid=lambda: 701, kill=mock.Mock(side_effect=lambda *args: calls.append("kill")), _exit=mock.Mock(), ) signals = SimpleNamespace( SIGKILL=9, SIG_DFL=0, SIG_UNBLOCK=1, signal=mock.Mock(side_effect=lambda *args: calls.append("disposition")), pthread_sigmask=mock.Mock(side_effect=lambda *args: calls.append("unblock")), ) with mock.patch.object(owned_process, "os", platform), mock.patch.object(owned_process, "signal", signals): owned_process._exit_like_payload(-signum) self.assertEqual(calls, (["disposition"] if signum != 9 else []) + ["unblock", "kill"]) platform.kill.assert_called_once_with(701, signum) signals.pthread_sigmask.assert_called_once_with(1, {signum}) @unittest.skipUnless(sys.platform == "linux", "requires a real Linux kernel and /proc") class LinuxOwnedProcessIntegrationTests(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.addCleanup(self.temp.cleanup) self.directory = Path(self.temp.name) self.release = self.directory / "release" self.sentinel = subprocess.Popen( [sys.executable, "-I", "-S", "-B", str(HELPER), "linger", str(self.directory / "sentinel.pid")], stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) self.addCleanup(self.stop_raw, self.sentinel) @staticmethod def stop_raw(process): if process.poll() is None: process.kill() process.wait(timeout=10) @staticmethod def stop_owned(process): process.terminate() process.wait(timeout=10) def wait_for_tree(self, host_levels=(1, 2)): deadline = time.monotonic() + 10 markers = [self.directory / f"payload-{n}.pid" for n in range(3)] markers += [self.directory / f"host-{n}.pid" for n in host_levels] while time.monotonic() < deadline: try: return [int(marker.read_text(encoding="ascii")) for marker in markers] except (OSError, ValueError): time.sleep(0.02) self.fail("isolated nested test tree did not become ready") def launch(self, args=None): process = owned_process.OwnedProcess( [sys.executable, "-I", "-S", "-B", str(HELPER)] + (args or ["nested-owned", "2", str(self.directory), str(self.release)]), stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) self.addCleanup(self.stop_owned, process) return process def assert_tree_gone(self, pids): for pid in pids: self.assertFalse(Path(f"/proc/{pid}").exists(), f"owned descendant {pid} was not reaped") self.assertIsNone(self.sentinel.poll()) def signal_payload(self, process, signum): descriptor = os.pidfd_open(process.pid) try: self.assert_process_identity(process.payload_identity, process.pid) signal.pidfd_send_signal(descriptor, signum) finally: os.close(descriptor) def assert_process_identity(self, expected, pid): current = owned_process._linux_process_identity(pid) self.assertEqual( {key: current[key] for key in ('pid', 'creation_time', 'executable')}, {key: expected[key] for key in ('pid', 'creation_time', 'executable')}, ) self.assertGreater(current['creation_time_unix'], 0) self.assertGreater(expected['creation_time_unix'], 0) def test_live_identity_matches_kernel_and_host_session_is_isolated(self): process = self.launch() self.wait_for_tree() self.assert_process_identity(process.payload_identity, process.pid) self.assert_process_identity(process.host_identity, process.host_pid) self.assertEqual(os.getsid(process.host_pid), process.host_pid) self.assertEqual(os.getsid(process.pid), process.pid) self.assertTrue(process.ownership_snapshot()["job_accounting"]["child_subreaper_verified"]) def test_normal_exit_waits_for_transitive_cleanup(self): process = self.launch() pids = self.wait_for_tree() self.release.touch() self.assertEqual(process.wait(timeout=10), 23) self.assert_tree_gone(pids) def test_parent_pipe_eof_waits_for_transitive_cleanup(self): process = self.launch() pids = self.wait_for_tree() with process._control_lock: descriptor, process._control_fd = process._control_fd, None os.close(descriptor) self.assertEqual(process.wait(timeout=10), 1) self.assert_tree_gone(pids) def test_stop_request_works_with_an_inherited_writer_and_drains_descendants(self): process = self.launch() pids = self.wait_for_tree() duplicate_writer = os.dup(process._control_fd) try: process.kill() self.assertEqual(process.wait(timeout=10), 1) self.assert_tree_gone(pids) finally: os.close(duplicate_writer) def test_stop_kills_and_reaps_grandchild_that_escaped_with_setsid(self): process = self.launch(['escaped-session', str(self.directory)]) escaped_path = self.directory / 'escaped.pid' deadline = time.monotonic() + 10 escaped_text = '' while time.monotonic() < deadline: if escaped_path.exists(): escaped_text = escaped_path.read_text(encoding='ascii').strip() if escaped_text: break time.sleep(0.02) self.assertTrue(escaped_text, 'escaped-session grandchild did not start') escaped_pid = int(escaped_text) self.assertEqual(os.getsid(escaped_pid), escaped_pid) process.kill() self.assertEqual(process.wait(timeout=10), 1) self.assert_tree_gone([process.pid, escaped_pid]) def test_forked_proxy_finalization_leaves_parent_job_alive(self): process = self.launch(["forked-proxy", str(self.directory), str(self.release)]) pids = self.wait_for_tree(host_levels=(1,)) detached = self.directory / "fork-detached" deadline = time.monotonic() + 10 while not detached.exists() and time.monotonic() < deadline: time.sleep(0.02) self.assertTrue(detached.exists(), "forked proxy cleanup acquired an inherited lock") self.assertIsNone(process.poll()) self.assertTrue(Path(f"/proc/{int((self.directory / 'payload-0.pid').read_text())}/exe").exists()) self.release.touch() self.assertEqual(process.wait(timeout=10), 23) self.assert_tree_gone(pids) def test_adopted_observers_are_reaped_while_root_stays_alive(self): process = self.launch(["adopted-observer", str(self.directory), str(self.release)]) pids = self.wait_for_tree(host_levels=(1,)) descendants = [pid for pid in pids if pid != process.pid] (self.directory / "orphan-release").touch() deadline = time.monotonic() + 10 while time.monotonic() < deadline and any(Path(f"/proc/{pid}").exists() for pid in descendants): time.sleep(0.02) self.assert_tree_gone(descendants) self.assertIsNone(process.poll()) def test_payload_sigkill_is_minus_nine_after_transitive_cleanup(self): process = self.launch() pids = self.wait_for_tree() self.signal_payload(process, signal.SIGKILL) self.assertEqual(process.wait(timeout=10), -signal.SIGKILL) self.assert_tree_gone(pids) def test_payload_sigterm_is_preserved_after_transitive_cleanup(self): process = self.launch() pids = self.wait_for_tree() self.signal_payload(process, signal.SIGTERM) self.assertEqual(process.wait(timeout=10), -signal.SIGTERM) self.assert_tree_gone(pids) def test_failed_start_does_not_return_before_transitive_cleanup(self): original_wait = owned_process.OwnedProcess._wait_for_startup pids = [] def reject(process, timeout): original_wait(process, timeout) pids.extend(self.wait_for_tree()) raise subprocess.TimeoutExpired("deliberately rejected handshake", timeout) with mock.patch.object(owned_process.OwnedProcess, "_wait_for_startup", reject): with self.assertRaises(subprocess.TimeoutExpired): self.launch() self.assert_tree_gone(pids) if __name__ == "__main__": unittest.main()