Initial server source import
This commit is contained in:
@@ -0,0 +1,226 @@
|
||||
import ctypes
|
||||
import os
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
APP_DIR = ROOT / "app"
|
||||
HELPER = Path(__file__).with_name("owned_process_helper.py")
|
||||
sys.path.insert(0, str(APP_DIR))
|
||||
|
||||
from owned_process import OwnedProcess, OwnedProcessError, run_owned
|
||||
|
||||
|
||||
CREATE_NO_WINDOW = 0x08000000
|
||||
|
||||
|
||||
def child_flags():
|
||||
return CREATE_NO_WINDOW if os.name == "nt" else 0
|
||||
|
||||
|
||||
def wait_for_pid(path, timeout=5):
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
value = Path(path).read_text(encoding="ascii").strip()
|
||||
if value:
|
||||
return int(value)
|
||||
except (OSError, ValueError):
|
||||
pass
|
||||
time.sleep(0.02)
|
||||
raise AssertionError(f"PID marker was not created: {path}")
|
||||
|
||||
|
||||
def windows_process_running(pid):
|
||||
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
|
||||
kernel32.OpenProcess.argtypes = [ctypes.c_ulong, ctypes.c_int, ctypes.c_ulong]
|
||||
kernel32.OpenProcess.restype = ctypes.c_void_p
|
||||
kernel32.GetExitCodeProcess.argtypes = [ctypes.c_void_p, ctypes.POINTER(ctypes.c_ulong)]
|
||||
kernel32.CloseHandle.argtypes = [ctypes.c_void_p]
|
||||
handle = kernel32.OpenProcess(0x1000, False, int(pid))
|
||||
if not handle:
|
||||
return False
|
||||
try:
|
||||
exit_code = ctypes.c_ulong()
|
||||
return bool(kernel32.GetExitCodeProcess(handle, ctypes.byref(exit_code))) and exit_code.value == 259
|
||||
finally:
|
||||
kernel32.CloseHandle(handle)
|
||||
|
||||
|
||||
def wait_not_running(pid, timeout=5):
|
||||
deadline = time.monotonic() + timeout
|
||||
while time.monotonic() < deadline:
|
||||
if not windows_process_running(pid):
|
||||
return True
|
||||
time.sleep(0.02)
|
||||
return not windows_process_running(pid)
|
||||
|
||||
|
||||
class OwnedProcessTests(unittest.TestCase):
|
||||
def test_stop_does_not_wait_for_a_duplicated_control_writer(self):
|
||||
process = OwnedProcess(
|
||||
[sys.executable, str(HELPER), "partial"],
|
||||
stdin=subprocess.DEVNULL, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
)
|
||||
duplicate_writer = os.dup(process._control_fd)
|
||||
try:
|
||||
process.terminate()
|
||||
self.assertEqual(process.wait(timeout=5), 1)
|
||||
finally:
|
||||
os.close(duplicate_writer)
|
||||
process.terminate()
|
||||
process.wait(timeout=5)
|
||||
|
||||
def test_stdout_stderr_and_return_code(self):
|
||||
process = OwnedProcess(
|
||||
[sys.executable, str(HELPER), "stdio", "37", "wait-stdin"],
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
creationflags=child_flags(),
|
||||
)
|
||||
stdout, stderr = process.communicate(timeout=5)
|
||||
self.assertEqual(process.returncode, 37)
|
||||
self.assertEqual(stdout, "owned stdout\n")
|
||||
self.assertEqual(stderr, "owned stderr\n")
|
||||
self.assertTrue(process.job_membership_verified)
|
||||
snapshot = process.ownership_snapshot()
|
||||
self.assertEqual(snapshot['payload_identity']['pid'], process.pid)
|
||||
self.assertEqual(snapshot['host_identity']['pid'], process.host_pid)
|
||||
self.assertTrue(snapshot['active_members'])
|
||||
|
||||
def test_timeout_keeps_partial_output(self):
|
||||
with self.assertRaises(subprocess.TimeoutExpired) as raised:
|
||||
run_owned(
|
||||
[sys.executable, str(HELPER), "partial"],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=0.5,
|
||||
creationflags=child_flags(),
|
||||
)
|
||||
self.assertIn("partial stdout", raised.exception.stdout)
|
||||
self.assertIn("partial stderr", raised.exception.stderr)
|
||||
|
||||
|
||||
@unittest.skipUnless(os.name == "nt", "Windows Job Object behavior")
|
||||
class WindowsOwnedProcessTests(unittest.TestCase):
|
||||
def start_sentinel(self, ready_path):
|
||||
process = subprocess.Popen(
|
||||
[sys.executable, str(HELPER), "linger", str(ready_path)],
|
||||
stdin=subprocess.DEVNULL,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
close_fds=True,
|
||||
creationflags=CREATE_NO_WINDOW,
|
||||
)
|
||||
wait_for_pid(ready_path)
|
||||
self.addCleanup(self.stop_sentinel, process)
|
||||
return process
|
||||
|
||||
@staticmethod
|
||||
def stop_sentinel(process):
|
||||
if process.poll() is not None:
|
||||
return
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=3)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait(timeout=3)
|
||||
|
||||
def test_normal_root_exit_removes_descendant_and_preserves_sentinel(self):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
descendant_marker = Path(temp_dir) / "descendant.pid"
|
||||
sentinel_marker = Path(temp_dir) / "sentinel.pid"
|
||||
sentinel = self.start_sentinel(sentinel_marker)
|
||||
process = OwnedProcess(
|
||||
[sys.executable, str(HELPER), "spawn-and-exit", str(descendant_marker), "23"],
|
||||
stdin=subprocess.DEVNULL,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
creationflags=CREATE_NO_WINDOW,
|
||||
)
|
||||
descendant_pid = wait_for_pid(descendant_marker)
|
||||
self.assertEqual(process.wait(timeout=5), 23)
|
||||
self.assertTrue(wait_not_running(descendant_pid), f"descendant {descendant_pid} survived")
|
||||
self.assertIsNone(sentinel.poll())
|
||||
|
||||
def test_forced_stop_removes_root_and_descendant_and_preserves_sentinel(self):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
descendant_marker = Path(temp_dir) / "descendant.pid"
|
||||
sentinel_marker = Path(temp_dir) / "sentinel.pid"
|
||||
sentinel = self.start_sentinel(sentinel_marker)
|
||||
process = OwnedProcess(
|
||||
[sys.executable, str(HELPER), "tree", str(descendant_marker)],
|
||||
stdin=subprocess.DEVNULL,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
creationflags=CREATE_NO_WINDOW,
|
||||
)
|
||||
root_pid = process.pid
|
||||
descendant_pid = wait_for_pid(descendant_marker)
|
||||
process.terminate()
|
||||
process.wait(timeout=5)
|
||||
self.assertTrue(wait_not_running(root_pid), f"root {root_pid} survived")
|
||||
self.assertTrue(wait_not_running(descendant_pid), f"descendant {descendant_pid} survived")
|
||||
self.assertIsNone(sentinel.poll())
|
||||
|
||||
def test_assignment_failure_prevents_payload_launch(self):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
marker = Path(temp_dir) / "payload.marker"
|
||||
with mock.patch.dict(os.environ, {"OWNED_PROCESS_TEST_FAIL_JOB_ASSIGNMENT": "1"}):
|
||||
with self.assertRaises(OwnedProcessError) as raised:
|
||||
OwnedProcess(
|
||||
[sys.executable, str(HELPER), "marker", str(marker)],
|
||||
creationflags=CREATE_NO_WINDOW,
|
||||
)
|
||||
self.assertEqual(raised.exception.stage, "job_assignment")
|
||||
self.assertFalse(marker.exists())
|
||||
|
||||
def test_job_cpu_weight_and_memory_priority_are_verified(self):
|
||||
process = OwnedProcess(
|
||||
[sys.executable, str(HELPER), 'stdio', '0'],
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
creationflags=CREATE_NO_WINDOW,
|
||||
job_memory_limit_bytes=256 * 1024 * 1024,
|
||||
job_cpu_weight=2,
|
||||
process_memory_priority=4,
|
||||
)
|
||||
process.communicate(timeout=5)
|
||||
accounting = process.ownership_snapshot()['job_accounting']
|
||||
self.assertEqual(accounting['job_cpu_weight'], 2)
|
||||
self.assertEqual(accounting['process_memory_priority'], 4)
|
||||
self.assertEqual(accounting['job_memory_limit_bytes'], 256 * 1024 * 1024)
|
||||
|
||||
|
||||
class StaticProcessSafetyTests(unittest.TestCase):
|
||||
def test_production_sources_do_not_restore_pid_tree_killing(self):
|
||||
forbidden = (
|
||||
"taskkill",
|
||||
"Win32_Process",
|
||||
"ParentProcessId",
|
||||
"windows_descendant_pids",
|
||||
"force_kill_windows_process_tree",
|
||||
"cleanup_orphaned_trufflehog_processes",
|
||||
)
|
||||
violations = []
|
||||
for path in APP_DIR.rglob("*.py"):
|
||||
source = path.read_text(encoding="utf-8", errors="replace")
|
||||
for token in forbidden:
|
||||
if token in source:
|
||||
violations.append(f"{path.relative_to(ROOT)}: {token}")
|
||||
self.assertEqual(violations, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user