From a7e7fa7cb30b3a105abf3028113a5a2af26aaf83 Mon Sep 17 00:00:00 2001 From: yuesu Date: Thu, 17 Sep 2026 14:03:58 +0800 Subject: [PATCH] fix(e2e): isolate download observation from execution timeouts --- scripts/e2e_eval/README.md | 35 ++ scripts/e2e_eval/hf_download_observer.py | 174 ++++++ scripts/e2e_eval/run_eval.py | 342 ++---------- scripts/e2e_eval/utils/download_observer.py | 195 +++++++ scripts/e2e_eval/utils/process_tree.py | 216 ++++++++ .../test_download_observer_integration.py | 414 ++++++++++++++ tests/unit/eval/test_e2e_process_tree.py | 400 +++++++++++++ tests/unit/eval/test_hf_download_observer.py | 283 ++++++++++ tests/unit/eval/test_run_eval_script.py | 524 +++++++++++------- 9 files changed, 2084 insertions(+), 499 deletions(-) create mode 100644 scripts/e2e_eval/hf_download_observer.py create mode 100644 scripts/e2e_eval/utils/download_observer.py create mode 100644 scripts/e2e_eval/utils/process_tree.py create mode 100644 tests/unit/eval/test_download_observer_integration.py create mode 100644 tests/unit/eval/test_e2e_process_tree.py create mode 100644 tests/unit/eval/test_hf_download_observer.py diff --git a/scripts/e2e_eval/README.md b/scripts/e2e_eval/README.md index d918f8c64..6f1f961df 100644 --- a/scripts/e2e_eval/README.md +++ b/scripts/e2e_eval/README.md @@ -131,6 +131,41 @@ uv run python scripts/e2e_eval/run_eval.py --update-baseline --eval-type accurac | `--retry-failed [TYPE ...]` | — | Re-run failed jobs (implies `--continue`); unknown types are rejected as argument errors. Retry criteria are not mutually exclusive: `HF_FETCH_FAIL` also checks failed perf and accuracy logs for `WinError 10060`, `we couldn't connect to 'https://huggingface.co'`, or `thrown while requesting HEAD https://huggingface.co`, even when the primary perf classification is another type or accuracy is `FAIL`. | | `--build-only` | off | Build with `--no-compile`, writing each stage's ONNX (no EP needed). Loops the EP matrix when `--ep`/`--device` omitted | +#### Download observation and process timeouts + +Each CLI invocation has a lightweight, disposable download-observer process. +Filesystem traversal and native `psutil.open_files()` calls never run in the +supervisor. The observer first looks for changed Hugging Face partial-download +files; it does not enumerate process handles when there are no new candidates. +Only files positively associated with the CLI process tree can pause execution +time. Cache roots honor `HF_HOME`, `XDG_CACHE_HOME`, `HF_HUB_CACHE` (and its legacy +alias), `HF_DATASETS_CACHE`, and `HF_XET_CACHE`. + +The observer sends bounded, nonblocking local status messages after each scan, +normally once per second. Missing observations for five seconds, observer exit, +or startup failure cause a warning and a fallback to the ordinary execution +timeout. The observer is not restarted within that CLI invocation. Fresh +heartbeats alone do not constitute download progress: the stall deadline still +uses the partial file's size/mtime changes. + +When a successful scan finds that all previously owned partial files have gone, +the download episode has ended and the execution budget resets once. The +observer does not infer final filenames or validate cache contents; download +success remains the CLI's responsibility. A scan error stops the observer, +and its unavailability is included in the captured stderr/result diagnostics. +Unknown observation never grants a fresh budget. This fallback may time out +a slow download, but cannot disable execution-timeout enforcement. + +CLI and observer trees are owned separately: Windows uses Job Objects assigned +before process resume, and POSIX uses new process groups. Normal exit, timeout, +and interruption clean up both trees with bounded waits, including descendants +that outlive the root. On POSIX, descendants must not deliberately leave their +assigned group. CLI output is captured in temporary files rather than reader +pipes, so inherited output handles cannot prevent EOF/pipe cleanup indefinitely. +No observation is requested after the CLI exits. The outer pipeline/job deadline +remains an additional infrastructure limit; it is not substituted for the +download-aware execution timeout. + ### `run_llm_eval.py` — Run GenAI Context Sweep Runs an existing ONNX Runtime GenAI bundle through `winml perf --runtime diff --git a/scripts/e2e_eval/hf_download_observer.py b/scripts/e2e_eval/hf_download_observer.py new file mode 100644 index 000000000..73e99d73b --- /dev/null +++ b/scripts/e2e_eval/hf_download_observer.py @@ -0,0 +1,174 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Disposable HF download observer: never import WinML or the E2E runner. + +All filesystem traversal and unsafe native handle inspection stay in this +process. The supervisor consumes bounded datagrams and can kill this process +without relying on its Python interpreter, native threads, or filesystem. +""" + +from __future__ import annotations + +import argparse +import json +import os +import socket +import time +from pathlib import Path + +import psutil + + +def _expand_path(value): + return Path(os.path.expandvars(os.fspath(value))).expanduser() + + +def _hf_cache_roots(env): + home = _expand_path(env.get("HF_HOME") or ( + _expand_path(env.get("XDG_CACHE_HOME", Path.home() / ".cache")) / "huggingface" + )) + return ( + _expand_path(env.get("HF_HUB_CACHE") or env.get("HUGGINGFACE_HUB_CACHE") or home / "hub"), + _expand_path(env.get("HF_DATASETS_CACHE") or home / "datasets"), + _expand_path(env.get("HF_XET_CACHE") or home / "xet"), + ) + + +def _snapshot_hf_downloads(env): + hub, datasets, xet = _hf_cache_roots(env) + result = {} + for root, patterns in ( + (hub, ("*/blobs/*.incomplete", "*.incomplete")), + (datasets, ("downloads/*.incomplete",)), + (xet, ("**/*.incomplete",)), + ): + for pattern in patterns: + for path in root.glob(pattern): + try: + stat = path.stat() + result[path] = (stat.st_size, stat.st_mtime_ns) + except FileNotFoundError: + continue # Atomic download completion raced with the scan. + return result + + +def _normalized_path(path): + return os.path.normcase(os.path.realpath(path)) + + +def _process_tree_open_paths(pid): + root = psutil.Process(pid) + result = set() + for process in [root, *root.children(recursive=True)]: + try: + result.update(_normalized_path(item.path) for item in process.open_files()) + except psutil.NoSuchProcess: + continue + return result + + +class DownloadTracker: + """Track only changed partial files positively associated with this tree. + + Ownership is cached for a download episode, avoiding repeated native scans + while an already-known file grows. A successful cache scan with no remaining + owned partials marks the end of a download episode. The CLI, not this timing + observer, determines whether the downloaded model/data is valid. + """ + + def __init__(self, env, pid): + self.env = env + self.pid = pid + self.previous = {} + self.owned = {} + self.epoch = 0 + self.completed_epoch = 0 + self.last_progress = 0.0 + + def poll(self, now=None): + """Return a completed scan; explicit time is for deterministic tests.""" + current = _snapshot_hf_downloads(self.env) + candidates = { + path for path, value in current.items() + if path not in self.owned and self.previous.get(path) != value + } + # Do not perform native handle inspection on ordinary no-download perf. + discovered = set() + if candidates: + opened = _process_tree_open_paths(self.pid) + discovered = {path for path in candidates if _normalized_path(path) in opened} + # Native inspection may be slow. Timestamp the observed progress after + # it returns, not when the scan started (which could falsely imply stall). + if now is None: + now = time.monotonic() + was_active = bool(self.owned) + if discovered and not was_active: + self.epoch += 1 + for path in discovered: + self.owned[path] = current[path] + self.last_progress = now + + for path, previous in list(self.owned.items()): + if path not in current: + del self.owned[path] + elif previous != current[path]: + self.last_progress = now + self.owned[path] = current[path] + if was_active and not self.owned: + self.completed_epoch = self.epoch + self.previous = current + return { + "state": "ACTIVE" if self.owned else "IDLE", + "epoch": self.epoch, + "completed_epoch": self.completed_epoch, + "last_progress": self.last_progress, + } + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--pid", type=int, required=True) + parser.add_argument("--created", type=float, required=True) + parser.add_argument("--parent", type=int, required=True) + parser.add_argument("--port", type=int, required=True) + parser.add_argument("--token", required=True) + parser.add_argument("--interval", type=float, default=1.0) + args = parser.parse_args() + if args.interval <= 0: + parser.error("interval must be positive") + try: + root = psutil.Process(args.pid) + parent = psutil.Process(args.parent) + if root.create_time() != args.created: + return # Do not inspect a reused PID. + except psutil.NoSuchProcess: + return + tracker = DownloadTracker(os.environ, args.pid) + with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as channel: + channel.connect(("127.0.0.1", args.port)) + channel.setblocking(False) + sequence = 0 + while parent.is_running() and root.is_running(): + # A native hang prevents this heartbeat: do not mask it with a thread. + try: + state = tracker.poll() + except (OSError, psutil.Error): + # Failure is terminal for this optional observer. The client + # reports UNKNOWN and falls back to normal execution timing. + return + sequence += 1 + message = {**state, "token": args.token, "seq": sequence, + "observed_at": time.monotonic()} + try: + channel.send(json.dumps(message, allow_nan=False).encode()) + except BlockingIOError: + pass # The next heartbeat repeats durable completion metadata. + except OSError: + return + time.sleep(args.interval) + + +if __name__ == "__main__": + main() diff --git a/scripts/e2e_eval/run_eval.py b/scripts/e2e_eval/run_eval.py index 729f0be39..502d7de22 100644 --- a/scripts/e2e_eval/run_eval.py +++ b/scripts/e2e_eval/run_eval.py @@ -44,6 +44,7 @@ import argparse import contextlib +import errno import functools import hashlib import json @@ -56,7 +57,6 @@ import subprocess import sys import tempfile -import threading import time from dataclasses import dataclass from datetime import date, datetime, timezone @@ -68,6 +68,8 @@ from utils.classifier import FailureType, matches_hf_fetch_retry from utils.dataset_config import get_dataset_config, register_from_registry +from utils.download_observer import DownloadObserver, ExecutionBudget +from utils.process_tree import ManagedProcess from utils.recipes import RecipeVariant, copy_recipe_target, discover_recipe_variants from utils.registry import ( ModelEntry, @@ -575,319 +577,83 @@ def _sanitize_output(text: str) -> str: return "\n".join(kept) -def _kill_process_tree(pid: int) -> None: - """Kill a process and all its children. +def _run_subprocess(args: list[str], timeout: int) -> dict: + """Run an owned CLI tree with isolated, optional HF download observation. - On Windows, taskkill /T may miss grandchildren spawned without job objects. - We use psutil if available for reliable tree kill, falling back to taskkill. + Only fresh positive download observations pause execution time; only an + explicitly completed download episode resets it. Observer failures resume + normal timing. File-backed output eliminates inherited-pipe EOF/close locks. """ - try: - import psutil - except ImportError: - psutil = None - - if psutil is not None: - try: - parent = psutil.Process(pid) - except (psutil.NoSuchProcess, psutil.AccessDenied): - return - try: - children = parent.children(recursive=True) - except psutil.Error: - # The process tree may change between Process() and children(). - # Fall through to the platform tree-kill as a best effort. - pass - else: - for child in children: - with contextlib.suppress(psutil.NoSuchProcess, psutil.AccessDenied): - child.kill() - with contextlib.suppress(psutil.NoSuchProcess, psutil.AccessDenied): - parent.kill() - # Wait briefly for processes to terminate - with contextlib.suppress(psutil.Error): - psutil.wait_procs([*children, parent], timeout=5) - return - - # Fallback: taskkill on Windows, killpg on Unix - if platform.system() == "Windows": - subprocess.run( # noqa: S603 - ["taskkill", "/F", "/T", "/PID", str(pid)], # noqa: S607 - capture_output=True, - ) - else: - import signal - - try: - os.killpg(os.getpgid(pid), signal.SIGKILL) - except ProcessLookupError: - pass # Process already exited; nothing to kill - - -def _expand_cache_path(path: str | os.PathLike[str]) -> Path: - return Path(os.path.expandvars(os.fspath(path))).expanduser() - - -def _hf_cache_roots(env: dict[str, str]) -> tuple[Path, Path, Path]: - """Resolve cache roots using Hugging Face's environment precedence.""" - if "HF_HOME" in env: - hf_home = _expand_cache_path(env["HF_HOME"]) - else: - xdg_cache = _expand_cache_path(env.get("XDG_CACHE_HOME", Path.home() / ".cache")) - hf_home = xdg_cache / "huggingface" - - hub_cache_value = env.get("HF_HUB_CACHE") - if hub_cache_value is None: - hub_cache_value = env.get("HUGGINGFACE_HUB_CACHE") - hub_cache = _expand_cache_path(hub_cache_value or hf_home / "hub") - datasets_cache = _expand_cache_path(env.get("HF_DATASETS_CACHE") or hf_home / "datasets") - xet_cache = _expand_cache_path(env.get("HF_XET_CACHE") or hf_home / "xet") - return hub_cache, datasets_cache, xet_cache - - -def _snapshot_hf_downloads(env: dict[str, str]) -> dict[Path, tuple[int, int]]: - """Return observable Hugging Face partial downloads as size/mtime pairs.""" - hub_cache, datasets_cache, xet_cache = _hf_cache_roots(env) - searches = ( - (hub_cache, ("*/blobs/*.incomplete", "*.incomplete")), - (datasets_cache, ("downloads/*.incomplete",)), - (xet_cache, ("**/*.incomplete",)), - ) - snapshot: dict[Path, tuple[int, int]] = {} - for root, patterns in searches: - if not root.is_dir(): - continue - for pattern in patterns: - try: - candidates = root.glob(pattern) - for path in candidates: - try: - stat = path.stat() - except OSError: - continue - snapshot[path] = (stat.st_size, stat.st_mtime_ns) - except OSError: - continue - return snapshot - - -def _normalized_path(path: str | os.PathLike[str]) -> str: - return os.path.normcase(os.path.realpath(os.fspath(path))) - - -def _process_tree_open_paths(pid: int) -> set[str]: - """Return normalized paths opened by a process and its descendants.""" - try: - import psutil - except ImportError: - return set() - - try: - root = psutil.Process(pid) - except psutil.Error: - return set() - - processes = [root] - with contextlib.suppress(psutil.Error): - processes.extend(root.children(recursive=True)) - - paths: set[str] = set() - for process in processes: + for attempt in range(2): try: - open_files = process.open_files() - except psutil.Error: + result = _run_subprocess_once(args, timeout) + except OSError as exc: + if attempt or (exc.errno != errno.ENOSPC and getattr(exc, "winerror", None) != 112): + raise + safe_print(" [disk-full] Clearing caches after output allocation failure...") + _clear_disk_caches() continue - paths.update(_normalized_path(open_file.path) for open_file in open_files) - return paths - - -class _HfDownloadTracker: - """Detect downloads owned by the monitored subprocess tree.""" - - def __init__(self, env: dict[str, str], now: float) -> None: - self._env = env - self._previous = _snapshot_hf_downloads(env) - self._active_paths: set[Path] = set() - self._pid: int | None = None - self.last_progress = now - - def bind(self, pid: int) -> None: - self._pid = pid - - def poll(self, now: float) -> bool: - open_paths = _process_tree_open_paths(self._pid) if self._pid is not None else set() - current = _snapshot_hf_downloads(self._env) - progressed = { - path - for path, state in current.items() - if self._previous.get(path) != state and _normalized_path(path) in open_paths - } - if progressed: - self._active_paths.update(progressed) - self.last_progress = now - self._active_paths.intersection_update( - path for path in current if _normalized_path(path) in open_paths - ) - self._previous = current - return bool(self._active_paths) + if result["exit_code"] == 0 or result["timeout"] or not _is_no_space_error(result): + break + if attempt == 0: + safe_print(" [disk-full] Clearing caches and retrying once...") + _clear_disk_caches() + return result -def _run_subprocess(args: list[str], timeout: int) -> dict: - """Run a subprocess with execution and HF-download-stall timeouts. - - ``timeout`` starts normally when no Hugging Face download is observed. If a - Hub download starts, the execution budget is suspended and reset to its - full value after the download completes. Downloads get an independent - inactivity budget: if an ``*.incomplete`` cache file stops changing for - ``_HF_DOWNLOAD_STALL_TIMEOUT`` seconds, the process is terminated as an HF - fetch failure. - - Windows fix: On Windows, child processes can inherit pipe handles, causing - pipe reads to block indefinitely even after ``taskkill`` kills the process - tree. We work around this by: - 1. Using ``CREATE_NO_WINDOW`` to prevent console inheritance issues. - 2. Reading stdout/stderr in background threads. - 3. Polling process state independently of pipe EOF. - """ +def _run_subprocess_once(args: list[str], timeout: int) -> dict: env = { **os.environ, "PYTHONIOENCODING": "utf-8", "HF_HUB_DOWNLOAD_TIMEOUT": str(int(_HF_DOWNLOAD_STALL_TIMEOUT)), } start = time.perf_counter() - timed_out = False - hf_download_stalled = False - execution_elapsed = 0.0 - last_poll = start - download_tracker = _HfDownloadTracker(env, start) - download_was_active = False - - popen_kwargs: dict = { - "stdout": subprocess.PIPE, - "stderr": subprocess.PIPE, - "env": env, - } - if platform.system() == "Windows": - popen_kwargs["creationflags"] = subprocess.CREATE_NO_WINDOW - else: - popen_kwargs["start_new_session"] = True - proc = subprocess.Popen(args, **popen_kwargs) # noqa: S603 - download_tracker.bind(proc.pid) - - # Read pipes in background threads so communicate() timeout works even - # when grandchild processes keep pipe handles alive (Windows issue). - stdout_chunks: list[bytes] = [] - stderr_chunks: list[bytes] = [] - - def _reader(pipe, dest: list[bytes]) -> None: - try: - while True: - chunk = pipe.read(8192) - if not chunk: - break - dest.append(chunk) - except (OSError, ValueError): - pass # Pipe closed or broken; stop reading - - stdout_thread = threading.Thread(target=_reader, args=(proc.stdout, stdout_chunks), daemon=True) - stderr_thread = threading.Thread(target=_reader, args=(proc.stderr, stderr_chunks), daemon=True) - stdout_thread.start() - stderr_thread.start() - - try: - while True: - remaining = max(0.01, timeout - execution_elapsed) - try: - proc.wait(timeout=min(_SUBPROCESS_POLL_INTERVAL, remaining)) - now = time.perf_counter() - download_active = download_tracker.poll(now) - if download_was_active and not download_active: - execution_elapsed = 0.0 - elif not download_active: - execution_elapsed += now - last_poll - exit_code = proc.returncode - break - except subprocess.TimeoutExpired: - now = time.perf_counter() - download_active = download_tracker.poll(now) - if download_active: - download_was_active = True - if now - download_tracker.last_progress >= _HF_DOWNLOAD_STALL_TIMEOUT: - hf_download_stalled = True - else: - if download_was_active: - execution_elapsed = 0.0 - download_was_active = False - else: - execution_elapsed += now - last_poll - if execution_elapsed >= timeout: - timed_out = True - last_poll = now - - if not timed_out and not hf_download_stalled: - continue - - _kill_process_tree(proc.pid) - with contextlib.suppress(OSError): - proc.kill() - exit_code = -1 - break - - # Give reader threads a moment to finish draining - stdout_thread.join(timeout=10) - stderr_thread.join(timeout=10) - except KeyboardInterrupt: - safe_print("\n [Ctrl+C] Killing subprocess...") - _kill_process_tree(proc.pid) - with contextlib.suppress(OSError): - proc.kill() - stdout_thread.join(timeout=5) - stderr_thread.join(timeout=5) - raise - finally: - # Force-close pipes to unblock any stuck reader threads - for pipe in (proc.stdout, proc.stderr): - if pipe: - try: - pipe.close() - except OSError: - pass # Pipe already closed - # Final attempt: if reader threads are still alive after pipe close, - # don't block forever — just proceed with whatever was collected. - if stdout_thread.is_alive(): - stdout_thread.join(timeout=2) - if stderr_thread.is_alive(): - stderr_thread.join(timeout=2) - - stdout = b"".join(stdout_chunks).decode("utf-8", errors="replace") - stderr = b"".join(stderr_chunks).decode("utf-8", errors="replace") - if hf_download_stalled: + budget = ExecutionBudget(timeout, _HF_DOWNLOAD_STALL_TIMEOUT, time.monotonic()) + with tempfile.TemporaryFile() as out_file, tempfile.TemporaryFile() as err_file: + with ManagedProcess(args, stdout=out_file, stderr=err_file, env=env) as tree: + proc = tree.process + with DownloadObserver(proc.pid, env) as observer: + while proc.poll() is None: + now = time.monotonic() + observation = observer.poll(now) + budget.update(time.monotonic(), observation) + if budget.timed_out or budget.hf_download_stalled: + tree.terminate() + break + remaining = max(0.01, timeout - budget.execution_elapsed) + with contextlib.suppress(subprocess.TimeoutExpired): + proc.wait(timeout=min(_SUBPROCESS_POLL_INTERVAL, remaining)) + # Never perform a final observer scan after the CLI has exited. + exit_code = -1 if budget.timed_out or budget.hf_download_stalled else proc.returncode + observer_failure = observer.failure + # Tree ownership ends before capturing output or removing temp files. + out_file.seek(0) + err_file.seek(0) + stdout = out_file.read().decode("utf-8", errors="replace") + stderr = err_file.read().decode("utf-8", errors="replace") + if budget.hf_download_stalled: stderr += ( "\nError while downloading from https://huggingface.co: " f"no cache progress for {_HF_DOWNLOAD_STALL_TIMEOUT:g} seconds " "(Hugging Face download stalled).\n" ) + if observer_failure: + stderr += ( + f"\nWarning: HF download observation unavailable ({observer_failure}); " + "using the execution timeout.\n" + ) elapsed = round(time.perf_counter() - start, 1) - result = { + return { "stdout": stdout, "stderr": stderr, "exit_code": exit_code, "elapsed": elapsed, - "timeout": timed_out, - "hf_download_stalled": hf_download_stalled, + "timeout": budget.timed_out, + "hf_download_stalled": budget.hf_download_stalled, "command": " ".join(str(a) for a in args), } - # Retry once after clearing caches if the failure was due to disk full. - if exit_code != 0 and not timed_out and _is_no_space_error(result): - safe_print(" [disk-full] Detected 'no space left' — clearing caches and retrying...") - _clear_disk_caches() - safe_print(f" [disk-full] Retrying: {result['command']}") - result = _run_subprocess(args, timeout) - - return result - - # --------------------------------------------------------------------------- # Build phase # --------------------------------------------------------------------------- diff --git a/scripts/e2e_eval/utils/download_observer.py b/scripts/e2e_eval/utils/download_observer.py new file mode 100644 index 000000000..69a2702ee --- /dev/null +++ b/scripts/e2e_eval/utils/download_observer.py @@ -0,0 +1,195 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- +"""Nonblocking client and pure timeout accounting for the disposable observer.""" + +from __future__ import annotations + +import json +import logging +import math +import os +import secrets +import socket +import subprocess +import sys +import time +from dataclasses import dataclass +from pathlib import Path + +import psutil + +from .process_tree import ManagedProcess + + +logger = logging.getLogger(__name__) +FRESHNESS_SECONDS = 5.0 +OBSERVER_SCRIPT = Path(__file__).resolve().parents[1] / "hf_download_observer.py" + + +@dataclass(frozen=True) +class Observation: + """Only positive, fresh observation can pause execution accounting.""" + + state: str + observed_at: float + epoch: int = 0 + completed_epoch: int = 0 + last_progress: float = 0.0 + + +class ExecutionBudget: + """Monotonic execution/stall deadlines, independent of native observation.""" + + def __init__(self, timeout, stall_timeout, now): + self.timeout = timeout + self.stall_timeout = stall_timeout + self.execution_elapsed = 0.0 + self.last_tick = now + self.previous = Observation("UNKNOWN", now) + self.active_epoch = None + self._uncertain_episode = False + self.completed_epoch = 0 + self.timed_out = False + self.hf_download_stalled = False + + def update(self, now, observation): + """Charge unobserved time and apply each confirmed completion once.""" + paused = 0.0 + if self.previous.state == "ACTIVE": + paused = max(0.0, min(now, self.previous.observed_at + FRESHNESS_SECONDS) - self.last_tick) + self.execution_elapsed += max(0.0, now - self.last_tick - paused) + fresh = now - observation.observed_at <= FRESHNESS_SECONDS + state = observation.state if fresh else "UNKNOWN" + if state == "UNKNOWN": + if self.active_epoch is not None: + self._uncertain_episode = True + self.active_epoch = None + else: + if observation.completed_epoch > self.completed_epoch: + if self.active_epoch == observation.completed_epoch and not self._uncertain_episode: + self.execution_elapsed = 0.0 + self.completed_epoch = observation.completed_epoch + if state == "ACTIVE": + self.active_epoch = observation.epoch + else: + self.active_epoch = None + self._uncertain_episode = False + self.hf_download_stalled = ( + state == "ACTIVE" and now - observation.last_progress >= self.stall_timeout + ) + self.timed_out = state != "ACTIVE" and self.execution_elapsed >= self.timeout + self.previous = observation if fresh else Observation("UNKNOWN", now) + self.last_tick = now + + +class DownloadObserver: + """One isolated observer per CLI; failure degrades to fixed execution time. + + Datagrams bound both memory and receive time: no pipe reads, background + readers or partial-message waits. A failed observer is not restarted during + this CLI invocation, preventing stale epochs or restart loops from granting + extra execution time. The next CLI gets a new observer and session token. + """ + + def __init__(self, pid, env): + self._socket = None + self._owner = None + self._token = secrets.token_hex(16) + self._seq = 0 + self._started = time.monotonic() + self._last = Observation("UNKNOWN", self._started) + self.failure = None + try: + self._socket = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + self._socket.bind(("127.0.0.1", 0)) + self._socket.setblocking(False) + created = psutil.Process(pid).create_time() + self._owner = ManagedProcess( + [sys.executable, str(OBSERVER_SCRIPT), "--pid", str(pid), + "--created", str(created), "--parent", str(os.getpid()), + "--port", str(self._socket.getsockname()[1]), "--token", self._token], + env=env, stdin=subprocess.DEVNULL, + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) + except (OSError, psutil.Error, ValueError) as exc: + self._fail(f"could not start ({type(exc).__name__})") + + def _fail(self, reason): + if self.failure is None: + self.failure = reason + logger.warning("HF download observer unavailable: %s; using execution timeout", reason) + self.close() + + def poll(self, now): + """Consume at most eight complete datagrams, without waiting for one.""" + if self._socket is None or self._owner is None: + return Observation("UNKNOWN", now) + # A dead helper must not leave a stale ACTIVE lease or completion behind. + if self._owner.process.poll() is not None: + self._fail("observer process exited") + return Observation("UNKNOWN", now) + for _ in range(8): + try: + payload, address = self._socket.recvfrom(4097) + except BlockingIOError: + break + except OSError: + self._fail("observation channel failed") + return Observation("UNKNOWN", now) + if len(payload) > 4096 or address[0] != "127.0.0.1": + continue + try: + value = json.loads(payload) + if not isinstance(value, dict) or value.get("token") != self._token: + continue + seq = value["seq"] + if type(seq) is not int or seq <= self._seq: + continue + observation = Observation(**{key: value[key] for key in ( + "state", "observed_at", "epoch", "completed_epoch", "last_progress" + )}) + if observation.state not in {"ACTIVE", "IDLE", "UNKNOWN"}: + continue + if any(type(v) not in (int, float) or not math.isfinite(v) for v in ( + observation.observed_at, observation.last_progress + )): + continue + if not self._started <= observation.observed_at <= now: + continue + if observation.observed_at < self._last.observed_at: + continue + if not 0 <= observation.last_progress <= observation.observed_at: + continue + if any(type(v) is not int or v < 0 for v in ( + observation.epoch, observation.completed_epoch + )) or observation.completed_epoch > observation.epoch: + continue + if observation.epoch < self._last.epoch or observation.completed_epoch < self._last.completed_epoch: + continue + self._seq, self._last = seq, observation + except (ValueError, TypeError, KeyError, RecursionError): + continue + if now - self._last.observed_at > FRESHNESS_SECONDS: + self._fail("no completed observation within 5 seconds") + return Observation("UNKNOWN", now) + return self._last + + def close(self): + """Kill and reap the disposable observer using bounded process waits.""" + if self._owner is not None: + owner, self._owner = self._owner, None + try: + owner.close() + except (OSError, subprocess.TimeoutExpired) as exc: + logger.warning("HF observer cleanup failed: %s", exc) + if self._socket is not None: + self._socket.close() + self._socket = None + + def __enter__(self): + return self + + def __exit__(self, *_exc): + self.close() diff --git a/scripts/e2e_eval/utils/process_tree.py b/scripts/e2e_eval/utils/process_tree.py new file mode 100644 index 000000000..32e988ddf --- /dev/null +++ b/scripts/e2e_eval/utils/process_tree.py @@ -0,0 +1,216 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Own disposable subprocess trees without enumerating processes or open files.""" + +from __future__ import annotations + +import ctypes +import os +import signal +import subprocess +import sys +from ctypes import wintypes +from typing import TYPE_CHECKING + + +if TYPE_CHECKING: + from collections.abc import Callable, Mapping + from types import TracebackType + from typing import IO + + +__all__ = ["ManagedProcess"] + +_WINDOWS = sys.platform == "win32" +_CREATE_SUSPENDED = 0x00000004 +_CREATE_NO_WINDOW = 0x08000000 +_KILL_ON_JOB_CLOSE = 0x00002000 +_JOB_OBJECT_EXTENDED_LIMIT_INFORMATION = 9 +_WAIT_SECONDS = 2.0 + + +def _cleanup(*actions: Callable[[], None], error: BaseException | None = None) -> None: + primary = error + for action in actions: + try: + action() + except BaseException as cleanup_error: + if primary is None: + primary = cleanup_error + else: + primary.add_note(f"Process cleanup also failed: {cleanup_error!r}") + if primary is not None and error is None: + raise primary + + +class _IOCounters(ctypes.Structure): + _fields_ = [(name, ctypes.c_uint64) for name in ( + "ReadOperationCount", "WriteOperationCount", "OtherOperationCount", + "ReadTransferCount", "WriteTransferCount", "OtherTransferCount", + )] + + +class _BasicLimitInformation(ctypes.Structure): + _fields_ = [ + ("PerProcessUserTimeLimit", ctypes.c_int64), + ("PerJobUserTimeLimit", ctypes.c_int64), + ("LimitFlags", wintypes.DWORD), + ("MinimumWorkingSetSize", ctypes.c_size_t), + ("MaximumWorkingSetSize", ctypes.c_size_t), + ("ActiveProcessLimit", wintypes.DWORD), + ("Affinity", ctypes.c_size_t), + ("PriorityClass", wintypes.DWORD), + ("SchedulingClass", wintypes.DWORD), + ] + + +class _ExtendedLimitInformation(ctypes.Structure): + _fields_ = [ + ("BasicLimitInformation", _BasicLimitInformation), + ("IoInfo", _IOCounters), + ("ProcessMemoryLimit", ctypes.c_size_t), + ("JobMemoryLimit", ctypes.c_size_t), + ("PeakProcessMemoryUsed", ctypes.c_size_t), + ("PeakJobMemoryUsed", ctypes.c_size_t), + ] + + +class _WindowsJob: + """Keep the sole, non-inheritable handle to a kill-on-close job.""" + + def __init__(self) -> None: + self._kernel32 = ctypes.WinDLL("kernel32", use_last_error=True) + self._ntdll = ctypes.WinDLL("ntdll", use_last_error=True) + self._kernel32.CreateJobObjectW.argtypes = [ctypes.c_void_p, wintypes.LPCWSTR] + self._kernel32.CreateJobObjectW.restype = wintypes.HANDLE + self._kernel32.SetInformationJobObject.argtypes = [ + wintypes.HANDLE, ctypes.c_int, ctypes.c_void_p, wintypes.DWORD, + ] + self._kernel32.SetInformationJobObject.restype = wintypes.BOOL + self._kernel32.AssignProcessToJobObject.argtypes = [wintypes.HANDLE, wintypes.HANDLE] + self._kernel32.AssignProcessToJobObject.restype = wintypes.BOOL + self._kernel32.TerminateJobObject.argtypes = [wintypes.HANDLE, wintypes.UINT] + self._kernel32.TerminateJobObject.restype = wintypes.BOOL + self._kernel32.CloseHandle.argtypes = [wintypes.HANDLE] + self._kernel32.CloseHandle.restype = wintypes.BOOL + self._ntdll.NtResumeProcess.argtypes = [wintypes.HANDLE] + self._ntdll.NtResumeProcess.restype = wintypes.LONG + + self._handle: int | None = self._kernel32.CreateJobObjectW(None, None) + if not self._handle: + raise ctypes.WinError(ctypes.get_last_error()) + try: + info = _ExtendedLimitInformation() + info.BasicLimitInformation.LimitFlags = _KILL_ON_JOB_CLOSE + if not self._kernel32.SetInformationJobObject( + self._handle, _JOB_OBJECT_EXTENDED_LIMIT_INFORMATION, + ctypes.byref(info), ctypes.sizeof(info), + ): + raise ctypes.WinError(ctypes.get_last_error()) + except BaseException as error: + _cleanup(self.close, error=error) + raise + + def assign_and_resume(self, process: subprocess.Popen) -> None: + handle = wintypes.HANDLE(process._handle) + if not self._kernel32.AssignProcessToJobObject(self._handle, handle): + raise ctypes.WinError(ctypes.get_last_error()) + status = self._ntdll.NtResumeProcess(handle) + if status < 0: + raise OSError(f"NtResumeProcess failed with NTSTATUS 0x{status & 0xFFFFFFFF:08X}") + + def terminate(self) -> None: + if self._handle is not None and not self._kernel32.TerminateJobObject(self._handle, 1): + raise ctypes.WinError(ctypes.get_last_error()) + + def close(self) -> None: + if self._handle is not None: + if not self._kernel32.CloseHandle(self._handle): + raise ctypes.WinError(ctypes.get_last_error()) + self._handle = None + + +class ManagedProcess: + """Own a subprocess tree with single-threaded, release-once cleanup. + + Windows contains the suspended root in a kill-on-close job, also covering owner death. + POSIX descendants must stay in the group; owner death alone does not kill it. + Cleanup preserves primary exceptions, bounds waits, and never drains caller-owned streams. + """ + + def __init__( + self, args: list[str], *, + stdin: IO[bytes] | int | None = None, + stdout: IO[bytes] | int | None = None, + stderr: IO[bytes] | int | None = None, + env: Mapping[str, str] | None = None, + ) -> None: + self._process: subprocess.Popen | None = None + self._closed = False + self._job = _WindowsJob() if _WINDOWS else None + try: + self._process = subprocess.Popen( # noqa: S603 + args, stdin=stdin, stdout=stdout, stderr=stderr, env=env, + creationflags=_CREATE_SUSPENDED | _CREATE_NO_WINDOW if _WINDOWS else 0, + start_new_session=not _WINDOWS, + ) + if self._job is not None: + self._job.assign_and_resume(self._process) + except BaseException as error: + # Assignment failure can leave the suspended root outside the job. + if self._process is not None: + _cleanup(self._process.kill, error=error) + _cleanup(self.close, error=error) + raise + + @property + def process(self) -> subprocess.Popen: + """Return the underlying process for polling, waiting and its exit code.""" + if self._process is None: + raise RuntimeError("The subprocess has not been created") + return self._process + + def terminate(self) -> None: + """Forcibly kill the owned tree, without waiting or inspecting descendants.""" + if self._closed: + return + if self._job is not None: + self._job.terminate() + elif self._process is not None: + try: + os.killpg(self._process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + + def _reap_root(self) -> None: + if self._process is None: + return + try: + self._process.wait(timeout=_WAIT_SECONDS) + except subprocess.TimeoutExpired: + self._process.kill() + self._process.wait(timeout=_WAIT_SECONDS) + + def close(self) -> None: + """Attempt tree cleanup and root reaping once, even when either fails.""" + if self._closed: + return + actions = [self.terminate] + if self._job is not None: + actions.append(self._job.close) + try: + _cleanup(*actions, self._reap_root) + finally: + self._closed = True + + def __enter__(self) -> ManagedProcess: + return self + + def __exit__( + self, exc_type: type[BaseException] | None, + exc_value: BaseException | None, traceback: TracebackType | None, + ) -> None: + _cleanup(self.close, error=exc_value) diff --git a/tests/unit/eval/test_download_observer_integration.py b/tests/unit/eval/test_download_observer_integration.py new file mode 100644 index 000000000..449e0821b --- /dev/null +++ b/tests/unit/eval/test_download_observer_integration.py @@ -0,0 +1,414 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Real disposable observer/CLI tests, with explicit cross-process readiness. + +All child programs and file contents are generated by pytest. No network model +downloads or hardware providers are used. Native handle inspection runs only in +the production observer process, never in the pytest interpreter. +""" + +from __future__ import annotations + +import socket +import subprocess +import sys +import textwrap +import time +from concurrent.futures import Future +from contextlib import contextmanager +from queue import Queue +from threading import Thread + +import psutil +import pytest + +from tests.unit.eval.test_run_eval_script import _load_run_eval + + +@pytest.fixture(scope="module") +def run_eval(): + return _load_run_eval() + + +@pytest.fixture +def client(run_eval): + return sys.modules[run_eval.DownloadObserver.__module__] + + +@pytest.fixture +def isolated_cache(monkeypatch, tmp_path): + for name in ( + "HF_HOME", "HF_HUB_CACHE", "HUGGINGFACE_HUB_CACHE", "HF_DATASETS_CACHE", + "HF_XET_CACHE", "XDG_CACHE_HOME", + ): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("HF_HOME", str(tmp_path)) + return tmp_path + + +@pytest.fixture +def listener(): + with socket.socket() as channel: + channel.bind(("127.0.0.1", 0)) + channel.listen() + channel.settimeout(10) + yield channel + + +@contextmanager +def _connection(listener): + channel, _address = listener.accept() + with channel: + channel.settimeout(10) + assert channel.recv(1) == b"R", "Child must report readiness before receiving commands" + yield channel + + +@pytest.fixture +def owners(run_eval, client, monkeypatch): + """Record real process owners, with failure-safe cleanup of only our trees.""" + managed_process = run_eval.ManagedProcess + recorded = [] + + def spawn(*args, **kwargs): + owner = managed_process(*args, **kwargs) + recorded.append(owner) + return owner + + monkeypatch.setattr(run_eval, "ManagedProcess", spawn) + monkeypatch.setattr(client, "ManagedProcess", spawn) + try: + yield recorded + finally: + for owner in reversed(recorded): + owner.close() + + +@contextmanager +def _background_call(call, owners): + """Bound test coordination; propagate exceptions from the real runner.""" + result = Future() + + def invoke(): + try: + result.set_result(call()) + except BaseException as exc: + result.set_exception(exc) + + thread = Thread(target=invoke, daemon=True) + thread.start() + try: + yield result + finally: + # Unblock the runner on failure without racing its context-manager close. + if thread.is_alive(): + for owner in reversed(owners): + owner.terminate() + thread.join(timeout=5) + assert not thread.is_alive(), "Runner did not return after its owned trees were closed" + + +def _assert_closed(owners): + assert owners, "The test must exercise real managed processes" + for owner in owners: + assert owner._closed, "Production code must close the owner before returning" + assert owner.process.poll() is not None, "A model or helper process leaked" + + +@pytest.fixture +def observations(run_eval, client, monkeypatch): + """Observe the real client without altering protocol or scan timing.""" + pending = Queue() + history = [] + + def start(pid, env): + observer = client.DownloadObserver(pid, env) + poll = observer.poll + + def record(now): + observation = poll(now) + history.append(observation) + pending.put(observation) + return observation + + observer.poll = record + return observer + + monkeypatch.setattr(run_eval, "DownloadObserver", start) + return pending, history + + +def _await_observation(observations, predicate): + pending, _history = observations + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + observation = pending.get(timeout=max(0.01, deadline - time.monotonic())) + if predicate(observation): + return observation + pytest.fail("The helper did not report the required observation") + + +def _child_program(source): + program = textwrap.dedent(source) + compile(program, "", "exec") + return program + + +def _download_program(listener, incomplete): + """Keep a real partial file open until pytest acknowledges its observation.""" + return _child_program(f""" + import socket + from pathlib import Path + + path = Path({str(incomplete)!r}) + with socket.create_connection({listener.getsockname()!r}, timeout=10) as control: + control.sendall(b'R') + assert control.recv(1) == b'S' + path.parent.mkdir(parents=True, exist_ok=True) + stream = path.open('wb') + try: + stream.write(bytes(range(256))) + stream.flush() + control.sendall(b'W') + while True: + command = control.recv(1) + if command == b'G': + stream.write(bytes(range(256))) + stream.flush() + elif command == b'C': + stream.close() + path.replace(path.with_suffix('')) + elif command == b'X': + break + else: + raise RuntimeError('Unexpected controller command') + control.sendall(command) + finally: + stream.close() + print('download child finished', flush=True) + """) + + +@pytest.mark.parametrize("cache_variable", ["HF_HOME", "XDG_CACHE_HOME"]) +@pytest.mark.parametrize("complete", [False, True], ids=["stalled", "completed"]) +def test_owned_download_completion_and_stall( + run_eval, owners, observations, listener, isolated_cache, monkeypatch, cache_variable, complete +): + monkeypatch.delenv("HF_HOME", raising=False) + monkeypatch.setenv(cache_variable, str(isolated_cache)) + home = isolated_cache / "huggingface" if cache_variable == "XDG_CACHE_HOME" else isolated_cache + incomplete = home / "hub" / "models--generated--download" / "blobs" / "weights.incomplete" + program = _download_program(listener, incomplete) + monkeypatch.setattr(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 4) + with _background_call( + lambda: run_eval._run_subprocess([sys.executable, "-c", program], timeout=20), owners + ) as future, _connection(listener) as control: + control.sendall(b"S") + assert control.recv(1) == b"W" + active = _await_observation(observations, lambda obs: obs.state == "ACTIVE") + assert active.epoch > 0 + # Progress is generated only after positive ownership, not on a timer. + control.sendall(b"G") + assert control.recv(1) == b"G" + _await_observation(observations, lambda obs: obs.last_progress > active.last_progress) + if complete: + control.sendall(b"C") + assert control.recv(1) == b"C" + completed = _await_observation( + observations, + lambda obs: obs.state == "IDLE" and obs.completed_epoch == active.epoch, + ) + assert completed.epoch == active.epoch + control.sendall(b"X") + result = future.result(timeout=12) + assert result["timeout"] is False, result + assert "observation unavailable" not in result["stderr"], result + assert result["hf_download_stalled"] is not complete, result + assert result["exit_code"] == (0 if complete else -1), result + if complete: + assert incomplete.with_suffix("").read_bytes() == bytes(range(256)) * 2 + assert "download child finished" in result["stdout"] + else: + assert "Hugging Face download stalled" in result["stderr"] + assert len(owners) == 2, "Exactly one CLI and one observer must be spawned" + _assert_closed(owners) + + +@pytest.mark.parametrize("cache_variable", ["HF_HOME", "XDG_CACHE_HOME"]) +def test_unrelated_real_download_does_not_suspend_execution_timeout( + run_eval, owners, observations, listener, isolated_cache, monkeypatch, cache_variable +): + monkeypatch.delenv("HF_HOME", raising=False) + monkeypatch.setenv(cache_variable, str(isolated_cache)) + home = isolated_cache / "huggingface" if cache_variable == "XDG_CACHE_HOME" else isolated_cache + incomplete = home / "hub" / "models--generated--unrelated" / "blobs" / "weights.incomplete" + with run_eval.ManagedProcess( + [sys.executable, "-c", _download_program(listener, incomplete)], + stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, + ) as sibling, _connection(listener) as control: + control.sendall(b"S") + assert control.recv(1) == b"W" + with _background_call( + lambda: run_eval._run_subprocess( + [sys.executable, "-c", "from threading import Event; Event().wait()"], timeout=8 + ), owners + ) as future: + first = _await_observation(observations, lambda obs: obs.state == "IDLE") + control.sendall(b"G") + assert control.recv(1) == b"G" + _await_observation( + observations, + lambda obs: obs.state == "IDLE" and obs.observed_at > first.observed_at, + ) + result = future.result(timeout=12) + assert result["timeout"] is True, result + assert result["hf_download_stalled"] is False, result + assert "observation unavailable" not in result["stderr"], result + assert result["exit_code"] == -1 + assert all(obs.state != "ACTIVE" for obs in observations[1]) + assert sibling.process.poll() is None, "Unrelated downloader must not be killed" + _assert_closed(owners[1:]) + control.sendall(b"X") + assert sibling.process.wait(timeout=5) == 0 + assert len(owners) == 3 + _assert_closed(owners) + + +@pytest.mark.parametrize("failure", ["hang", "exit", "killed"]) +def test_real_runner_falls_back_when_external_helper_is_unavailable( + run_eval, client, owners, observations, listener, isolated_cache, monkeypatch, tmp_path, failure +): + worker = tmp_path / "unavailable_observer.py" + worker.write_text(_child_program(f""" + import os + import socket + from threading import Event + + with socket.create_connection({listener.getsockname()!r}, timeout=10) as control: + control.sendall(b'R') + assert control.recv(1) == b'G' + control.sendall(b'A') + if {failure!r} == 'exit': + os._exit(19) + Event().wait() + """), encoding="utf-8") + monkeypatch.setattr(client, "OBSERVER_SCRIPT", worker) + timeout = 8 + started = time.monotonic() + with _background_call( + lambda: run_eval._run_subprocess( + [sys.executable, "-c", "from threading import Event; Event().wait()"], timeout=timeout + ), owners + ) as future, _connection(listener) as control: + _await_observation(observations, lambda obs: obs.state == "UNKNOWN") + assert len(owners) == 2 + helper = owners[1] + assert helper.process.stdout is None, "Helper output must not create a drainable pipe" + assert helper.process.stderr is None + control.sendall(b"G") + assert control.recv(1) == b"A" + if failure == "killed": + helper.process.kill() + helper.process.wait(timeout=5) + result = future.result(timeout=timeout + 6) + assert result["exit_code"] == -1 + assert result["timeout"] is True + assert result["hf_download_stalled"] is False + assert "observation unavailable" in result["stderr"], result + expected_error = "no completed observation" if failure == "hang" else "process exited" + assert expected_error in result["stderr"] + assert timeout - 0.2 <= result["elapsed"] <= timeout + 6, result + assert time.monotonic() - started < timeout + 6 + assert len(owners) == 2, "Failed observers must not be restarted" + _assert_closed(owners) + + +@pytest.mark.parametrize("exit_code", [0, 7]) +def test_normal_exit_captures_file_backed_output_without_pipe_deadlock( + run_eval, owners, isolated_cache, exit_code +): + payload = bytes(range(256)) * 1024 + program = ( + "import sys; data = bytes(range(256)) * 1024; " + "sys.stdout.buffer.write(data); sys.stderr.buffer.write(data[::-1]); " + f"sys.exit({exit_code})" + ) + with _background_call( + lambda: run_eval._run_subprocess([sys.executable, "-c", program], timeout=15), owners + ) as future: + result = future.result(timeout=20) + assert result["exit_code"] == exit_code + assert result["stdout"] == payload.decode("utf-8", errors="replace") + assert result["stderr"] == payload[::-1].decode("utf-8", errors="replace") + assert result["timeout"] is False + assert result["hf_download_stalled"] is False + assert owners[0].process.stdout is None + assert owners[0].process.stderr is None + _assert_closed(owners) + + +def test_parent_exit_with_grandchild_holding_stdout_returns_and_cleans_tree( + run_eval, owners, listener, isolated_cache +): + grandchild = _child_program(f""" + import os + import socket + from threading import Event + + with socket.create_connection({listener.getsockname()!r}, timeout=10) as control: + control.sendall(b'R') + assert control.recv(1) == b'P' + control.sendall(str(os.getpid()).encode() + b'\\n') + assert control.recv(1) == b'G' + print('grandchild owns stdout', flush=True) + control.sendall(b'A') + Event().wait() + """) + parent = _child_program(f""" + import socket + import subprocess + import sys + + subprocess.Popen([sys.executable, '-c', {grandchild!r}]) + with socket.create_connection({listener.getsockname()!r}, timeout=10) as control: + control.sendall(b'R') + assert control.recv(1) == b'P' + control.sendall(b'parent\\n') + assert control.recv(1) == b'X' + print('parent finished', flush=True) + """) + with _background_call( + lambda: run_eval._run_subprocess([sys.executable, "-c", parent], timeout=20), owners + ) as future, _connection(listener) as first, _connection(listener) as second: + identities = [] + for control in (first, second): + control.sendall(b"P") + with control.makefile("rb") as reader: + identities.append(reader.readline().strip()) + parent_control, child_control = ( + (first, second) if identities[0] == b"parent" else (second, first) + ) + descendant = psutil.Process(int(next(value for value in identities if value != b"parent"))) + try: + child_control.sendall(b"G") + assert child_control.recv(1) == b"A" + assert descendant.is_running() + started = time.monotonic() + parent_control.sendall(b"X") + result = future.result(timeout=8) + assert time.monotonic() - started < 8 + assert result["exit_code"] == 0, result + assert result["timeout"] is False + assert "parent finished" in result["stdout"] + assert "grandchild owns stdout" in result["stdout"] + _assert_closed(owners) + descendant.wait(timeout=5) + assert not descendant.is_running() + finally: + # Preserve the failure while ensuring a broken containment cannot leak this child. + if descendant.is_running(): + descendant.kill() + descendant.wait(timeout=5) diff --git a/tests/unit/eval/test_e2e_process_tree.py b/tests/unit/eval/test_e2e_process_tree.py new file mode 100644 index 000000000..6d0a2e802 --- /dev/null +++ b/tests/unit/eval/test_e2e_process_tree.py @@ -0,0 +1,400 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Exercise real disposable process trees and injected OS failure paths.""" + +from __future__ import annotations + +import ctypes +import importlib.util +import os +import socket +import subprocess +import sys +import textwrap +from contextlib import ExitStack, nullcontext, suppress +from pathlib import Path +from types import MappingProxyType, SimpleNamespace +from unittest.mock import Mock + +import pytest + + +MODULE_PATH = ( + Path(__file__).resolve().parents[3] / "scripts" / "e2e_eval" / "utils" / "process_tree.py" +) +WINDOWS_ONLY = pytest.mark.skipif(sys.platform != "win32", reason="Requires Windows Job Objects") + + +@pytest.fixture(scope="module") +def process_tree(): + spec = importlib.util.spec_from_file_location("_e2e_process_tree", MODULE_PATH) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.fixture +def live_tree(tmp_path): + """Socket EOF proves descendant death; test teardown also releases workers.""" + with ExitStack() as stack: + listener = stack.enter_context(socket.socket()) + listener.bind(("127.0.0.1", 0)) + listener.listen() + listener.settimeout(5) + host, port = listener.getsockname() + worker = textwrap.dedent("""\ + import socket + import subprocess + import sys + + if sys.argv[3]: + subprocess.Popen([sys.executable, "-c", sys.argv[3], *sys.argv[1:3], ""]) + with socket.create_connection((sys.argv[1], int(sys.argv[2])), timeout=15) as conn: + conn.sendall(b"R") + conn.recv(1) + """) + root = textwrap.dedent(f"""\ + import os + import subprocess + import sys + import time + + worker = {worker!r} + subprocess.Popen([sys.executable, "-c", worker, {host!r}, {str(port)!r}, worker]) + if sys.argv[1] == "exit": + os._exit(len(sys.argv[1])) + time.sleep(15) + """) + compile(worker, "", "exec") + compile(root, "", "exec") + connections = [] + + def ready(): + for _ in range(2): + conn, _address = listener.accept() + stack.enter_context(conn) + conn.settimeout(5) + assert conn.recv(1) == b"R", "Descendant exited before announcing readiness" + connections.append(conn) + + def assert_stopped(): + assert len(connections) == 2 + for conn in connections: + # Windows may reset a socket when terminating its owner. + with suppress(ConnectionResetError): + assert conn.recv(1) == b"" + + yield SimpleNamespace( + args=[sys.executable, "-c", root], + kwargs={ + "stdin": stack.enter_context((tmp_path / "stdin").open("w+b")), + "stdout": stack.enter_context((tmp_path / "stdout").open("w+b")), + "stderr": stack.enter_context((tmp_path / "stderr").open("w+b")), + }, + ready=ready, + assert_stopped=assert_stopped, + ) + + +@pytest.mark.parametrize("input_size", [0, 31]) +def test_normal_output_and_exit_code_match_popen(process_tree, tmp_path, input_size): + code = ( + "import os, sys; data = sys.stdin.buffer.read(); " + "sys.stdout.buffer.write(os.environ['E2E_PROCESS_TREE_VALUE'].encode() + data[::-1]); " + "sys.stderr.buffer.write(data.hex().encode()); " + "sys.exit(len(data) % 17)" + ) + compile(code, "", "exec") + args = [sys.executable, "-c", code] + env = MappingProxyType({**os.environ, "E2E_PROCESS_TREE_VALUE": tmp_path.name}) + with ExitStack() as stack: + source = stack.enter_context((tmp_path / "input").open("w+b")) + source.write(bytes(range(input_size))) + results = [] + for name in ("baseline", "managed"): + source.seek(0) + stdout = stack.enter_context((tmp_path / f"{name}.out").open("w+b")) + stderr = stack.enter_context((tmp_path / f"{name}.err").open("w+b")) + kwargs = {"stdin": source, "stdout": stdout, "stderr": stderr, "env": env} + if name == "baseline": + with subprocess.Popen(args, **kwargs) as process: # noqa: S603 + exit_code = process.wait(timeout=5) + else: + with process_tree.ManagedProcess(args, **kwargs) as owner: + assert isinstance(owner.process, subprocess.Popen) + exit_code = owner.process.wait(timeout=5) + assert owner._closed + assert owner.process.poll() == exit_code + assert owner.process.wait(timeout=0) == exit_code + if sys.platform == "win32": + assert not owner.process._handle.closed + assert not source.closed and not stdout.closed and not stderr.closed + stdout.seek(0) + stderr.seek(0) + results.append((exit_code, stdout.read(), stderr.read())) + assert results[1] == results[0] + + +@pytest.mark.parametrize("scenario", ["terminate", "timeout", "root-exit", "exception"]) +def test_cleanup_reaches_children_and_grandchildren(process_tree, live_tree, scenario): + mode = "exit" if scenario == "root-exit" else "wait" + owner = process_tree.ManagedProcess([*live_tree.args, mode], **live_tree.kwargs) + try: + live_tree.ready() + if scenario == "root-exit": + exit_code = owner.process.wait(timeout=5) + assert exit_code == len(mode) + owner.close() + assert owner.process.returncode == exit_code + elif scenario == "terminate": + owner.terminate() + owner.terminate() + live_tree.assert_stopped() + owner.close() + elif scenario == "timeout": + with pytest.raises(subprocess.TimeoutExpired), owner: + owner.process.wait(timeout=0.01) + else: + with pytest.raises(RuntimeError, match="caller failed"), owner: + raise RuntimeError("caller failed") + live_tree.assert_stopped() + assert owner.process.returncode is not None + assert owner._closed + finally: + owner.close() + + +def test_spawn_error_is_not_hidden(process_tree, tmp_path): + with pytest.raises(FileNotFoundError): + process_tree.ManagedProcess([str(tmp_path / "missing-executable")]) + + +@pytest.mark.parametrize("final_timeout", [False, True]) +def test_root_waits_are_bounded_and_close_is_release_once(process_tree, monkeypatch, final_timeout): + process = Mock(pid=os.getpid(), returncode=None) + error = subprocess.TimeoutExpired("disposable", 2) + process.wait.side_effect = [error, error if final_timeout else 0] + monkeypatch.setattr(process_tree.subprocess, "Popen", Mock(return_value=process)) + # No real process or group may be targeted by this injected wait failure. + monkeypatch.setattr(process_tree, "_WINDOWS", True) + job = Mock() + monkeypatch.setattr(process_tree, "_WindowsJob", Mock(return_value=job)) + owner = process_tree.ManagedProcess(["disposable"]) + process_tree.subprocess.Popen.assert_called_once_with( + ["disposable"], stdin=None, stdout=None, stderr=None, env=None, + creationflags=process_tree._CREATE_SUSPENDED | process_tree._CREATE_NO_WINDOW, + start_new_session=False, + ) + job.assign_and_resume.assert_called_once_with(process) + with pytest.raises(subprocess.TimeoutExpired) if final_timeout else nullcontext(): + owner.close() + assert [call.kwargs["timeout"] for call in process.wait.call_args_list] == [2.0, 2.0] + process.kill.assert_called_once_with() + assert owner._closed + owner.close() + owner.terminate() + assert process.wait.call_count == 2 + job.terminate.assert_called_once_with() + job.close.assert_called_once_with() + + +@pytest.mark.parametrize("root_exited", [False, True]) +@pytest.mark.parametrize("group_error", [None, ProcessLookupError, PermissionError]) +def test_posix_group_cleanup_reaps_once_regardless_of_root_liveness( + process_tree, monkeypatch, root_exited, group_error, +): + monkeypatch.setattr(process_tree, "_WINDOWS", False) + process = Mock(pid=os.getpid(), returncode=0 if root_exited else None) + popen = Mock(return_value=process) + killpg = Mock(side_effect=group_error) + monkeypatch.setattr(process_tree.subprocess, "Popen", popen) + monkeypatch.setattr(process_tree.os, "killpg", killpg, raising=False) + monkeypatch.setattr(process_tree.signal, "SIGKILL", 9, raising=False) + owner = process_tree.ManagedProcess(["disposable"]) + popen.assert_called_once_with( + ["disposable"], stdin=None, stdout=None, stderr=None, env=None, + creationflags=0, start_new_session=True, + ) + with pytest.raises(PermissionError) if group_error is PermissionError else nullcontext(): + owner.close() + assert owner._closed + owner.close() + owner.terminate() + killpg.assert_called_once_with(process.pid, process_tree.signal.SIGKILL) + process.poll.assert_not_called() + process.wait.assert_called_once_with(timeout=2.0) + + +@pytest.fixture +def windows_api(monkeypatch, process_tree): + """Wrap real DLLs so injected failures still exercise native resource cleanup.""" + kernel32 = ctypes.WinDLL("kernel32", use_last_error=True) + ntdll = ctypes.WinDLL("ntdll", use_last_error=True) + monkeypatch.setattr(ctypes, "WinDLL", lambda name, **_kwargs: { + "kernel32": kernel32, "ntdll": ntdll, + }[name]) + created = [] + real_popen = subprocess.Popen + + def spawn(*args, **kwargs): + process = real_popen(*args, **kwargs) + created.append(process) + return process + + monkeypatch.setattr(process_tree.subprocess, "Popen", Mock(side_effect=spawn)) + real_close = kernel32.CloseHandle + real_close.argtypes = [ctypes.wintypes.HANDLE] + real_close.restype = ctypes.wintypes.BOOL + monkeypatch.setattr(kernel32, "CloseHandle", Mock(side_effect=real_close)) + yield SimpleNamespace(kernel32=kernel32, ntdll=ntdll, created=created) + # Independent fallback ensures a failing test never leaves suspended roots. + for process in created: + if process.poll() is None: + process.kill() + process.wait(timeout=2) + + +@WINDOWS_ONLY +@pytest.mark.parametrize("stage", ["assign", "resume"]) +@pytest.mark.parametrize("interrupt", [False, True]) +def test_windows_initialization_failure_cleans_resources( + process_tree, windows_api, monkeypatch, tmp_path, stage, interrupt, +): + api = windows_api + marker = tmp_path / "must-not-run" + code = "from pathlib import Path; import sys; Path(sys.argv[1]).touch()" + compile(code, "", "exec") + args = [sys.executable, "-c", code, str(marker)] + error_type = KeyboardInterrupt if interrupt else OSError + dll, function, result = { + "assign": (api.kernel32, "AssignProcessToJobObject", 0), + "resume": (api.ntdll, "NtResumeProcess", ctypes.c_long(0xC0000001).value), + }[stage] + failure = Mock(return_value=result, side_effect=KeyboardInterrupt if interrupt else None) + monkeypatch.setattr(dll, function, failure) + with pytest.raises(error_type): + process_tree.ManagedProcess(args, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) + assert not marker.exists() + api.kernel32.CloseHandle.assert_called_once() + process, = api.created + assert process.returncode is not None + assert not process._handle.closed + assert process.wait(timeout=0) == process.returncode + + +@WINDOWS_ONLY +@pytest.mark.parametrize("close_failure", [False, True]) +def test_windows_signatures_full_width_handles_and_single_close( + process_tree, monkeypatch, close_failure, +): + from ctypes import wintypes + + job_handle = (1 << (ctypes.sizeof(ctypes.c_void_p) * 8 - 2)) + 123 + process_handle = job_handle + 1 + kernel32 = SimpleNamespace(**{name: Mock(return_value=1) for name in ( + "SetInformationJobObject", "AssignProcessToJobObject", "TerminateJobObject", "CloseHandle", + )}) + kernel32.CreateJobObjectW = Mock(return_value=job_handle) + ntdll = SimpleNamespace(NtResumeProcess=Mock(return_value=0)) + monkeypatch.setattr(ctypes, "WinDLL", lambda name, **_kwargs: { + "kernel32": kernel32, "ntdll": ntdll, + }[name]) + events = Mock() + events.attach_mock(kernel32.AssignProcessToJobObject, "assign") + events.attach_mock(ntdll.NtResumeProcess, "resume") + job = process_tree._WindowsJob() + job.assign_and_resume(SimpleNamespace(_handle=process_handle)) + job.terminate() + kernel32.CloseHandle.return_value = not close_failure + with pytest.raises(OSError) if close_failure else nullcontext(): + job.close() + if not close_failure: + job.close() + signatures = [ + (kernel32.CreateJobObjectW, [ctypes.c_void_p, wintypes.LPCWSTR], wintypes.HANDLE), + (kernel32.SetInformationJobObject, + [wintypes.HANDLE, ctypes.c_int, ctypes.c_void_p, wintypes.DWORD], wintypes.BOOL), + (kernel32.AssignProcessToJobObject, [wintypes.HANDLE, wintypes.HANDLE], wintypes.BOOL), + (kernel32.TerminateJobObject, [wintypes.HANDLE, wintypes.UINT], wintypes.BOOL), + (kernel32.CloseHandle, [wintypes.HANDLE], wintypes.BOOL), + (ntdll.NtResumeProcess, [wintypes.HANDLE], wintypes.LONG), + ] + for function, argtypes, restype in signatures: + assert function.argtypes == argtypes + assert function.restype is restype + kernel32.CreateJobObjectW.assert_called_once_with(None, None) + assert [call[0] for call in events.mock_calls] == ["assign", "resume"] + assert kernel32.AssignProcessToJobObject.call_args.args[0] == job_handle + assert kernel32.AssignProcessToJobObject.call_args.args[1].value == process_handle + assert ntdll.NtResumeProcess.call_args.args[0].value == process_handle + kernel32.CloseHandle.assert_called_once_with(job_handle) + info_call = kernel32.SetInformationJobObject.call_args.args + info_pointer = ctypes.cast(info_call[2], ctypes.POINTER(process_tree._ExtendedLimitInformation)) + info = info_pointer.contents + assert info.BasicLimitInformation.LimitFlags == process_tree._KILL_ON_JOB_CLOSE + assert info_call[3] == ctypes.sizeof(info) + if ctypes.sizeof(ctypes.c_void_p) == 8: + assert ctypes.sizeof(info) == 144 + + +@WINDOWS_ONLY +@pytest.mark.parametrize("error_type", [None, RuntimeError, KeyboardInterrupt]) +def test_windows_terminate_failure_still_closes_job_and_reaps( + process_tree, windows_api, monkeypatch, live_tree, error_type, +): + owner = process_tree.ManagedProcess([*live_tree.args, "wait"], **live_tree.kwargs) + try: + live_tree.ready() + monkeypatch.setattr(windows_api.kernel32, "TerminateJobObject", Mock(return_value=0)) + if error_type is None: + with pytest.raises(OSError): + owner.close() + else: + primary = error_type("caller failed") + with pytest.raises(error_type) as caught, owner: + raise primary + assert caught.value is primary + assert len(primary.__notes__) == 1 + live_tree.assert_stopped() + assert owner.process.returncode is not None + assert not owner.process._handle.closed + assert owner._closed + owner.close() + owner.terminate() + windows_api.kernel32.TerminateJobObject.assert_called_once() + windows_api.kernel32.CloseHandle.assert_called_once() + finally: + owner.close() + + +@WINDOWS_ONLY +@pytest.mark.parametrize("root_mode", ["wait", "exit"]) +def test_windows_parent_abnormal_exit_kills_job(live_tree, root_mode): + # No context cleanup: the OS must release the killed parent's job handle. + code = textwrap.dedent("""\ + import importlib.util + import sys + import time + + spec = importlib.util.spec_from_file_location("process_tree", sys.argv[1]) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + owner = module.ManagedProcess(sys.argv[2:]) + time.sleep(15) + """) + compile(code, "", "exec") + args = [sys.executable, "-c", code, str(MODULE_PATH), *live_tree.args, root_mode] + with subprocess.Popen(args, **live_tree.kwargs) as parent: # noqa: S603 + try: + live_tree.ready() + parent.kill() + parent.wait(timeout=2) + live_tree.assert_stopped() + finally: + if parent.poll() is None: + parent.kill() + parent.wait(timeout=2) diff --git a/tests/unit/eval/test_hf_download_observer.py b/tests/unit/eval/test_hf_download_observer.py new file mode 100644 index 000000000..20d91fd37 --- /dev/null +++ b/tests/unit/eval/test_hf_download_observer.py @@ -0,0 +1,283 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +"""Regression tests for isolated download observation and timeout accounting.""" + +from __future__ import annotations + +import importlib.util +import json +import sys +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + + +@pytest.fixture +def observer_modules(monkeypatch): + directory = Path(__file__).resolve().parents[3] / "scripts" / "e2e_eval" + monkeypatch.syspath_prepend(str(directory)) + modules = [] + for name, path in ( + ("_hf_worker_test", directory / "hf_download_observer.py"), + ("utils._hf_client_test", directory / "utils" / "download_observer.py"), + ): + spec = importlib.util.spec_from_file_location(name, path) + module = importlib.util.module_from_spec(spec) + monkeypatch.setitem(sys.modules, name, module) + spec.loader.exec_module(module) + modules.append(module) + return modules + + +def test_no_partial_files_never_enumerates_handles(observer_modules, monkeypatch, tmp_path): + worker, _ = observer_modules + monkeypatch.setattr(worker, "_process_tree_open_paths", lambda *_: pytest.fail("handle scan")) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + assert tracker.poll(0)["state"] == "IDLE" + assert tracker.poll(1)["state"] == "IDLE" + + +def test_unknown_does_not_reset_or_indefinitely_pause_budget(observer_modules): + _, client = observer_modules + budget = client.ExecutionBudget(timeout=10, stall_timeout=600, now=0) + budget.update(4, client.Observation("ACTIVE", 4, 1, 0, 4)) + budget.update(6, client.Observation("UNKNOWN", 6)) + assert budget.execution_elapsed == pytest.approx(4) + budget.update(11, client.Observation("UNKNOWN", 11)) + assert budget.execution_elapsed == pytest.approx(9) + budget.update(12, client.Observation("IDLE", 12, 1, 1, 4)) + assert budget.timed_out + + +def test_confirmed_completion_resets_once(observer_modules): + _, client = observer_modules + budget = client.ExecutionBudget(timeout=10, stall_timeout=600, now=0) + budget.update(4, client.Observation("ACTIVE", 4, 1, 0, 4)) + budget.update(6, client.Observation("IDLE", 6, 1, 1, 5)) + assert budget.execution_elapsed == 0 + budget.update(9, client.Observation("IDLE", 9, 1, 1, 5)) + assert budget.execution_elapsed == 3 + + +def test_unrelated_download_never_becomes_active(observer_modules, monkeypatch, tmp_path): + worker, _ = observer_modules + partial = tmp_path / "hub" / "models--generated" / "blobs" / "weights.incomplete" + partial.parent.mkdir(parents=True) + partial.write_bytes(bytes(range(16))) + monkeypatch.setattr(worker, "_process_tree_open_paths", lambda *_: set()) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + assert tracker.poll(0)["state"] == "IDLE" + + +def test_owned_download_ends_when_partial_disappears(observer_modules, monkeypatch, tmp_path): + worker, _ = observer_modules + partial = tmp_path / "hub" / "models--generated" / "blobs" / "weights.incomplete" + partial.parent.mkdir(parents=True) + partial.write_bytes(bytes(range(32))) + scans = [] + monkeypatch.setattr(worker, "_process_tree_open_paths", lambda pid: ( + scans.append(pid) or {worker._normalized_path(partial)} + )) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + first = tracker.poll(0) + assert first["state"] == "ACTIVE" + assert tracker.poll(1)["last_progress"] == 0 + assert len(scans) == 1 + partial.rename(partial.with_suffix("")) + finished = tracker.poll(2) + assert finished["state"] == "IDLE" + assert finished["completed_epoch"] == first["epoch"] + assert tracker.poll(3)["completed_epoch"] == finished["completed_epoch"] + + +@pytest.mark.parametrize("existing_final", [False, True]) +def test_download_end_does_not_depend_on_final_file_layout( + observer_modules, monkeypatch, tmp_path, existing_final +): + worker, _ = observer_modules + partial = tmp_path / "hub" / "models--generated" / "blobs" / "weights.incomplete" + partial.parent.mkdir(parents=True) + partial.write_bytes(bytes(range(8))) + if existing_final: + partial.with_suffix("").write_bytes(bytes(range(32))) + monkeypatch.setattr(worker, "_process_tree_open_paths", lambda *_: { + worker._normalized_path(partial) + }) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + assert tracker.poll(0)["state"] == "ACTIVE" + partial.unlink() + state = tracker.poll(1) + assert state["state"] == "IDLE" + assert state["completed_epoch"] == 1 + + +def test_one_ended_file_does_not_discard_another_active_download( + observer_modules, monkeypatch, tmp_path +): + worker, _ = observer_modules + directory = tmp_path / "hub" / "models--generated" / "blobs" + directory.mkdir(parents=True) + partials = [directory / f"weights-{index}.incomplete" for index in range(2)] + for partial in partials: + partial.write_bytes(bytes(range(8))) + monkeypatch.setattr(worker, "_process_tree_open_paths", lambda *_: { + worker._normalized_path(partial) for partial in partials + }) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + assert tracker.poll(0)["state"] == "ACTIVE" + partials[0].unlink() + assert tracker.poll(1)["state"] == "ACTIVE" + assert tracker.poll(2)["completed_epoch"] == 0 + partials[1].unlink() + assert tracker.poll(3)["completed_epoch"] == 1 + + +def test_failed_scan_does_not_report_download_end(observer_modules, monkeypatch, tmp_path): + worker, _ = observer_modules + partial = tmp_path / "hub" / "models--generated" / "blobs" / "weights.incomplete" + partial.parent.mkdir(parents=True) + partial.write_bytes(bytes(range(8))) + monkeypatch.setattr(worker, "_process_tree_open_paths", lambda *_: { + worker._normalized_path(partial) + }) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + assert tracker.poll(0)["state"] == "ACTIVE" + monkeypatch.setattr(worker, "_snapshot_hf_downloads", MagicMock(side_effect=OSError)) + with pytest.raises(OSError): + tracker.poll(1) + assert tracker.completed_epoch == 0 + + +def test_recovered_active_episode_cannot_reset_after_unknown(observer_modules): + _, client = observer_modules + budget = client.ExecutionBudget(timeout=10, stall_timeout=600, now=0) + budget.update(4, client.Observation("ACTIVE", 4, 1, 0, 4)) + budget.update(6, client.Observation("UNKNOWN", 6)) + budget.update(9, client.Observation("ACTIVE", 9, 1, 0, 9)) + budget.update(11, client.Observation("IDLE", 11, 1, 1, 9)) + assert budget.execution_elapsed == 7 + budget.update(14, client.Observation("IDLE", 14, 1, 1, 9)) + assert budget.timed_out + + +def test_stale_active_lease_expires_even_if_no_new_observation(observer_modules): + _, client = observer_modules + budget = client.ExecutionBudget(timeout=10, stall_timeout=600, now=0) + active = client.Observation("ACTIVE", 4, 1, 0, 4) + budget.update(4, active) + budget.update(16, active) + assert budget.execution_elapsed == 11 + assert budget.timed_out + assert not budget.hf_download_stalled + + +def test_heartbeat_does_not_count_as_download_progress(observer_modules): + _, client = observer_modules + budget = client.ExecutionBudget(timeout=10, stall_timeout=5, now=0) + for now in range(1, 7): + budget.update(now, client.Observation("ACTIVE", now, 1, 0, 1)) + assert budget.hf_download_stalled + assert not budget.timed_out + + +@pytest.fixture +def protocol_client(observer_modules): + _, client = observer_modules + observer = client.DownloadObserver.__new__(client.DownloadObserver) + observer._socket = MagicMock() + observer._owner = MagicMock() + observer._owner.process.poll.return_value = None + observer._token = "test-session" # noqa: S105 - Fixed local protocol fixture, not a credential. + observer._seq = 0 + observer._started = 10.0 + observer._last = client.Observation("UNKNOWN", 10.0) + observer.failure = None + return observer + + +def _message(**overrides): + message = {"token": "test-session", "seq": 1, "state": "ACTIVE", + "observed_at": 11.0, "epoch": 1, "completed_epoch": 0, "last_progress": 11.0} + return json.dumps(message | overrides).encode(), ("127.0.0.1", 9999) + + +@pytest.mark.parametrize("overrides", [ + {"token": "other-session"}, {"seq": True}, {"seq": 0}, {"state": "invalid"}, + {"observed_at": float("nan")}, {"last_progress": float("inf")}, + {"observed_at": 20}, {"observed_at": 9}, {"last_progress": 12}, + {"epoch": -1}, {"completed_epoch": 2}, {"epoch": True}, +]) +def test_invalid_protocol_messages_are_ignored(protocol_client, overrides): + protocol_client._socket.recvfrom.side_effect = [_message(**overrides), BlockingIOError] + assert protocol_client.poll(11).state == "UNKNOWN" + assert protocol_client._seq == 0 + + +def test_message_processing_is_bounded_even_with_flood(protocol_client): + protocol_client._socket.recvfrom.return_value = _message() + assert protocol_client.poll(11).state == "ACTIVE" + assert protocol_client._socket.recvfrom.call_count == 8 + + +def test_replayed_or_backward_messages_cannot_extend_active_lease(protocol_client): + protocol_client._socket.recvfrom.side_effect = [_message(), BlockingIOError] + protocol_client.poll(11) + protocol_client._socket.recvfrom.side_effect = [ + _message(seq=1, observed_at=12), _message(seq=2, observed_at=10.5), + _message(seq=3, epoch=0), BlockingIOError, + ] + assert protocol_client.poll(12).observed_at == 11 + assert protocol_client._seq == 1 + + +def test_helper_exit_discards_queued_active_state(protocol_client): + owner = protocol_client._owner + channel = protocol_client._socket + owner.process.poll.return_value = 17 + assert protocol_client.poll(11).state == "UNKNOWN" + channel.recvfrom.assert_not_called() + owner.close.assert_called_once() + assert protocol_client.failure == "observer process exited" + + +def test_silent_helper_is_reaped_and_client_fails_closed(protocol_client): + owner = protocol_client._owner + protocol_client._socket.recvfrom.side_effect = BlockingIOError + assert protocol_client.poll(16).state == "UNKNOWN" + owner.close.assert_called_once() + assert "no completed observation" in protocol_client.failure + + +@pytest.mark.parametrize("cache_var", [ + "HF_HUB_CACHE", "HUGGINGFACE_HUB_CACHE", "HF_DATASETS_CACHE", "HF_XET_CACHE" +]) +def test_cache_root_overrides(observer_modules, tmp_path, cache_var): + worker, _ = observer_modules + roots = worker._hf_cache_roots({"HF_HOME": str(tmp_path / "home"), cache_var: str(tmp_path)}) + index = 1 if cache_var == "HF_DATASETS_CACHE" else 2 if cache_var == "HF_XET_CACHE" else 0 + assert roots[index] == tmp_path + + +def test_slow_ownership_probe_timestamps_progress_after_scan( + observer_modules, monkeypatch, tmp_path +): + worker, _ = observer_modules + partial = tmp_path / "hub" / "models--generated" / "blobs" / "weights.incomplete" + partial.parent.mkdir(parents=True) + partial.write_bytes(bytes(range(8))) + clock = [10.0] + monkeypatch.setattr(worker.time, "monotonic", lambda: clock[0]) + + def open_paths(_pid): + clock[0] += 3.0 + return {worker._normalized_path(partial)} + + monkeypatch.setattr(worker, "_process_tree_open_paths", open_paths) + tracker = worker.DownloadTracker({"HF_HOME": str(tmp_path)}, 123) + state = tracker.poll() + assert state["state"] == "ACTIVE" + assert state["last_progress"] == clock[0] diff --git a/tests/unit/eval/test_run_eval_script.py b/tests/unit/eval/test_run_eval_script.py index d67393f86..b5d7f6823 100644 --- a/tests/unit/eval/test_run_eval_script.py +++ b/tests/unit/eval/test_run_eval_script.py @@ -14,14 +14,15 @@ from __future__ import annotations import argparse +import errno import importlib.util import json +import os import sys -import time -from io import BytesIO from pathlib import Path from types import SimpleNamespace from unittest.mock import MagicMock, patch +from unittest.mock import call as mock_call import pytest @@ -191,192 +192,292 @@ def test_other_eps_do_not_skip(self, run_eval, ep): assert run_eval._should_skip_winml_quant(ep) is False -class TestKillProcessTree: - def test_access_denied_child_does_not_escape(self, run_eval): - import psutil - - child = MagicMock() - child.kill.side_effect = psutil.AccessDenied(pid=456) - parent = MagicMock() - parent.children.return_value = [child] +class TestRunSubprocessDiskFullRetry: + """Exercise the retry boundary without launching processes or deleting caches.""" + @pytest.fixture + def retry_harness(self, run_eval): + calls = MagicMock() + args, timeout = ["controlled-child"], 30 + success = {"stdout": "", "stderr": "", "exit_code": 0, "timeout": False} with ( - patch.object(psutil, "Process", return_value=parent), - patch.object(psutil, "wait_procs") as wait_procs, + patch.object(run_eval, "_run_subprocess_once", calls.run_once), + patch.object(run_eval, "_clear_disk_caches", calls.clear_caches), + patch.object(run_eval, "_run_subprocess", wraps=run_eval._run_subprocess) as run, ): - run_eval._kill_process_tree(123) + yield SimpleNamespace( + calls=calls, args=args, timeout=timeout, success=success, run=run + ) + # Recursive retries would re-enter the patched module-level function. + run.assert_called_once_with(args, timeout) + + @pytest.fixture(params=["enospc", "winerror112"]) + def disk_full_errors(self, request): + error_number = errno.ENOSPC if request.param == "enospc" else errno.EIO + errors = [OSError(error_number, os.strerror(error_number)) for _ in range(2)] + if request.param == "winerror112": + for error in errors: + error.winerror = 112 + return errors + + @pytest.mark.parametrize("stream", ["stdout", "stderr"]) + def test_disk_full_output_retries_once_then_returns_result(self, retry_harness, stream): + harness = retry_harness + failure = { + **harness.success, + "exit_code": 1, + stream: str(OSError(errno.ENOSPC, os.strerror(errno.ENOSPC))), + } + harness.calls.run_once.side_effect = [failure, harness.success] - parent.kill.assert_called_once_with() - wait_procs.assert_called_once_with([child, parent], timeout=5) + result = harness.run(harness.args, harness.timeout) - def test_children_race_falls_back_to_platform_kill(self, run_eval): - import psutil + assert result is harness.success + assert harness.calls.mock_calls == [ + mock_call.run_once(harness.args, harness.timeout), + mock_call.clear_caches(), + mock_call.run_once(harness.args, harness.timeout), + ] - parent = MagicMock() - parent.children.side_effect = psutil.NoSuchProcess(pid=123) + @pytest.mark.parametrize("stream", ["stdout", "stderr"]) + def test_repeated_disk_full_output_stops_after_two_attempts(self, retry_harness, stream): + harness = retry_harness + failures = [ + { + **harness.success, + "exit_code": 1, + stream: str(OSError(errno.ENOSPC, os.strerror(errno.ENOSPC))), + } + for _ in range(2) + ] + harness.calls.run_once.side_effect = failures - with ( - patch.object(psutil, "Process", return_value=parent), - patch.object(run_eval.platform, "system", return_value="Windows"), - patch.object(run_eval.subprocess, "run") as subprocess_run, - ): - run_eval._kill_process_tree(123) + result = harness.run(harness.args, harness.timeout) - subprocess_run.assert_called_once_with( - ["taskkill", "/F", "/T", "/PID", "123"], - capture_output=True, - ) + assert result is failures[-1] + assert harness.calls.mock_calls == [ + mock_call.run_once(harness.args, harness.timeout), + mock_call.clear_caches(), + mock_call.run_once(harness.args, harness.timeout), + ] + def test_disk_full_exception_cleans_once_then_retries_successfully( + self, retry_harness, disk_full_errors + ): + harness = retry_harness + harness.calls.run_once.side_effect = [disk_full_errors[0], harness.success] -class TestRunSubprocessTimeouts: - _HF_CACHE_ENV_VARS = ( - "HF_HOME", - "HF_HUB_CACHE", - "HUGGINGFACE_HUB_CACHE", - "HF_DATASETS_CACHE", - "HF_XET_CACHE", - "XDG_CACHE_HOME", - ) + result = harness.run(harness.args, harness.timeout) - @classmethod - def _cache_env(cls, run_eval, **overrides): - env = { - name: value - for name, value in run_eval.os.environ.items() - if name not in cls._HF_CACHE_ENV_VARS - } - env.update({name: str(value) for name, value in overrides.items()}) - return patch.dict(run_eval.os.environ, env, clear=True) - - @staticmethod - def _download_script( - incomplete: Path, - delays: list[float], - *, - before_download: float = 0.0, - after_download: float = 0.0, - ) -> str: - return "\n".join( - [ - "import time", - "from pathlib import Path", - f"path = Path({str(incomplete)!r})", - f"time.sleep({before_download})", - "path.parent.mkdir(parents=True, exist_ok=True)", - "with path.open('wb') as stream:", - *[ - line - for delay in delays - for line in ( - " stream.write(b'x')", - " stream.flush()", - f" time.sleep({delay})", - ) - ], - "path.unlink()", - f"time.sleep({after_download})", - ] - ) + assert result is harness.success + assert harness.calls.mock_calls == [ + mock_call.run_once(harness.args, harness.timeout), + mock_call.clear_caches(), + mock_call.run_once(harness.args, harness.timeout), + ] - def _run_download_timeline( - self, run_eval, tmp_path, *, cache_variable, after_download=0.55, scan_delay=0.0 + def test_repeated_disk_full_exception_raises_after_two_attempts( + self, retry_harness, disk_full_errors ): - """Exercise real cache detection/accounting without subsecond OS scheduling races. + harness = retry_harness + harness.calls.run_once.side_effect = disk_full_errors - The old test allowed only 150 ms for Python startup before a 500-ms - deadline, and depended on Windows open_files observing a brief handle. - This clock controls only process wait/handle discovery; cache-root - resolution, snapshots, ownership matching and budget accounting are real. - """ - cache_home = tmp_path / "huggingface" if cache_variable == "XDG_CACHE_HOME" else tmp_path - incomplete = cache_home / "hub" / "models--acme--model" / "blobs" / "model.incomplete" - now = 0.0 - download_start, download_end = 0.35, 1.35 - finish = download_end + after_download - observed = [] - proc = MagicMock(pid=123, returncode=0, stdout=BytesIO(), stderr=BytesIO()) - - def open_paths(pid): - nonlocal now - assert pid == proc.pid - now += scan_delay - if download_start <= now < download_end: - incomplete.parent.mkdir(parents=True, exist_ok=True) - with incomplete.open("ab") as stream: - stream.write(b"x") - observed.append(incomplete) - return {run_eval._normalized_path(incomplete)} - incomplete.unlink(missing_ok=True) - return set() - - def wait(timeout): - nonlocal now - now = min(now + timeout, max(now, finish)) - if now >= finish: - return 0 - raise run_eval.subprocess.TimeoutExpired("controlled-child", timeout) - - proc.wait.side_effect = wait - with ( - self._cache_env(run_eval, **{cache_variable: tmp_path}), - patch.object(run_eval, "time", SimpleNamespace(perf_counter=lambda: now)), - patch.object(run_eval.subprocess, "Popen", return_value=proc), - patch.object(run_eval, "_process_tree_open_paths", side_effect=open_paths), - patch.object(run_eval, "_kill_process_tree") as kill_tree, - patch.object(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 2.0), - ): - result = run_eval._run_subprocess(["controlled-child"], timeout=0.5) + with pytest.raises(OSError) as exc_info: + harness.run(harness.args, harness.timeout) - assert observed, "The actual cache snapshot must detect the owned download" - if result["timeout"]: - kill_tree.assert_called_once_with(proc.pid) - proc.kill.assert_called_once() - else: - kill_tree.assert_not_called() - proc.kill.assert_not_called() - return result - - @pytest.mark.parametrize("cache_variable", ["HF_HOME", "XDG_CACHE_HOME"]) - def test_execution_timeout_restarts_after_hf_download(self, run_eval, tmp_path, cache_variable): - result = self._run_download_timeline(run_eval, tmp_path, cache_variable=cache_variable) - assert result["exit_code"] == 0, result - assert result["elapsed"] == 1.9 - assert result["timeout"] is False - assert result["hf_download_stalled"] is False + assert exc_info.value is disk_full_errors[-1] + assert harness.calls.mock_calls == [ + mock_call.run_once(harness.args, harness.timeout), + mock_call.clear_caches(), + mock_call.run_once(harness.args, harness.timeout), + ] - def test_execution_timeout_restarts_after_hf_download_with_slow_handle_scan( - self, run_eval, tmp_path + @pytest.mark.parametrize("error_number", [errno.EACCES, errno.EIO]) + def test_unrelated_oserror_is_not_retried(self, retry_harness, error_number): + harness = retry_harness + error = OSError(error_number, os.strerror(error_number)) + harness.calls.run_once.side_effect = error + + with pytest.raises(OSError) as exc_info: + harness.run(harness.args, harness.timeout) + + assert exc_info.value is error + assert harness.calls.mock_calls == [mock_call.run_once(harness.args, harness.timeout)] + + +class TestExecutionBudget: + """Keep timing assertions independent of process startup and native scans.""" + + @pytest.fixture + def protocol(self, run_eval): + return sys.modules[run_eval.DownloadObserver.__module__] + + @pytest.mark.parametrize("state", ["IDLE", "UNKNOWN"]) + def test_no_owned_download_uses_original_timeout(self, protocol, state): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=30, now=0) + budget.update(1, protocol.Observation(state, 1)) + assert not budget.timed_out + budget.update(2, protocol.Observation(state, 2)) + assert budget.execution_elapsed == 2 + assert budget.timed_out + assert not budget.hf_download_stalled + + def test_only_confirmed_completion_restarts_execution_budget_once(self, protocol): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=30, now=0) + budget.update(1, protocol.Observation("ACTIVE", 1, epoch=1, last_progress=1)) + for now in range(2, 9): + budget.update(now, protocol.Observation("ACTIVE", now, epoch=1, last_progress=now)) + assert budget.execution_elapsed == 1 + assert not budget.timed_out + budget.update(9, protocol.Observation("IDLE", 9, epoch=1, completed_epoch=1)) + assert budget.execution_elapsed == 0 + budget.update(10, protocol.Observation("IDLE", 10, epoch=1, completed_epoch=1)) + assert budget.execution_elapsed == 1 + assert not budget.timed_out + budget.update(11, protocol.Observation("IDLE", 11, epoch=1, completed_epoch=1)) + assert budget.timed_out + + def test_slow_scan_pauses_only_until_last_observation_expires(self, protocol): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=30, now=0) + active = protocol.Observation("ACTIVE", 1, epoch=1, last_progress=1) + budget.update(1, active) + expiry = active.observed_at + protocol.FRESHNESS_SECONDS + budget.update(expiry, active) + assert budget.execution_elapsed == 1 + assert not budget.timed_out + budget.update(expiry + 1, active) + assert budget.execution_elapsed == 2 + assert budget.timed_out + assert not budget.hf_download_stalled + + @pytest.mark.parametrize("state", ["UNKNOWN", "IDLE"]) + def test_loss_of_ownership_is_not_completion(self, protocol, state): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=30, now=0) + budget.update(1, protocol.Observation("ACTIVE", 1, epoch=1, last_progress=1)) + budget.update(2, protocol.Observation(state, 2, epoch=1)) + assert budget.execution_elapsed == 1 + # A delayed completion after loss of ownership must not grant more time. + budget.update(3, protocol.Observation("IDLE", 3, epoch=1, completed_epoch=1)) + assert budget.execution_elapsed == 2 + assert budget.timed_out + + def test_unseen_completion_does_not_grant_a_fresh_budget(self, protocol): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=30, now=0) + budget.update(2, protocol.Observation("IDLE", 2, epoch=1, completed_epoch=1)) + assert budget.timed_out + assert budget.execution_elapsed == 2 + + def test_download_stall_has_an_independent_deadline(self, protocol): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=4, now=0) + for now in range(1, 5): + budget.update(now, protocol.Observation("ACTIVE", now, epoch=1, last_progress=1)) + assert not budget.hf_download_stalled + assert not budget.timed_out + budget.update(5, protocol.Observation("ACTIVE", 5, epoch=1, last_progress=1)) + assert budget.hf_download_stalled + assert not budget.timed_out + + def test_progress_advances_stall_deadline(self, protocol): + budget = protocol.ExecutionBudget(timeout=2, stall_timeout=4, now=0) + budget.update(1, protocol.Observation("ACTIVE", 1, epoch=1, last_progress=1)) + budget.update(4, protocol.Observation("ACTIVE", 4, epoch=1, last_progress=4)) + budget.update(7, protocol.Observation("ACTIVE", 7, epoch=1, last_progress=4)) + assert not budget.hf_download_stalled + budget.update(8, protocol.Observation("ACTIVE", 8, epoch=1, last_progress=4)) + assert budget.hf_download_stalled + + +@pytest.fixture +def subprocess_harness(run_eval, monkeypatch): + """Fake only the owned process and observer; use real budget and spool files.""" + protocol = sys.modules[run_eval.DownloadObserver.__module__] + state = SimpleNamespace(now=0.0, finish=float("inf"), streams=[], cleanup=[]) + proc = MagicMock(pid=123, returncode=None, stdout=None, stderr=None) + tree = MagicMock(process=proc) + tree.__enter__.return_value = tree + observer = MagicMock(failure=None) + observer.__enter__.return_value = observer + observer.poll.side_effect = lambda now: protocol.Observation("IDLE", now) + + def poll(): + if state.now >= state.finish: + proc.returncode = 0 + return proc.returncode + + def wait(timeout): + state.now = min(state.now + timeout, state.finish) + if poll() is not None: + return proc.returncode + raise run_eval.subprocess.TimeoutExpired("controlled-child", timeout) + + def spawn(_args, *, stdout, stderr, env): + state.streams.extend([stdout, stderr]) + assert stdout.fileno() != stderr.fileno() + assert env["PYTHONIOENCODING"] == "utf-8" + assert env["HF_HUB_DOWNLOAD_TIMEOUT"] == str(int(run_eval._HF_DOWNLOAD_STALL_TIMEOUT)) + stdout.write(b"child output\n") + stderr.write(b"child diagnostic\n") + return tree + + def exit_observer(*_exc): + state.cleanup.append("observer") + + def exit_tree(*_exc): + assert all(not stream.closed for stream in state.streams) + state.cleanup.append("tree") + + proc.poll.side_effect = poll + proc.wait.side_effect = wait + tree.__exit__.side_effect = exit_tree + observer.__exit__.side_effect = exit_observer + monkeypatch.setattr(run_eval, "ManagedProcess", MagicMock(side_effect=spawn)) + monkeypatch.setattr(run_eval, "DownloadObserver", MagicMock(return_value=observer)) + monkeypatch.setattr( + run_eval, "time", + SimpleNamespace(monotonic=lambda: state.now, perf_counter=lambda: state.now), + ) + state.process, state.tree, state.observer, state.protocol = proc, tree, observer, protocol + yield state + assert state.cleanup == ["observer", "tree"] + assert all(stream.closed for stream in state.streams) + proc.communicate.assert_not_called() + + +class TestRunSubprocessTimeouts: + @pytest.mark.parametrize("finish", [2.25, 4.0]) + def test_execution_restarts_then_still_times_out_after_completed_download( + self, run_eval, subprocess_harness, finish ): - result = self._run_download_timeline( - run_eval, tmp_path, cache_variable="HF_HOME", scan_delay=0.7 - ) - assert result["exit_code"] == 0, result - assert result["timeout"] is False + harness = subprocess_harness + harness.finish = finish + + def observe(now): + if now < 0.25: + return harness.protocol.Observation("IDLE", now) + if now < 1.5: + return harness.protocol.Observation("ACTIVE", now, epoch=1, last_progress=now) + return harness.protocol.Observation("IDLE", now, epoch=1, completed_epoch=1) + + harness.observer.poll.side_effect = observe + result = run_eval._run_subprocess(["controlled-child"], timeout=1) + timed_out = finish > 2.5 + assert result["timeout"] is timed_out + assert result["exit_code"] == (-1 if timed_out else 0) assert result["hf_download_stalled"] is False + assert harness.tree.terminate.call_count == int(timed_out) + assert result["stdout"] == "child output\n" + assert result["stderr"] == "child diagnostic\n" - @pytest.mark.parametrize("cache_variable", ["HF_HOME", "XDG_CACHE_HOME"]) - def test_execution_still_times_out_after_completed_download( - self, run_eval, tmp_path, cache_variable + def test_stalled_download_is_retryable_not_execution_timeout( + self, run_eval, subprocess_harness, monkeypatch ): - result = self._run_download_timeline( - run_eval, tmp_path, cache_variable=cache_variable, after_download=2.0 + harness = subprocess_harness + monkeypatch.setattr(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 1) + harness.observer.poll.side_effect = lambda now: harness.protocol.Observation( + "ACTIVE", now, epoch=1, last_progress=0 ) - assert result["exit_code"] == -1, result - assert result["timeout"] is True - assert result["hf_download_stalled"] is False - - def test_stalled_hf_download_uses_independent_timeout(self, run_eval, tmp_path): - incomplete = tmp_path / "hub" / "models--acme--model" / "blobs" / "model.incomplete" - script = self._download_script(incomplete, [5.0]) - - with ( - self._cache_env(run_eval, HF_HOME=tmp_path), - patch.object(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 0.5), - ): - result = run_eval._run_subprocess([sys.executable, "-c", script], timeout=5) - + result = run_eval._run_subprocess(["controlled-child"], timeout=0.5) assert result["exit_code"] == -1 - assert result["elapsed"] < 3 assert result["timeout"] is False assert result["hf_download_stalled"] is True assert "Hugging Face download stalled" in result["stderr"] @@ -384,60 +485,61 @@ def test_stalled_hf_download_uses_independent_timeout(self, run_eval, tmp_path): assert classifier.classify_failure(result["stderr"], result["exit_code"]) is ( classifier.FailureType.HF_FETCH_FAIL ) - assert classifier.matches_hf_fetch_retry(result["stderr"]) is True - - def test_unrelated_hf_download_does_not_suspend_execution_timeout(self, run_eval, tmp_path): - incomplete = tmp_path / "hub" / "models--acme--model" / "blobs" / "model.incomplete" - sibling_script = self._download_script(incomplete, [0.1] * 20) - - with ( - self._cache_env(run_eval, HF_HOME=tmp_path), - patch.object(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 1.0), - ): - sibling = run_eval.subprocess.Popen([sys.executable, "-c", sibling_script]) - try: - deadline = time.perf_counter() + 2 - while not incomplete.exists() and time.perf_counter() < deadline: - time.sleep(0.01) - assert incomplete.exists() - - result = run_eval._run_subprocess( - [sys.executable, "-c", "import time; time.sleep(1.2)"], - timeout=0.5, - ) - finally: - sibling.kill() - sibling.wait(timeout=5) + assert classifier.matches_hf_fetch_retry(result["stderr"]) + harness.tree.terminate.assert_called_once_with() + @pytest.mark.parametrize("observer_failed", [False, True]) + def test_without_download_uses_original_timeout( + self, run_eval, subprocess_harness, observer_failed + ): + harness = subprocess_harness + if observer_failed: + harness.observer.failure = "observer process exited" + harness.observer.poll.side_effect = lambda now: harness.protocol.Observation( + "UNKNOWN", now + ) + result = run_eval._run_subprocess(["controlled-child"], timeout=1) assert result["exit_code"] == -1 assert result["timeout"] is True assert result["hf_download_stalled"] is False + assert result["elapsed"] == 1 + if observer_failed: + assert harness.observer.failure in result["stderr"] + assert "using the execution timeout" in result["stderr"] + else: + assert "observation unavailable" not in result["stderr"] + harness.tree.terminate.assert_called_once_with() + run_eval.DownloadObserver.assert_called_once() - def test_stalled_xdg_hf_download_uses_independent_timeout(self, run_eval, tmp_path): - incomplete = ( - tmp_path / "huggingface" / "hub" / "models--acme--model" / "blobs" / "model.incomplete" - ) - script = self._download_script(incomplete, [5.0]) + @pytest.mark.parametrize("finish", [0, 0.25]) + def test_exited_child_never_triggers_postexit_observer_read( + self, run_eval, subprocess_harness, finish + ): + harness = subprocess_harness + harness.finish = finish - with ( - self._cache_env(run_eval, XDG_CACHE_HOME=tmp_path), - patch.object(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 0.5), - ): - result = run_eval._run_subprocess([sys.executable, "-c", script], timeout=5) + def observe(now): + assert harness.process.returncode is None, "No post-exit handle scan is allowed" + return harness.protocol.Observation("IDLE", now) - assert result["exit_code"] == -1 + harness.observer.poll.side_effect = observe + result = run_eval._run_subprocess(["controlled-child"], timeout=1) + assert result["exit_code"] == 0 assert result["timeout"] is False - assert result["hf_download_stalled"] is True - - def test_execution_without_hf_download_uses_original_timeout(self, run_eval): - with patch.object(run_eval, "_HF_DOWNLOAD_STALL_TIMEOUT", 1.0): - result = run_eval._run_subprocess( - [sys.executable, "-c", "import time; time.sleep(5)"], timeout=0.5 - ) + assert harness.observer.poll.call_count == int(finish > 0) + harness.tree.terminate.assert_not_called() - assert result["exit_code"] == -1 - assert result["timeout"] is True - assert result["hf_download_stalled"] is False + @pytest.mark.parametrize("interrupt_at", ["wait", "observer"]) + def test_keyboard_interrupt_unwinds_observer_tree_and_spools( + self, run_eval, subprocess_harness, interrupt_at + ): + harness = subprocess_harness + target = harness.process.wait if interrupt_at == "wait" else harness.observer.poll + target.side_effect = KeyboardInterrupt + with pytest.raises(KeyboardInterrupt): + run_eval._run_subprocess(["controlled-child"], timeout=1) + assert harness.observer.__exit__.call_args.args[0] is KeyboardInterrupt + assert harness.tree.__exit__.call_args.args[0] is KeyboardInterrupt def test_curated_target_models_preserve_existing_priorities(run_eval):