diff --git a/AGENTS.md b/AGENTS.md index 1e304b15..9c173fd1 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -140,6 +140,19 @@ The aggregator is a separate process (`python -m inference_endpoint.async_utils. - **Post-run service wait**: `drain_and_build_report` waits for service subprocesses **unbounded** (`wait_for_exit(None)`). `wait_for_exit` SIGKILLs on expiry, and the aggregator writes `final_snapshot.json` — the Report's primary source — as the last thing it does, so a parent-side deadline here trades a wedged-service hang for a lost snapshot. The abort path keeps its own ceiling (`interrupted_teardown_grace_s`, 30s), and the whole-run watchdog stays armed throughout. - **Histogram bucket edges are dynamic per snapshot**: log-spaced over the observed `[min, max]`. Bucket count is fixed at construction; consumers MUST re-render from the snapshot's `(lo, hi, count)` triples each frame and MUST NOT track bucket-by-index across snapshots. +### SWE-bench Pyxis command transport + +`evaluation/swebench_service/swebench_service/pyxis_persistent.py` owns one +long-lived Slurm step per agent environment. The container runs the packaged +`pyxis_command_worker.sh`, staged in a private control mount separate from tool +`/tmp`. Worker startup retries only confirmed prelaunch Slurm failures before any +request is published. A lock serializes callers through one atomically published +request directory and completion marker. +Commands run in fresh PID namespaces; accepted requests are never replayed after an +uncertain failure. `pyxis_environment.py` owns container creation, command result +mapping, and worker-before-container cleanup. All tool commands use the persistent +worker. Evaluation remains in `pyxis_worker.py`. + ### CLI Modes CLI is auto-generated from `config/schema.py` Pydantic models via cyclopts. Fields annotated with `cyclopts.Parameter(alias="--flag")` get flat shorthands; all other fields get auto-generated dotted flags (kebab-case). @@ -283,7 +296,7 @@ src/inference_endpoint/ │ ├── types.py # Pydantic: VideoPathRequest, VideoPathResponse, VideoPayloadResponse │ └── adapter.py # VideoGenAdapter (HttpRequestAdapter) + VideoGenAccumulator (no-op) ├── evaluation/ # Accuracy evaluation (extractor, scoring, livecodebench) -│ └── swebench_service/ # Isolated uv service for Docker-backed SWE-bench runs +│ └── swebench_service/ # Isolated uv service for Docker/Pyxis SWE-bench runs ├── compliance/ # Submission compliance checks (config-lock, accuracy gate, run validity) │ ├── __init__.py │ └── checker.py # check_submission() + Check/ComplianceReport (Edge-Agentic ruleset) diff --git a/examples/10_Agentic_Inference/README.md b/examples/10_Agentic_Inference/README.md index a8550f3b..a9d9696b 100644 --- a/examples/10_Agentic_Inference/README.md +++ b/examples/10_Agentic_Inference/README.md @@ -260,7 +260,7 @@ Reference mean values are shown in parentheses. | Metric | Kimi K3 | Qwen3.6-35B-A3B | DeepSeek-V4.1-Flash | | ------------------ | ------------------------ | ------------------------ | ------------------------ | -| Inline accuracy | `>= 58.32%` (`58.9%`) | `>= 55.86%` (`56.43%`) | `>= 52.36%` (`53.16%`) | +| Inline accuracy | `>= 58.32%` (`58.9%`) | `>= 55.86%` (`56.43%`) | `>= 52.36%` (`53.16%`) | | OSL per-turn mean¹ | `425-520` tokens (`472`) | `344-422` tokens (`383`) | `793-970` tokens (`882`) | | SWE-bench accuracy | `>= 93.5%` (`94.83%`) | `>= 69%` (`71.7%`) | `>= 96.4%` (`97.5%`) | diff --git a/pyproject.toml b/pyproject.toml index 611f3013..bb6cdaed 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -96,7 +96,7 @@ dependencies = [ "colorama==0.4.6", # Fix pytz-2024 import warning "pytz==2026.1.post1", - "urllib3==2.7.0", + "urllib3==2.8.0", "pyyaml==6.0.3", # anyio is pulled in transitively (openai -> httpx -> anyio); pinned here to # force past CVE-2026-63374 / CVE-2026-64847 (both fixed in 4.14.2), which @@ -130,7 +130,7 @@ dev = [ # patched here, pinned only inside the bfcl fork. Closes CVE-2025-68146 / # CVE-2026-22701 (filelock) and CVE-2026-22702 (virtualenv). "filelock>=3.20.3", - "virtualenv>=20.36.1", + "virtualenv==21.7.13", # pip is pulled in transitively (pip-audit -> pip-api -> pip); force past # PYSEC-2026-3721 (fixed in 26.2), which otherwise fails `uv run pip-audit`. "pip==26.2", diff --git a/src/inference_endpoint/evaluation/swebench_service/README.md b/src/inference_endpoint/evaluation/swebench_service/README.md index cd11904b..0e6afbad 100644 --- a/src/inference_endpoint/evaluation/swebench_service/README.md +++ b/src/inference_endpoint/evaluation/swebench_service/README.md @@ -66,10 +66,32 @@ never forwarded. During generation, the service still uses mini-swe-agent for the agent loop and model requests, but replaces its Docker environment with `PyxisEnvironment`. Every -trajectory receives a named, writable Pyxis container. Each tool call becomes an -overlapping `srun` step in that container, preserving filesystem changes across -turns. Tool commands run in private PID namespaces so one trajectory cannot signal -processes belonging to another trajectory. +trajectory receives a named, writable Pyxis container and one long-lived overlapping +`srun` command worker. Tool calls use atomic request and response files in the +container's private `/tmp` mount, avoiding Slurm step creation and Enroot startup +on every turn. Each worker handles one request at a time; it publishes a completion +marker only after the command exits and its output is closed. +Filesystem changes persist, while each command runs in a fresh shell +and private PID namespace. Commands cannot signal the worker or other trajectories; +remaining child processes are removed when the command's PID namespace exits. +The service must stay on the allocated node, and node-local `TMPDIR` is recommended +for the request files. + +The command worker is the packaged `swebench_service/pyxis_command_worker.sh`. +The Python service copies it into the private `/tmp` mount and starts it with Bash +inside the task container; no service Python installation is needed in task images. + +Command failures preserve their exit status and merged stdout/stderr. A command +timeout terminates its process group with a five-second kill grace; loss of the +worker, invalid responses, and driver deadlines fail the run as infrastructure +errors. An accepted request is never automatically replayed because its execution +may already have changed the repository. The worker is stopped and reaped before +its container and temporary files are removed. + +All tool commands use the persistent worker. +Container initialization, worker startup, evaluation, and cleanup still use Slurm +steps. This reduces per-command scheduler traffic; it does not bypass allocation +limits or guarantee any particular end-to-end evaluation time. After generation, the Pyxis worker evaluates each prediction in a fresh `srun` container step because the Docker-based SWE-bench evaluator cannot run on the diff --git a/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_command_worker.sh b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_command_worker.sh new file mode 100644 index 00000000..0b9106a5 --- /dev/null +++ b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_command_worker.sh @@ -0,0 +1,49 @@ +#!/usr/bin/env bash +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +set -uo pipefail +root=$1 +shift +interpreter=("$@") +request="$root/running" +touch "$root/ready" || exit 70 + +while [ ! -e "$root/stop" ]; do + if [ ! -d "$root/request" ]; then + sleep 0.05 + continue + fi + # The host serializes callers and removes the previous result before publishing. + [ ! -e "$request" ] || exit 70 + mv -- "$root/request" "$request" || exit 70 + timeout_s=$(cat "$request/timeout") || exit 70 + case "$timeout_s" in ''|*[!0-9]*|0) exit 70 ;; esac + # A sentinel preserves trailing newlines in the command and working directory. + cwd=$(cat "$request/cwd" && printf x) || exit 70 + # Keep tool text out of supervisor argv so pkill -f cannot match it there. + unshare --pid --fork --mount-proc \ + timeout -k 5 "$timeout_s" bash -c ' + status=$1; cwd=$2; shift 2 + command=$(cat <&3 && printf x) || exit 70 + exec 3<&- + if cd -- "$cwd"; then + "$@" "${command%x}" + rc=$? + else + rc=125 + fi + printf "%s\n" "$rc" > "$status" || exit 70 + exit "$rc" + ' command-status "$request/command_status" "${cwd%x}" \ + "${interpreter[@]}" 3< "$request/command" > "$request/output" 2>&1 + returncode=$? + timed_out=0 + # Explicit exits 124/137 are command results, not timeout notifications. + if [ "$(cat "$request/command_status" 2>/dev/null)" != "$returncode" ]; then + case "$returncode" in 124|137) timed_out=1 ;; *) exit 70 ;; esac + fi + size=$(wc -c < "$request/output") || exit 70 + printf '%s %s %s\n' "$returncode" "$timed_out" "$((size))" > "$request/complete.tmp" || exit 70 + mv -- "$request/complete.tmp" "$request/complete" || exit 70 +done diff --git a/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_environment.py b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_environment.py index fc3b3a62..8b9366f6 100644 --- a/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_environment.py +++ b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_environment.py @@ -6,18 +6,23 @@ import logging import os import platform -import random import re +import shutil import subprocess import tempfile import threading -import time import uuid from pathlib import Path from typing import Any from pydantic import AliasChoices, BaseModel, Field +from .pyxis_persistent import PersistentExecChannel +from .pyxis_slurm import SRUN_MAX_ATTEMPTS as _SRUN_MAX_ATTEMPTS +from .pyxis_slurm import ( + is_retryable_prelaunch_failure as _is_retryable_prelaunch_failure, +) +from .pyxis_slurm import wait_for_prelaunch_retry from .runner import RunnerError logger = logging.getLogger(__name__) @@ -45,18 +50,8 @@ "SLURM_CONF", ) _STEP_STATUS = "/tmp/.mlperf_srun_status" -_SRUN_MAX_ATTEMPTS = 5 -_RETRYABLE_PRELAUNCH_ERRORS = ( - "spank_sybil: rpc request error", - "required plugin spank_sybil.so", - "failed to connect to any sack sockets", - "failed to create token", - "curl: (56) connect tunnel failed", - "unable to confirm allocation for job", -) -_IGNORABLE_SRUN_PREAMBLE_LINES = { - "srun: lua: Checking requeue policy with options:", -} +_PERSISTENT_ROOT = "/.mlperf_persistent_exec" +_COMMAND_WORKER = Path(__file__).with_name("pyxis_command_worker.sh") _STEP_SCRIPT = r"""set +e status_path=$1 timeout_s=$2 @@ -73,14 +68,6 @@ def safe_srun_env() -> dict[str, str]: return {name: os.environ[name] for name in _SAFE_SRUN_ENV if name in os.environ} -def _is_retryable_prelaunch_failure(status: str, output: str) -> bool: - """Return whether Slurm rejected the step before its command started.""" - if status != "pending": - return False - lowered = output.lower() - return any(marker in lowered for marker in _RETRYABLE_PRELAUNCH_ERRORS) - - def build_srun_command( *, argv: list[str], @@ -187,15 +174,7 @@ def run_srun_step( ).strip() retryable = _is_retryable_prelaunch_failure(status, output) if retryable and attempt < _SRUN_MAX_ATTEMPTS: - backoff_s = min(2**attempt, 16) - delay_s = backoff_s + random.uniform(0.0, backoff_s) - logger.warning( - "Retrying Pyxis pre-launch failure in %.1fs (attempt %d/%d)", - delay_s, - attempt, - _SRUN_MAX_ATTEMPTS, - ) - time.sleep(delay_s) + wait_for_prelaunch_retry(attempt) continue if failure_path is not None: @@ -232,6 +211,7 @@ class PyxisEnvironmentConfig(BaseModel): env: dict[str, str] = Field(default_factory=dict) timeout_s: int = Field( default=30, + gt=0, validation_alias=AliasChoices("timeout_s", "timeout"), serialization_alias="timeout", ) @@ -246,70 +226,89 @@ def __init__(self, **kwargs: Any): self.name = f"mswe_{safe_run_id}_{uuid.uuid4().hex[:8]}" self._tmp = tempfile.TemporaryDirectory(prefix=f"pyxis_{self.name}_") self._tmp_dir = Path(self._tmp.name) - self._tmp_dir.chmod(0o1777) self._lock = threading.Lock() self._cleaned = False try: + (self._tmp_dir / "tmp").mkdir(mode=0o700) + shutil.copyfile(_COMMAND_WORKER, self._tmp_dir / _COMMAND_WORKER.name) # A no-op initializes and validates the named persistent container. run_srun_step( image=self.config.image, name=self.name, - mounts=[(self._tmp_dir, "/tmp")], + mounts=[(self._tmp_dir / "tmp", "/tmp")], workdir=self.config.cwd, argv=["true"], - status_path=self._tmp_dir / Path(_STEP_STATUS).name, + status_path=self._tmp_dir / "tmp" / Path(_STEP_STATUS).name, timeout_s=self.config.timeout_s, failure_path=self.config.infrastructure_failure_path, ) - except RunnerError as exc: + self._persistent_channel = PersistentExecChannel( + self._tmp_dir / "channel", + self._persistent_server_command(), + safe_srun_env(), + failure_path=self.config.infrastructure_failure_path, + launch_timeout_s=self.config.timeout_s + 30, + ) + self._persistent_channel.start() + except (RunnerError, OSError) as exc: self.cleanup() raise RunnerError( f"failed to start Pyxis container for {self.config.image}" ) from exc + def _persistent_server_command(self) -> list[str]: + return build_srun_command( + name=self.name, + # Tool cleanup of /tmp must not remove the worker or its protocol files. + mounts=[ + (self._tmp_dir / "tmp", "/tmp"), + (self._tmp_dir, _PERSISTENT_ROOT), + ], + workdir=self.config.cwd, + argv=[ + "env", + *(f"{key}={value}" for key, value in self.config.env.items()), + "unshare", + "--pid", + "--fork", + "--mount-proc", + "--kill-child", + "bash", + f"{_PERSISTENT_ROOT}/{_COMMAND_WORKER.name}", + f"{_PERSISTENT_ROOT}/channel", + *self.config.interpreter, + ], + ) + def execute( self, action: dict[str, Any], cwd: str = "", *, timeout: int | None = None ) -> dict[str, Any]: command = action.get("command", "") logger.debug("Executing Pyxis command: %s", command) - argv = ["env"] - argv.extend(f"{key}={value}" for key, value in self.config.env.items()) - argv.extend([*self.config.interpreter, command]) - result = run_srun_step( - argv=argv, - status_path=self._tmp_dir / Path(_STEP_STATUS).name, - timeout_s=timeout or self.config.timeout_s, - failure_path=self.config.infrastructure_failure_path, - name=self.name, - mounts=[(self._tmp_dir, "/tmp")], - workdir=cwd or self.config.cwd, + timeout_s = self.config.timeout_s if timeout is None else timeout + result = self._persistent_channel.execute( + command=command, + cwd=cwd or self.config.cwd, + timeout_s=timeout_s, ) output: dict[str, Any] - if result.returncode == 124: + if result.timed_out: output = { - "output": result.stdout, + "output": result.output, "returncode": -1, "exception_info": "The command timed out", "extra": { "exception_type": "TimeoutExpired", - "exception": ( - f"command timed out after {timeout or self.config.timeout_s}s" - ), + "exception": f"command timed out after {timeout_s}s", }, } else: output = { - "output": result.stdout, + "output": result.output, "returncode": result.returncode, "exception_info": "", } lines = output.get("output", "").lstrip().splitlines(keepends=True) - # Some Slurm cli_filter plugins write informational messages to stderr. - # run_srun_step merges stderr into stdout so command errors remain visible, - # which can place this cluster-generated preamble before mini-swe-agent's - # otherwise first-line submission marker. - while lines and lines[0].strip() in _IGNORABLE_SRUN_PREAMBLE_LINES: - lines.pop(0) if ( lines and lines[0].strip() == "COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT" @@ -353,6 +352,14 @@ def cleanup(self) -> None: return self._cleaned = True try: + channel = getattr(self, "_persistent_channel", None) + if channel is not None: + try: + channel.close() + except (OSError, subprocess.SubprocessError): + logger.warning( + "Could not stop Pyxis worker %s", self.name, exc_info=True + ) if os.environ.get("SLURM_JOB_ID", "").strip(): try: subprocess.run( diff --git a/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_persistent.py b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_persistent.py new file mode 100644 index 00000000..0370a61f --- /dev/null +++ b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_persistent.py @@ -0,0 +1,188 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""One command at a time in a long-lived, allocation-local Pyxis worker.""" + +from __future__ import annotations + +import logging +import os +import shutil +import subprocess +import threading +import time +from dataclasses import dataclass +from pathlib import Path + +from .pyxis_slurm import ( + SRUN_MAX_ATTEMPTS, + is_retryable_prelaunch_failure, + wait_for_prelaunch_retry, +) +from .runner import RunnerError + +logger = logging.getLogger(__name__) +_POLL_S = 0.05 + + +@dataclass(frozen=True, slots=True) +class CommandResult: + returncode: int + output: str + timed_out: bool + + +class PersistentExecChannel: + """Serialize commands; never replay them after an uncertain failure.""" + + def __init__( + self, + protocol_dir: Path, + command: list[str], + env: dict[str, str], + failure_path: Path | None = None, + launch_timeout_s: float = 30, + driver_grace_s: float = 30, + shutdown_grace_s: float = 3, + ) -> None: + self.protocol_dir = protocol_dir + self._command = command + self._env = env + self._failure_path = failure_path + self._launch_timeout_s = launch_timeout_s + self._driver_grace_s = driver_grace_s + self._shutdown_grace_s = shutdown_grace_s + self._process: subprocess.Popen[bytes] | None = None + self._lock = threading.Lock() + self._closed = False + # A new directory prevents stale requests or replies from being reused. + protocol_dir.mkdir(mode=0o700) + + def _read_log(self) -> str: + try: + with (self.protocol_dir / "server.log").open("rb") as log: + log.seek(max(0, log.seek(0, os.SEEK_END) - 8000)) + return log.read().decode("utf-8", errors="replace") + except OSError: + # Startup can fail before the log is created. + return "" + + def _fail(self, detail: str) -> RunnerError: + detail += "\n" + self._read_log() + try: + if self._failure_path is not None: + self._failure_path.touch() + except OSError: + logger.warning("Could not mark Pyxis infrastructure failure", exc_info=True) + try: + self._stop() + except (OSError, subprocess.SubprocessError): + logger.warning("Could not stop failed Pyxis worker", exc_info=True) + return RunnerError(f"persistent Pyxis infrastructure failure: {detail}") + + def _wait_for(self, path: Path, timeout_s: float) -> None: + deadline = time.monotonic() + timeout_s + while not path.exists(): + if self._process is None or self._process.poll() is not None: + raise RunnerError("worker is not running; execution is uncertain") + if time.monotonic() >= deadline: + raise RunnerError( + "worker exceeded its driver deadline; execution is uncertain" + ) + time.sleep(_POLL_S) + + def start(self) -> None: + with self._lock: + if self._closed or self._process is not None: + raise RunnerError("persistent Pyxis channel already started or closed") + log_path = self.protocol_dir / "server.log" + ready = self.protocol_dir / "ready" + for attempt in range(1, SRUN_MAX_ATTEMPTS + 1): + try: + with log_path.open("wb") as log: + self._process = subprocess.Popen( + self._command, + stdin=subprocess.DEVNULL, + stdout=log, + stderr=subprocess.STDOUT, + env=self._env, + ) + self._wait_for(ready, self._launch_timeout_s) + return + except (OSError, RunnerError) as exc: + # No requests can be published while start holds the lock. + # Never relaunch a ready worker or a possibly running step. + if ( + attempt < SRUN_MAX_ATTEMPTS + and self._process is not None + and self._process.poll() not in (None, 0) + and not ready.exists() + and is_retryable_prelaunch_failure("pending", self._read_log()) + ): + wait_for_prelaunch_retry(attempt) + continue + raise self._fail(f"worker did not become ready: {exc}") from exc + + def _read_completion(self, request: Path) -> CommandResult: + try: + returncode, timed_out, size = map( + int, (request / "complete").read_text().split() + ) + except ValueError as exc: + raise RunnerError("invalid Pyxis completion marker") from exc + if not 0 <= returncode <= 255 or timed_out not in (0, 1) or size < 0: + raise RunnerError("invalid Pyxis completion marker") + output = (request / "output").read_bytes() + if len(output) != size: + raise RunnerError("incomplete Pyxis command output") + return CommandResult( + returncode, output.decode("utf-8", errors="replace"), bool(timed_out) + ) + + def execute(self, *, command: str, cwd: str, timeout_s: int) -> CommandResult: + if timeout_s <= 0: + raise ValueError("Pyxis command timeout must be positive") + with self._lock: + if self._closed: + raise RunnerError("persistent Pyxis channel is closed") + try: + if self._process is None or self._process.poll() is not None: + raise RunnerError("worker is not running") + staging = self.protocol_dir / ".request" + staging.mkdir() + for name, value in ( + ("command", command), + ("cwd", cwd), + ("timeout", str(timeout_s)), + ): + (staging / name).write_text(value) + os.replace(staging, self.protocol_dir / "request") + running = self.protocol_dir / "running" + self._wait_for(running / "complete", timeout_s + self._driver_grace_s) + result = self._read_completion(running) + shutil.rmtree(running) + return result + except (OSError, RunnerError) as exc: + raise self._fail(str(exc)) from exc + + def _stop(self) -> None: + self._closed = True + try: + (self.protocol_dir / "stop").touch() + finally: + process = self._process + if process is not None: + try: + process.wait(timeout=self._shutdown_grace_s) + except subprocess.TimeoutExpired: + process.terminate() + try: + process.wait(timeout=self._shutdown_grace_s) + except subprocess.TimeoutExpired: + process.kill() + process.wait(timeout=self._shutdown_grace_s) + + def close(self) -> None: + with self._lock: + if not self._closed: + self._stop() diff --git a/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_slurm.py b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_slurm.py new file mode 100644 index 00000000..30179963 --- /dev/null +++ b/src/inference_endpoint/evaluation/swebench_service/swebench_service/pyxis_slurm.py @@ -0,0 +1,39 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Shared retry policy for Slurm steps that have not started their command.""" + +import logging +import random +import time + +logger = logging.getLogger(__name__) +SRUN_MAX_ATTEMPTS = 5 +_RETRYABLE_PRELAUNCH_ERRORS = ( + "spank_sybil: rpc request error", + "required plugin spank_sybil.so", + "failed to connect to any sack sockets", + "failed to create token", + "curl: (56) connect tunnel failed", + "unable to confirm allocation for job", +) + + +def is_retryable_prelaunch_failure(status: str, output: str) -> bool: + """Return whether Slurm rejected the step before its command started.""" + if status != "pending": + return False + lowered = output.lower() + return any(marker in lowered for marker in _RETRYABLE_PRELAUNCH_ERRORS) + + +def wait_for_prelaunch_retry(attempt: int) -> None: + backoff_s = min(2**attempt, 16) + delay_s = backoff_s + random.uniform(0.0, backoff_s) + logger.warning( + "Retrying Pyxis pre-launch failure in %.1fs (attempt %d/%d)", + delay_s, + attempt, + SRUN_MAX_ATTEMPTS, + ) + time.sleep(delay_s) diff --git a/tests/unit/evaluation/swebench_service/test_pyxis_persistent.py b/tests/unit/evaluation/swebench_service/test_pyxis_persistent.py new file mode 100644 index 00000000..cf42934b --- /dev/null +++ b/tests/unit/evaluation/swebench_service/test_pyxis_persistent.py @@ -0,0 +1,548 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import os +import shlex +import shutil +import subprocess +import sys +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path + +import pytest +from inference_endpoint.evaluation.swebench_service.swebench_service import ( + pyxis_environment as environment_mod, +) +from inference_endpoint.evaluation.swebench_service.swebench_service import ( + pyxis_persistent as transport, +) +from inference_endpoint.evaluation.swebench_service.swebench_service.runner import ( + RunnerError, +) + +pytestmark = pytest.mark.unit + + +@pytest.fixture +def portable_unshare(tmp_path, monkeypatch): + """Only bypass namespace creation; run the actual shell worker.""" + if not all(shutil.which(tool) for tool in ("timeout", "bash")): + pytest.skip("requires GNU timeout and bash") + bin_dir = tmp_path / "bin" + bin_dir.mkdir() + unshare = bin_dir / "unshare" + unshare.write_text( + '#!/bin/sh\nwhile [ "${1#--}" != "$1" ]; do shift; done\nexec "$@"\n' + ) + unshare.chmod(0o700) + monkeypatch.setenv("PATH", f"{bin_dir}:{os.environ['PATH']}") + + +@pytest.fixture +def worker(tmp_path, monkeypatch, portable_unshare): + root = tmp_path / "protocol" + commands = [] + + command = [ + "bash", + str(Path(transport.__file__).with_name("pyxis_command_worker.sh")), + str(root), + "bash", + "-c", + ] + popen = subprocess.Popen + + def launch(argv, **kwargs): + commands.append(argv) + return popen(argv, **kwargs) + + monkeypatch.setattr(subprocess, "Popen", launch) + + channel = transport.PersistentExecChannel( + root, + command, + environment_mod.safe_srun_env(), + failure_path=tmp_path / "failed", + launch_timeout_s=2, + driver_grace_s=1, + shutdown_grace_s=0.2, + ) + channel.start() + yield channel, commands + channel.close() + + +def execute(channel, command, cwd, timeout=2): + return channel.execute(command=command, cwd=str(cwd), timeout_s=timeout) + + +@pytest.mark.parametrize("code", [0, 7, 124, 137]) +def test_preserves_exit_code_and_merged_output(worker, tmp_path, code): + channel, launches = worker + result = execute( + channel, f"printf stdout; printf stderr >&2; exit {code}", tmp_path + ) + assert result.output == "stdoutstderr" + assert result.returncode == code + assert result.timed_out is False + assert len(launches) == 1 + + +def test_files_persist_but_shell_state_does_not(worker, tmp_path): + channel, launches = worker + execute(channel, "printf persisted > state; export PRIVATE=value; cd /", tmp_path) + result = execute( + channel, 'cat state; printf "|%s|%s" "$PWD" "${PRIVATE-unset}"', tmp_path + ) + assert result.output == f"persisted|{tmp_path}|unset" + assert result.returncode == 0 + assert len(launches) == 1 + + +def test_preserves_multiline_command_and_binary_output(worker, tmp_path): + channel, _ = worker + result = execute( + channel, "printf 'a\\000b\\377\\n'\n# trailing comment\n", tmp_path + ) + assert result.output == "a\x00b\ufffd\n" + assert result.returncode == 0 + + +def test_timeout_is_reported_and_worker_remains_usable(worker, tmp_path): + channel, launches = worker + result = execute(channel, "printf before; exec sleep 5", tmp_path, timeout=1) + assert result.output == "before" + assert result.timed_out is True + assert result.returncode == 124 + assert execute(channel, "printf recovered", tmp_path).output == "recovered" + assert len(launches) == 1 + + +def test_missing_cwd_is_a_command_error_and_worker_remains_usable(worker, tmp_path): + channel, launches = worker + result = execute(channel, "echo must-not-run", tmp_path / "missing") + assert result.returncode == 125 + assert not result.timed_out + assert "No such file or directory" in result.output + assert execute(channel, "printf recovered", tmp_path).output == "recovered" + assert len(launches) == 1 + + +def test_environment_survives_tool_tmp_cleanup(monkeypatch, tmp_path, portable_unshare): + """Use production mount sources and worker, substituting local paths for mounts.""" + scratch = None + + def initialize(**kwargs): + nonlocal scratch + scratch = next(source for source, dest in kwargs["mounts"] if dest == "/tmp") + + def local_command(*, argv, mounts, **kwargs): + result = [] + for arg in argv: + for source, dest in mounts: + if arg == dest or arg.startswith(dest + "/"): + arg = str(source) + arg[len(dest) :] + break + result.append(arg) + return result + + monkeypatch.delenv("SLURM_JOB_ID", raising=False) + monkeypatch.setattr(environment_mod, "run_srun_step", initialize) + monkeypatch.setattr(environment_mod, "build_srun_command", local_command) + env = environment_mod.PyxisEnvironment( + image="image", run_id="cleanup", cwd=str(tmp_path) + ) + try: + command = ( + "import os, shutil, tempfile; " + f"location = tempfile.mkdtemp(dir={str(scratch)!r}); " + "shutil.rmtree(os.path.dirname(location)); print('cleaned')" + ) + result = env.execute( + {"command": f"{shlex.quote(sys.executable)} -c {shlex.quote(command)}"} + ) + assert result["returncode"] == 0 + assert result["output"] == "cleaned\n" + assert env.execute({"command": "printf recovered"})["output"] == "recovered" + finally: + env.cleanup() + + +@pytest.mark.parametrize( + "message", + [ + "curl: (56) CONNECT tunnel failed, response 403", + "srun: error: Unable to confirm allocation for job 123", + ], +) +def test_startup_retries_confirmed_prelaunch_failure( + monkeypatch, tmp_path, portable_unshare, message +): + root = tmp_path / "protocol" + launches = [] + delays = [] + popen = subprocess.Popen + + def launch(argv, **kwargs): + launches.append(argv) + assert not (root / "request").exists() + assert not (tmp_path / "failed").exists() + if len(launches) < 3: + return popen( + ["bash", "-c", 'printf "%s\\n" "$1" >&2; exit 1', "prelaunch", message], + **kwargs, + ) + return popen(argv, **kwargs) + + monkeypatch.setattr(subprocess, "Popen", launch) + monkeypatch.setattr(transport, "wait_for_prelaunch_retry", delays.append) + channel = transport.PersistentExecChannel( + root, + ["bash", str(environment_mod._COMMAND_WORKER), str(root), "bash", "-c"], + environment_mod.safe_srun_env(), + failure_path=tmp_path / "failed", + launch_timeout_s=2, + ) + try: + channel.start() + assert ( + execute(channel, "printf once >> count; cat count", tmp_path).output + == "once" + ) + assert len(launches) == 3 + assert delays == [1, 2] + assert not (tmp_path / "failed").exists() + finally: + channel.close() + + +@pytest.mark.parametrize( + "finish, attempts", + [ + ("exit 1", 5), + ("exit 0", 1), + ("exec sleep 10", 1), + ], +) +def test_startup_retry_boundaries(monkeypatch, tmp_path, finish, attempts): + launches = [] + popen = subprocess.Popen + + def launch(argv, **kwargs): + launches.append(argv) + return popen(argv, **kwargs) + + monkeypatch.setattr(subprocess, "Popen", launch) + monkeypatch.setattr(transport, "wait_for_prelaunch_retry", lambda attempt: None) + channel = transport.PersistentExecChannel( + tmp_path / "protocol", + ["bash", "-c", "echo 'unable to confirm allocation for job' >&2; " + finish], + environment_mod.safe_srun_env(), + failure_path=tmp_path / "failed", + launch_timeout_s=0.2, + shutdown_grace_s=0.1, + ) + try: + with pytest.raises(RunnerError, match="did not become ready"): + channel.start() + assert len(launches) == attempts + assert channel._process.poll() is not None + assert (tmp_path / "failed").exists() + finally: + channel.close() + + +def test_ready_worker_is_never_relaunched(monkeypatch, tmp_path): + delays = [] + monkeypatch.setattr(transport, "wait_for_prelaunch_retry", delays.append) + root = tmp_path / "protocol" + channel = transport.PersistentExecChannel( + root, + [ + "bash", + "-c", + 'touch "$1/ready"; echo "unable to confirm allocation for job"; exit 1', + "worker", + str(root), + ], + environment_mod.safe_srun_env(), + failure_path=tmp_path / "failed", + ) + try: + channel.start() + channel._process.wait(timeout=2) + with pytest.raises(RunnerError, match="not running"): + execute(channel, "touch unexpected", tmp_path) + assert delays == [] + assert not (tmp_path / "unexpected").exists() + assert (tmp_path / "failed").exists() + finally: + channel.close() + + +def test_process_cleanup_is_a_command_failure(tmp_path): + """Run destructive process matching only inside real, private PID namespaces.""" + if sys.platform != "linux" or not all( + shutil.which(tool) for tool in ("unshare", "timeout", "pkill") + ): + pytest.skip("requires Linux PID namespaces and procps") + isolation = ["unshare", "--pid", "--fork", "--mount-proc", "--kill-child"] + probe = subprocess.run([*isolation, "true"], capture_output=True, timeout=5) + if probe.returncode: + pytest.skip("PID namespaces are not permitted") + root = tmp_path / "protocol" + channel = transport.PersistentExecChannel( + root, + [ + *isolation, + "bash", + str(environment_mod._COMMAND_WORKER), + str(root), + "bash", + "-c", + ], + environment_mod.safe_srun_env(), + launch_timeout_s=2, + ) + try: + channel.start() + result = execute(channel, "sleep 20 & pkill -f 'sleep 20'", tmp_path) + assert result.returncode == 143 + assert not result.timed_out + assert execute(channel, "printf recovered", tmp_path).output == "recovered" + finally: + channel.close() + + +@pytest.mark.parametrize("timeout", [0, -1]) +def test_invalid_timeout_does_not_execute_command(worker, tmp_path, timeout): + channel, _ = worker + with pytest.raises(ValueError, match="positive"): + execute(channel, "touch unexpected", tmp_path, timeout=timeout) + assert not (tmp_path / "unexpected").exists() + assert execute(channel, "true", tmp_path).returncode == 0 + + +def test_concurrent_callers_are_serialized(worker, tmp_path): + channel, launches = worker + + def increment(_): + return execute( + channel, + "n=$(cat count 2>/dev/null || echo 0); sleep 0.05; echo $((n+1)) > count", + tmp_path, + ) + + with ThreadPoolExecutor(max_workers=4) as pool: + results = list(pool.map(increment, range(4))) + assert all(result.returncode == 0 for result in results) + assert (tmp_path / "count").read_text() == "4\n" + assert len(launches) == 1 + + +def test_dead_worker_does_not_restart_or_replay(worker, tmp_path): + channel, launches = worker + channel._process.kill() + channel._process.wait(timeout=2) + with pytest.raises(RunnerError, match="not running"): + execute(channel, "echo duplicated >> count", tmp_path) + with pytest.raises(RunnerError, match="closed"): + execute(channel, "echo duplicated >> count", tmp_path) + assert (tmp_path / "failed").exists() + assert not (tmp_path / "count").exists() + assert len(launches) == 1 + + +def test_driver_deadline_never_replays_an_active_command(worker, tmp_path, monkeypatch): + channel, launches = worker + channel._driver_grace_s = 0.2 + wait_for = channel._wait_for + monkeypatch.setattr( + channel, + "_wait_for", + lambda path, timeout: wait_for(path.with_name("missing"), timeout), + ) + with pytest.raises(RunnerError, match="driver deadline"): + execute(channel, "echo once >> count; sleep 0.4", tmp_path, timeout=1) + with pytest.raises(RunnerError, match="closed"): + execute(channel, "echo twice >> count", tmp_path) + assert (tmp_path / "count").read_text() == "once\n" + assert (tmp_path / "failed").exists() + assert len(launches) == 1 + + +@pytest.mark.parametrize("fault", ["metadata", "truncated", "missing"]) +def test_completion_corruption_fails_the_run(worker, tmp_path, monkeypatch, fault): + channel, launches = worker + read_completion = channel._read_completion + + def corrupt(request): + if fault == "metadata": + (request / "complete").write_text("garbled") + elif fault == "truncated": + (request / "output").write_bytes(b"") + else: + (request / "output").unlink() + return read_completion(request) + + monkeypatch.setattr(channel, "_read_completion", corrupt) + with pytest.raises(RunnerError, match="infrastructure failure"): + execute(channel, "echo once >> count; printf result", tmp_path) + with pytest.raises(RunnerError, match="closed"): + execute(channel, "echo twice >> count", tmp_path) + assert (tmp_path / "count").read_text() == "once\n" + assert (tmp_path / "failed").exists() + assert len(launches) == 1 + + +def test_close_reaps_worker_and_is_idempotent(worker): + channel, _ = worker + process = channel._process + channel.close() + channel.close() + assert process.poll() == 0 + + +def test_failure_marker_error_does_not_skip_worker_shutdown( + worker, tmp_path, monkeypatch +): + channel, _ = worker + touch = Path.touch + + def fail_marker(path, *args, **kwargs): + if path == tmp_path / "failed": + raise OSError("disk full") + return touch(path, *args, **kwargs) + + def bad_reply(request): + raise RunnerError("invalid reply") + + monkeypatch.setattr(Path, "touch", fail_marker) + monkeypatch.setattr(channel, "_read_completion", bad_reply) + with pytest.raises(RunnerError, match="invalid reply"): + execute(channel, "true", tmp_path) + assert channel._closed + assert channel._process.poll() is not None + + +def test_startup_failure_is_infrastructure_failure(tmp_path): + channel = transport.PersistentExecChannel( + tmp_path / "protocol", + ["bash", "-c", "echo denied >&2; exit 1"], + environment_mod.safe_srun_env(), + failure_path=tmp_path / "failed", + ) + try: + with pytest.raises(RunnerError, match="(?s)did not become ready.*denied"): + channel.start() + assert (tmp_path / "failed").exists() + finally: + channel.close() + + +def test_environment_routes_commands_to_one_worker(monkeypatch, tmp_path): + monkeypatch.setenv("SLURM_JOB_ID", "123") + monkeypatch.setenv("SLURMD_NODENAME", "node") + monkeypatch.setenv("OPENAI_API_KEY", "must-not-reach-worker") + steps = [] + channels = [] + + def step(**kwargs): + steps.append(kwargs) + return subprocess.CompletedProcess([], 0, "", "") + + class Channel: + def __init__(self, protocol_dir, command, env, **kwargs): + self.argv = command + self.env = env + self.closed = False + self.commands = 0 + channels.append(self) + + def start(self): + pass + + def execute(self, **kwargs): + self.commands += 1 + return transport.CommandResult(7, "failed", False) + + def close(self): + self.closed = True + + def cleanup(command, **kwargs): + assert channels[0].closed + return subprocess.CompletedProcess(command, 0) + + monkeypatch.setattr(environment_mod, "run_srun_step", step) + monkeypatch.setattr(environment_mod, "PersistentExecChannel", Channel) + monkeypatch.setattr(subprocess, "run", cleanup) + env = environment_mod.PyxisEnvironment(image="image", run_id="run") + try: + for _ in range(3): + assert env.execute({"command": "exit 7"})["returncode"] == 7 + assert len(steps) == 1 + assert len(channels) == 1 + assert channels[0].commands == 3 + assert "OPENAI_API_KEY" not in channels[0].env + assert "--kill-child" in channels[0].argv + assert "--jobid=123" in channels[0].argv + assert (env._tmp_dir.stat().st_mode & 0o777) == 0o700 + + def fail(**kwargs): + raise RunnerError("worker died; execution is uncertain") + + monkeypatch.setattr(channels[0], "execute", fail) + with pytest.raises(RunnerError, match="execution is uncertain"): + env.execute({"command": "touch state"}) + assert len(steps) == 1 + assert len(channels) == 1 + finally: + env.cleanup() + assert not env._tmp_dir.exists() + + +@pytest.mark.parametrize( + "manifest", ["--1 0 0", "0 0 -1", "0 2 0", "256 0 0", "truncated"] +) +def test_rejects_corrupt_completion_metadata(tmp_path, manifest): + channel = transport.PersistentExecChannel(tmp_path / "protocol", [], {}) + request = tmp_path / "request" + request.mkdir() + (request / "complete").write_text(manifest) + with pytest.raises(RunnerError, match="completion marker"): + channel._read_completion(request) + + +def test_missing_worker_executable_marks_infrastructure_failure(tmp_path): + channel = transport.PersistentExecChannel( + tmp_path / "protocol", + [str(tmp_path / "missing")], + {}, + failure_path=tmp_path / "failed", + ) + try: + with pytest.raises(RunnerError, match="did not become ready"): + channel.start() + assert channel._process is None + assert (tmp_path / "failed").exists() + finally: + channel.close() + + +def test_worker_that_never_becomes_ready_is_reaped(tmp_path): + channel = transport.PersistentExecChannel( + tmp_path / "protocol", + ["sleep", "10"], + environment_mod.safe_srun_env(), + launch_timeout_s=0.1, + shutdown_grace_s=0.1, + failure_path=tmp_path / "failed", + ) + try: + with pytest.raises(RunnerError, match="did not become ready"): + channel.start() + assert channel._process.poll() is not None + finally: + channel.close() + assert channel._process.poll() is not None + assert (tmp_path / "failed").exists() diff --git a/tests/unit/evaluation/swebench_service/test_runner.py b/tests/unit/evaluation/swebench_service/test_runner.py index 7faf7458..03c248fe 100644 --- a/tests/unit/evaluation/swebench_service/test_runner.py +++ b/tests/unit/evaluation/swebench_service/test_runner.py @@ -9,6 +9,7 @@ import types from pathlib import Path from typing import Literal, get_type_hints +from unittest.mock import Mock import msgspec.json import pytest @@ -16,6 +17,9 @@ from inference_endpoint.evaluation.swebench_service.swebench_service import ( pyxis_environment as pyxis_env_mod, ) +from inference_endpoint.evaluation.swebench_service.swebench_service import ( + pyxis_slurm as slurm_mod, +) from inference_endpoint.evaluation.swebench_service.swebench_service import ( pyxis_worker as worker_mod, ) @@ -28,6 +32,10 @@ resolve_image, safe_srun_env, ) +from inference_endpoint.evaluation.swebench_service.swebench_service.pyxis_persistent import ( + CommandResult, + PersistentExecChannel, +) from inference_endpoint.evaluation.swebench_service.swebench_service.runner import ( CancellationToken, PyxisSweBenchRunner, @@ -43,12 +51,25 @@ pytestmark = pytest.mark.unit +@pytest.fixture +def pyxis_channel(monkeypatch): + channel = Mock(spec=PersistentExecChannel) + channel.execute.return_value = CommandResult(0, "ok\n", False) + monkeypatch.setattr( + pyxis_env_mod, "PersistentExecChannel", Mock(return_value=channel) + ) + return channel + + def test_pyxis_implementation_is_confined_to_environment_and_worker_modules(): package_dir = Path(runner_mod.__file__).parent assert {path.name for path in package_dir.glob("pyxis_*") if path.is_file()} == { "pyxis_environment.py", "pyxis_worker.py", + "pyxis_persistent.py", + "pyxis_slurm.py", + "pyxis_command_worker.sh", } @@ -873,7 +894,7 @@ def test_pyxis_srun_environment_withholds_inherited_step_identity(monkeypatch, n def test_pyxis_environment_retries_connect_tunnel_prelaunch_failure( - monkeypatch, tmp_path + monkeypatch, tmp_path, pyxis_channel ): monkeypatch.setenv("SLURM_JOB_ID", "1738605") monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") @@ -894,13 +915,14 @@ def fake_run(command, **kwargs): return subprocess.CompletedProcess(command, 0, stdout="ok\n", stderr="") monkeypatch.setattr(subprocess, "run", fake_run) - monkeypatch.setattr(pyxis_env_mod.random, "uniform", lambda _a, _b: 0.0) - monkeypatch.setattr(pyxis_env_mod.time, "sleep", delays.append) + monkeypatch.setattr(slurm_mod.random, "uniform", lambda _a, _b: 0.0) + monkeypatch.setattr(slurm_mod.time, "sleep", delays.append) environment = PyxisEnvironment(image=tmp_path / "task.sqsh", run_id="run-1") assert calls == 3 assert delays == [2, 4] + pyxis_channel.start.assert_called_once_with() environment.cleanup() @@ -925,8 +947,8 @@ def fake_run(command, **kwargs): return subprocess.CompletedProcess(command, 0, stdout="ok\n", stderr="") monkeypatch.setattr(subprocess, "run", fake_run) - monkeypatch.setattr(pyxis_env_mod.random, "uniform", lambda _a, _b: 0.0) - monkeypatch.setattr(pyxis_env_mod.time, "sleep", delays.append) + monkeypatch.setattr(slurm_mod.random, "uniform", lambda _a, _b: 0.0) + monkeypatch.setattr(slurm_mod.time, "sleep", delays.append) result = pyxis_env_mod.run_srun_step( argv=["true"], @@ -966,8 +988,8 @@ def fake_run(command, **kwargs): return subprocess.CompletedProcess(command, 0, stdout="ok\n", stderr="") monkeypatch.setattr(subprocess, "run", fake_run) - monkeypatch.setattr(pyxis_env_mod.random, "uniform", lambda _a, _b: 0.0) - monkeypatch.setattr(pyxis_env_mod.time, "sleep", delays.append) + monkeypatch.setattr(slurm_mod.random, "uniform", lambda _a, _b: 0.0) + monkeypatch.setattr(slurm_mod.time, "sleep", delays.append) result = pyxis_env_mod.run_srun_step( argv=["true"], @@ -1003,7 +1025,7 @@ def test_pyxis_environment_does_not_retry_non_retryable_prelaunch_failure( def test_pyxis_environment_reuses_named_writable_container( - monkeypatch, tmp_path, caplog + monkeypatch, tmp_path, caplog, pyxis_channel ): monkeypatch.setenv("SLURM_JOB_ID", "1738605") monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") @@ -1035,6 +1057,8 @@ def fake_run(command, **kwargs): ): first = environment.execute({"command": "touch state"}) second = environment.execute({"command": "test -f state"}) + worker_command = environment._persistent_server_command() + worker_script = (environment._tmp_dir / "pyxis_command_worker.sh").read_bytes() environment.cleanup() container_name = next( @@ -1043,15 +1067,27 @@ def fake_run(command, **kwargs): if argument.startswith("--container-name=") ) assert f"--container-image={image.resolve()}" in calls[0][0] - for command, kwargs in calls[1:3]: - assert f"--container-name={container_name}" in command - assert not any(arg.startswith("--container-image=") for arg in command) - assert "--no-container-mount-home" in command + assert f"--container-name={container_name}" in worker_command + assert not any(arg.startswith("--container-image=") for arg in worker_command) + assert "--no-container-mount-home" in worker_command + assert "PAGER=cat" in worker_command + assert worker_command[-2:] == ["bash", "-c"] + assert "--kill-child" in worker_command + assert "/.mlperf_persistent_exec/pyxis_command_worker.sh" in worker_command + assert worker_script == ( + Path(pyxis_env_mod.__file__).with_name("pyxis_command_worker.sh").read_bytes() + ) + for _command, kwargs in calls: assert kwargs["env"].get("OPENAI_API_KEY") is None - assert calls[1][0][-5:] == ["env", "PAGER=cat", "bash", "-c", "touch state"] - assert any( - "unshare --pid --fork --mount-proc" in argument for argument in calls[1][0] + assert len(calls) == 2 # Container initialization and cleanup. + assert pyxis_channel.execute.call_count == 2 + pyxis_channel.execute.assert_any_call( + command="touch state", cwd="/testbed", timeout_s=30 + ) + pyxis_channel.execute.assert_any_call( + command="test -f state", cwd="/testbed", timeout_s=30 ) + pyxis_channel.close.assert_called_once_with() assert calls[-1][0][-4:] == [ "enroot", "remove", @@ -1062,7 +1098,9 @@ def fake_run(command, **kwargs): assert "Executing Pyxis command: touch state" in caplog.text -def test_pyxis_environment_mounts_persistent_tmp_on_every_step(monkeypatch, tmp_path): +def test_pyxis_environment_shares_private_tmp_with_worker( + monkeypatch, tmp_path, pyxis_channel +): monkeypatch.setenv("SLURM_JOB_ID", "1738605") monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") image = tmp_path / "task.sqsh" @@ -1078,33 +1116,31 @@ def fake_run(command, **kwargs): environment = PyxisEnvironment(image=image, run_id="run-1") environment.execute({"command": "touch /tmp/state"}) environment.execute({"command": "test -f /tmp/state"}) + assert len(calls) == 1 + worker_command = environment._persistent_server_command() tmp_mounts = [ next(arg for arg in command if arg.startswith("--container-mounts=")) - for command in calls[:3] + for command in [calls[0], worker_command] ] - assert tmp_mounts[0] == tmp_mounts[1] == tmp_mounts[2] + assert tmp_mounts[1].startswith(tmp_mounts[0] + ",") source, destination = ( tmp_mounts[0].removeprefix("--container-mounts=").split(":", 1) ) assert destination == "/tmp" persistent_tmp = Path(source) assert persistent_tmp.is_dir() - assert stat.S_IMODE(persistent_tmp.stat().st_mode) == 0o1777 + assert stat.S_IMODE(persistent_tmp.stat().st_mode) == 0o700 environment.cleanup() assert not persistent_tmp.exists() -@pytest.mark.parametrize( - "preamble", - [ - "", - "srun: lua: Checking requeue policy with options:\n", - ], -) -def test_pyxis_environment_extracts_submission(monkeypatch, tmp_path, preamble): +@pytest.mark.parametrize("preamble", ["", "\n "]) +def test_pyxis_environment_extracts_submission( + monkeypatch, tmp_path, pyxis_channel, preamble +): class Submitted(Exception): pass @@ -1113,74 +1149,45 @@ class Submitted(Exception): exceptions.Submitted = Submitted monkeypatch.setitem(sys.modules, "minisweagent", minisweagent) monkeypatch.setitem(sys.modules, "minisweagent.exceptions", exceptions) - monkeypatch.setenv("SLURM_JOB_ID", "1738605") - monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") - calls = 0 - - def fake_run(command, **kwargs): - nonlocal calls - calls += 1 - output = ( - "ok\n" - if calls == 1 - else ( - f"{preamble}COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT\n" - "diff --git a/a b/a\n" - ) - ) - _finish_srun_step(command, 0) - return subprocess.CompletedProcess(command, 0, stdout=output, stderr="") - - monkeypatch.setattr(subprocess, "run", fake_run) - environment = PyxisEnvironment(image=tmp_path / "task.sqsh", run_id="run-1") + environment = object.__new__(PyxisEnvironment) + environment.config = pyxis_env_mod.PyxisEnvironmentConfig( + image="image", run_id="run" + ) + environment._persistent_channel = pyxis_channel + monkeypatch.setattr(environment, "cleanup", lambda: None) + pyxis_channel.execute.return_value = CommandResult( + 0, + f"{preamble}COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT\ndiff --git a/a b/a\n", + False, + ) with pytest.raises(Submitted) as exc_info: environment.execute({"command": "submit"}) assert exc_info.value.args[0]["extra"]["submission"] == "diff --git a/a b/a\n" - environment.cleanup() - - -def test_pyxis_environment_decodes_timeout_output(monkeypatch, tmp_path): - monkeypatch.setenv("SLURM_JOB_ID", "1738605") - monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") - calls = 0 - def fake_run(command, **kwargs): - nonlocal calls - calls += 1 - if calls == 2: - _finish_srun_step(command, 124) - return subprocess.CompletedProcess( - command, 124, stdout="partial�", stderr="" - ) - _finish_srun_step(command, 0) - return subprocess.CompletedProcess(command, 0, stdout="ok\n", stderr="") - monkeypatch.setattr(subprocess, "run", fake_run) - environment = PyxisEnvironment(image=tmp_path / "task.sqsh", run_id="run-1") +def test_pyxis_environment_decodes_timeout_output(monkeypatch, pyxis_channel): + environment = object.__new__(PyxisEnvironment) + environment.config = pyxis_env_mod.PyxisEnvironmentConfig( + image="image", run_id="run" + ) + environment._persistent_channel = pyxis_channel + monkeypatch.setattr(environment, "cleanup", lambda: None) + pyxis_channel.execute.return_value = CommandResult(124, "partial�", True) - output = environment.execute({"command": "sleep 60"}) + output = environment.execute({"command": "sleep 60"}, cwd="/other", timeout=5) assert output["returncode"] == -1 assert output["output"] == "partial�" assert output["extra"]["exception_type"] == "TimeoutExpired" - environment.cleanup() + pyxis_channel.execute.assert_called_once_with( + command="sleep 60", cwd="/other", timeout_s=5 + ) -def test_pyxis_environment_raises_when_srun_never_starts_command(monkeypatch, tmp_path): +def test_pyxis_srun_step_raises_when_command_never_starts(monkeypatch, tmp_path): failure_path = tmp_path / ".pyxis_infrastructure_failure" - environment = object.__new__(PyxisEnvironment) - environment.config = types.SimpleNamespace( - cwd="/testbed", - env={}, - timeout_s=30, - interpreter=["bash", "-c"], - infrastructure_failure_path=failure_path, - ) - environment.name = "mswe_run-1_abcd1234" - environment._tmp_dir = tmp_path - monkeypatch.setenv("SLURM_JOB_ID", "1738605") monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") monkeypatch.setattr( @@ -1192,36 +1199,40 @@ def test_pyxis_environment_raises_when_srun_never_starts_command(monkeypatch, tm ) with pytest.raises(RunnerError, match="before the command completed"): - environment.execute({"command": "pytest -q"}) + pyxis_env_mod.run_srun_step( + argv=["true"], + status_path=tmp_path / ".mlperf_srun_status", + timeout_s=30, + failure_path=failure_path, + ) assert failure_path.exists() -def test_pyxis_environment_preserves_command_failure(monkeypatch, tmp_path): - monkeypatch.setenv("SLURM_JOB_ID", "1738605") - monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") - calls = 0 - - def fake_run(command, **kwargs): - nonlocal calls - calls += 1 - returncode = 0 if calls == 1 else 1 - _finish_srun_step(command, returncode) - return subprocess.CompletedProcess( - command, returncode, stdout="command failed\n", stderr="" - ) - - monkeypatch.setattr(subprocess, "run", fake_run) - environment = PyxisEnvironment(image=tmp_path / "task.sqsh", run_id="run-1") +@pytest.mark.parametrize("returncode", [1, 124, 137]) +def test_pyxis_environment_preserves_command_failure( + monkeypatch, pyxis_channel, returncode +): + environment = object.__new__(PyxisEnvironment) + environment.config = pyxis_env_mod.PyxisEnvironmentConfig( + image="image", run_id="run" + ) + environment._persistent_channel = pyxis_channel + monkeypatch.setattr(environment, "cleanup", lambda: None) + pyxis_channel.execute.return_value = CommandResult( + returncode, "command failed\n", False + ) - output = environment.execute({"command": "false"}) + output = environment.execute({"command": f"exit {returncode}"}) - assert output["returncode"] == 1 + assert output["returncode"] == returncode assert output["output"] == "command failed\n" - environment.cleanup() + assert output["exception_info"] == "" -def test_pyxis_cleanup_is_best_effort_outside_allocation(monkeypatch, tmp_path): +def test_pyxis_cleanup_is_best_effort_outside_allocation( + monkeypatch, tmp_path, pyxis_channel +): monkeypatch.setenv("SLURM_JOB_ID", "1738605") monkeypatch.setenv("SLURMD_NODENAME", "gb-nvl-053-compute04") @@ -1238,6 +1249,7 @@ def fake_run(command, **kwargs): environment.cleanup() + pyxis_channel.close.assert_called_once_with() assert not persistent_tmp.exists() diff --git a/uv.lock b/uv.lock index 98a669bc..a8e011a8 100644 --- a/uv.lock +++ b/uv.lock @@ -1483,9 +1483,9 @@ requires-dist = [ { name = "tiktoken", specifier = "==0.13.0" }, { name = "transformers", specifier = "==5.16.1" }, { name = "typing-extensions", specifier = "==4.15.0" }, - { name = "urllib3", specifier = "==2.7.0" }, + { name = "urllib3", specifier = "==2.8.0" }, { name = "uvloop", specifier = "==0.22.1" }, - { name = "virtualenv", marker = "extra == 'dev'", specifier = ">=20.36.1" }, + { name = "virtualenv", marker = "extra == 'dev'", specifier = "==21.7.13" }, { name = "websocket-client", specifier = "==1.9.0" }, ] provides-extras = ["sql", "dev", "test", "performance", "rouge", "bfcl"] @@ -2981,15 +2981,14 @@ wheels = [ [[package]] name = "python-discovery" -version = "1.2.1" +version = "1.6.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "filelock", version = "3.25.2", source = { registry = "https://pypi.org/simple" }, marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'x86_64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, - { name = "platformdirs", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'x86_64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/b9/88/815e53084c5079a59df912825a279f41dd2e0df82281770eadc732f5352c/python_discovery-1.2.1.tar.gz", hash = "sha256:180c4d114bff1c32462537eac5d6a332b768242b76b69c0259c7d14b1b680c9e", size = 58457, upload-time = "2026-03-26T22:30:44.496Z" } +sdist = { url = "https://files.pythonhosted.org/packages/0c/57/250bd238b966cece44328235eb85290045d059265fdaf7527a3a958123db/python_discovery-1.6.1.tar.gz", hash = "sha256:cf87d3627dfb4412437fdd5b13eae402607722998d21567993aedbc59b23c15e", size = 84338, upload-time = "2026-09-18T01:31:53.971Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/67/0f/019d3949a40280f6193b62bc010177d4ce702d0fce424322286488569cd3/python_discovery-1.2.1-py3-none-any.whl", hash = "sha256:b6a957b24c1cd79252484d3566d1b49527581d46e789aaf43181005e56201502", size = 31674, upload-time = "2026-03-26T22:30:43.396Z" }, + { url = "https://files.pythonhosted.org/packages/16/7d/e9ffbadfbf89c93848412d04594135c4ae8c1d37d9e053b9c3ed718fabc4/python_discovery-1.6.1-py3-none-any.whl", hash = "sha256:d43fcdef879fe795352bd13ccf8d185ba5a9f86f36cfcd00529f596e737442b3", size = 38664, upload-time = "2026-09-18T01:31:52.448Z" }, ] [[package]] @@ -4085,11 +4084,11 @@ wheels = [ [[package]] name = "urllib3" -version = "2.7.0" +version = "2.8.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602, upload-time = "2026-05-07T16:13:18.596Z" } +sdist = { url = "https://files.pythonhosted.org/packages/e3/05/b17359e1cefb4f909b5e40b1b90a496d987258916dbbf88e842c729f510e/urllib3-2.8.0.tar.gz", hash = "sha256:63bf2ead4c879426ebf22ef2a781eeb4aa3b4ae798a0435506f8687fd5bb9b63", size = 458972, upload-time = "2026-09-15T19:29:36.253Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087, upload-time = "2026-05-07T16:13:17.151Z" }, + { url = "https://files.pythonhosted.org/packages/92/9d/c4e665119135114480843e7ab388fa94d8480650450e6f8e26b70d323a4c/urllib3-2.8.0-py3-none-any.whl", hash = "sha256:0cf3cae568d36aa9576b28dfb35f11328f1cb974ca7647d9475ebb86c75ac6e3", size = 135717, upload-time = "2026-09-15T19:29:34.577Z" }, ] [[package]] @@ -4126,7 +4125,7 @@ wheels = [ [[package]] name = "virtualenv" -version = "21.2.0" +version = "21.7.13" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "distlib", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'x86_64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, @@ -4134,9 +4133,9 @@ dependencies = [ { name = "platformdirs", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'x86_64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, { name = "python-discovery", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'x86_64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux')" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/aa/92/58199fe10049f9703c2666e809c4f686c54ef0a68b0f6afccf518c0b1eb9/virtualenv-21.2.0.tar.gz", hash = "sha256:1720dc3a62ef5b443092e3f499228599045d7fea4c79199770499df8becf9098", size = 5840618, upload-time = "2026-03-09T17:24:38.013Z" } +sdist = { url = "https://files.pythonhosted.org/packages/13/50/c9b84eb106d0db420b9878dac0a386c5791726f127ce20b06897f6e3a1e9/virtualenv-21.7.13.tar.gz", hash = "sha256:0355558b6f33619aab31347e43643b0ebc97f61ea3acf617b2b69e1f8a843d11", size = 5360306, upload-time = "2026-09-18T04:35:49.354Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/c6/59/7d02447a55b2e55755011a647479041bc92a82e143f96a8195cb33bd0a1c/virtualenv-21.2.0-py3-none-any.whl", hash = "sha256:1bd755b504931164a5a496d217c014d098426cddc79363ad66ac78125f9d908f", size = 5825084, upload-time = "2026-03-09T17:24:35.378Z" }, + { url = "https://files.pythonhosted.org/packages/9a/ce/e74453531b0c49a58d0e8a4f9bab4495705859fee4a4f7de27d58f4a791a/virtualenv-21.7.13-py3-none-any.whl", hash = "sha256:1bea5af7463f59c4719db48fe739579a2a4f569c96f26c086edda85c96da9f59", size = 5328680, upload-time = "2026-09-18T04:35:47.194Z" }, ] [[package]]