diff --git a/.github/workflows/unit-tests.yml b/.github/workflows/unit-tests.yml index 85124757a4..caee7b7296 100644 --- a/.github/workflows/unit-tests.yml +++ b/.github/workflows/unit-tests.yml @@ -135,7 +135,7 @@ jobs: curl -LsSf https://astral.sh/uv/install.sh | sh uv venv --python 3.12 source .venv/bin/activate - uv sync --extra dev + uv sync --extra dev --extra sandbox - name: Test if: steps.changes.outputs.run_full == 'true' || steps.changes.outputs.run_servers == 'true' diff --git a/nemo_gym/sandbox/__init__.py b/nemo_gym/sandbox/__init__.py new file mode 100644 index 0000000000..4754ff483a --- /dev/null +++ b/nemo_gym/sandbox/__init__.py @@ -0,0 +1,51 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Public sandbox API for NeMo Gym.""" + +from nemo_gym.sandbox.api import AsyncSandbox, Sandbox +from nemo_gym.sandbox.providers import ( + ExecResult, + SandboxCreateError, + SandboxCreateVerificationError, + SandboxExecResult, + SandboxHandle, + SandboxProvider, + SandboxSpec, + SandboxStatus, + create_provider, + get_provider_class, + list_providers, + register_provider, +) +from nemo_gym.sandbox.utils import rewrite_image + + +__all__ = [ + "Sandbox", + "AsyncSandbox", + "ExecResult", + "SandboxCreateError", + "SandboxCreateVerificationError", + "SandboxExecResult", + "SandboxHandle", + "SandboxProvider", + "SandboxSpec", + "SandboxStatus", + "create_provider", + "get_provider_class", + "list_providers", + "register_provider", + "rewrite_image", +] diff --git a/nemo_gym/sandbox/api.py b/nemo_gym/sandbox/api.py new file mode 100644 index 0000000000..1d58c931c7 --- /dev/null +++ b/nemo_gym/sandbox/api.py @@ -0,0 +1,294 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Provider-neutral public sandbox API.""" + +import asyncio +import tempfile +import threading +from collections.abc import Awaitable, Callable, Mapping +from concurrent.futures import Future +from pathlib import Path +from typing import Any, TypeVar + +from nemo_gym.sandbox.providers import ( + SandboxExecResult, + SandboxHandle, + SandboxProvider, + SandboxSpec, + SandboxStatus, + create_provider, +) + + +T = TypeVar("T") + + +class AsyncSandbox: + """Async sandbox object backed by a runtime provider.""" + + def __init__( + self, + provider: Mapping[str, Any] | SandboxProvider, + spec: SandboxSpec | None = None, + *, + delete_on_stop: bool = False, + ) -> None: + self._provider = create_provider(provider) if isinstance(provider, Mapping) else provider + self._spec = spec + self._handle: SandboxHandle | None = None + self._delete_on_stop = delete_on_stop + self._stopped = True + self._closed = False + + def _require_handle(self) -> SandboxHandle: + if self._handle is None or self._stopped: + raise RuntimeError("Sandbox has not been started") + return self._handle + + async def _write_inline_file(self, handle: SandboxHandle, target_path: str, data: str | bytes) -> None: + with tempfile.TemporaryDirectory(prefix="nemo-gym-sandbox-upload-") as tmp_dir: + source_path = Path(tmp_dir) / "contents" + if isinstance(data, str): + source_path.write_text(data, encoding="utf-8") + else: + source_path.write_bytes(data) + await self._provider.upload_file(handle, source_path, target_path) + + async def _write_initial_files(self, handle: SandboxHandle, files: dict[str, str]) -> None: + for target_path, contents in files.items(): + await self._write_inline_file(handle, target_path, contents) + + async def start( + self, + spec: SandboxSpec | None = None, + *, + delete_on_stop: bool | None = None, + ) -> "AsyncSandbox": + if self._closed: + raise RuntimeError("Sandbox has been stopped") + if self._handle is not None and not self._stopped: + raise RuntimeError("Sandbox is already started") + requested_spec = spec if spec is not None else self._spec + if requested_spec is None: + raise ValueError("Sandbox.start() requires a SandboxSpec") + + handle = await self._provider.create(requested_spec) + try: + await self._write_initial_files(handle, requested_spec.files) + except Exception: + await self._provider.close(handle, delete=True) + await self._provider.aclose() + self._closed = True + raise + + self._spec = requested_spec + self._handle = handle + self._delete_on_stop = self._delete_on_stop if delete_on_stop is None else delete_on_stop + self._stopped = False + return self + + async def exec( + self, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = 180, + user: str | int | None = None, + ) -> SandboxExecResult: + return await self._provider.exec( + self._require_handle(), + command, + cwd=cwd if cwd is not None else self._spec.workdir if self._spec is not None else None, + env=env, + timeout_s=timeout_s, + user=user, + ) + + async def upload(self, local_path: Path | str, remote_path: str) -> None: + await self._provider.upload_file(self._require_handle(), Path(local_path), remote_path) + + async def download(self, remote_path: str, local_path: Path | str) -> None: + await self._provider.download_file(self._require_handle(), remote_path, Path(local_path)) + + async def status(self) -> SandboxStatus: + if self._handle is None: + return SandboxStatus.UNKNOWN + if self._stopped: + return SandboxStatus.STOPPED + return await self._provider.status(self._handle) + + async def stop(self, *, delete: bool | None = None) -> None: + if self._closed: + return + try: + if self._handle is not None and not self._stopped: + self._stopped = True + await self._provider.close( + self._handle, + delete=self._delete_on_stop if delete is None else delete, + ) + finally: + await self._provider.aclose() + self._closed = True + + async def __aenter__(self) -> "AsyncSandbox": + return self + + async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: + await self.stop() + + +class _AsyncLoopRunner: + """Run async sandbox operations for sync callers.""" + + def __init__(self) -> None: + self._loop = asyncio.new_event_loop() + self._ready = threading.Event() + self._closed = False + self._thread = threading.Thread(target=self._run_loop, name="nemo-gym-sandbox-sync-loop", daemon=True) + self._thread.start() + self._ready.wait() + + def _run_loop(self) -> None: + asyncio.set_event_loop(self._loop) + self._ready.set() + self._loop.run_forever() + + def _ensure_can_block(self, operation: str) -> None: + if self._closed or self._loop.is_closed(): + raise RuntimeError("Sandbox sync loop is closed") + try: + asyncio.get_running_loop() + except RuntimeError: + return + raise RuntimeError(f"Sandbox.{operation}() is blocking; use AsyncSandbox in async code instead.") + + def call(self, operation: str, func: Callable[[], T]) -> T: + self._ensure_can_block(operation) + future: Future[T] = Future() + + def invoke() -> None: + try: + future.set_result(func()) + except BaseException as e: + future.set_exception(e) + + self._loop.call_soon_threadsafe(invoke) + return future.result() + + def run(self, operation: str, awaitable_factory: Callable[[], Awaitable[T]]) -> T: + self._ensure_can_block(operation) + future = asyncio.run_coroutine_threadsafe(awaitable_factory(), self._loop) + return future.result() + + def close(self) -> None: + if self._closed: + return + self._closed = True + if not self._loop.is_closed(): + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(timeout=5) + self._loop.close() + + +class Sandbox: + """Synchronous wrapper around ``AsyncSandbox``.""" + + def __init__( + self, + provider: Mapping[str, Any] | SandboxProvider, + spec: SandboxSpec | None = None, + *, + delete_on_stop: bool = False, + ) -> None: + self._runner = _AsyncLoopRunner() + try: + self._async_sandbox = self._runner.call( + "__init__", + lambda: AsyncSandbox(provider, spec, delete_on_stop=delete_on_stop), + ) + except BaseException: + self._runner.close() + raise + self._closed = False + + def start( + self, + spec: SandboxSpec | None = None, + *, + delete_on_stop: bool | None = None, + ) -> "Sandbox": + self._runner.run( + "start", + lambda: self._async_sandbox.start( + spec, + delete_on_stop=delete_on_stop, + ), + ) + return self + + def exec( + self, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = 180, + user: str | int | None = None, + ) -> SandboxExecResult: + return self._runner.run( + "exec", + lambda: self._async_sandbox.exec( + command, + cwd=cwd, + env=env, + timeout_s=timeout_s, + user=user, + ), + ) + + def upload(self, local_path: Path | str, remote_path: str) -> None: + self._runner.run("upload", lambda: self._async_sandbox.upload(local_path, remote_path)) + + def download(self, remote_path: str, local_path: Path | str) -> None: + self._runner.run("download", lambda: self._async_sandbox.download(remote_path, local_path)) + + def status(self) -> SandboxStatus: + if self._closed: + return SandboxStatus.STOPPED + return self._runner.run("status", self._async_sandbox.status) + + def stop(self, *, delete: bool | None = None) -> None: + if self._closed: + return + self._closed = True + try: + self._runner.run("stop", lambda: self._async_sandbox.stop(delete=delete)) + finally: + self._runner.close() + + def __enter__(self) -> "Sandbox": + return self + + def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: + self.stop() + + def __del__(self) -> None: # pragma: no cover + if hasattr(self, "_closed") and not self._closed: + try: + self.stop() + except Exception: + pass diff --git a/nemo_gym/sandbox/providers/__init__.py b/nemo_gym/sandbox/providers/__init__.py new file mode 100644 index 0000000000..3614eac34f --- /dev/null +++ b/nemo_gym/sandbox/providers/__init__.py @@ -0,0 +1,48 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Sandbox provider registry.""" + +from nemo_gym.sandbox.providers.base import ( + ExecResult, + SandboxCreateError, + SandboxCreateVerificationError, + SandboxExecResult, + SandboxHandle, + SandboxProvider, + SandboxSpec, + SandboxStatus, +) +from nemo_gym.sandbox.providers.registry import ( + create_provider, + get_provider_class, + list_providers, + register_provider, +) + + +__all__ = [ + "ExecResult", + "SandboxCreateError", + "SandboxCreateVerificationError", + "SandboxExecResult", + "SandboxHandle", + "SandboxProvider", + "SandboxSpec", + "SandboxStatus", + "create_provider", + "get_provider_class", + "list_providers", + "register_provider", +] diff --git a/nemo_gym/sandbox/providers/base.py b/nemo_gym/sandbox/providers/base.py new file mode 100644 index 0000000000..f4629d9871 --- /dev/null +++ b/nemo_gym/sandbox/providers/base.py @@ -0,0 +1,136 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Provider-facing sandbox protocol.""" + +from dataclasses import dataclass, field +from enum import Enum +from pathlib import Path +from typing import Any, Protocol + + +class SandboxStatus(str, Enum): + """Provider-neutral sandbox lifecycle status.""" + + STARTING = "starting" + RUNNING = "running" + STOPPED = "stopped" + ERROR = "error" + UNKNOWN = "unknown" + + +@dataclass(frozen=True) +class SandboxSpec: + """Sandbox creation request.""" + + image: str | None = None + timeout_s: int | float | None = None + ready_timeout_s: int | float | None = None + workdir: str | None = None + env: dict[str, str] = field(default_factory=dict) + files: dict[str, str] = field(default_factory=dict) + metadata: dict[str, str] = field(default_factory=dict) + resources: dict[str, str] = field(default_factory=dict) + entrypoint: list[str] | None = None + provider_options: dict[str, Any] = field(default_factory=dict) + + +@dataclass +class SandboxHandle: + """Provider-neutral handle to a created sandbox. + + ``raw`` is provider-owned opaque state. Public code should pass it back to + the provider through this handle rather than inspecting or mutating it + directly. + """ + + sandbox_id: str + provider_name: str + raw: Any + + +@dataclass(frozen=True) +class SandboxExecResult: + """Provider-neutral process execution result. + + ``return_code`` is the process exit code when the sandbox actually ran the + command. Providers may use a non-process sentinel with ``error_type`` set + when the sandbox runtime reports an execution failure without a process + exit code. + """ + + stdout: str | None + stderr: str | None + return_code: int + error_type: str | None = None + + +ExecResult = SandboxExecResult + + +class SandboxCreateError(RuntimeError): + """Raised when a provider cannot create a sandbox.""" + + +class SandboxCreateVerificationError(SandboxCreateError): + """Raised when a newly-created sandbox fails provider readiness checks.""" + + +class SandboxProvider(Protocol): + """Runtime/infra provider contract used by the public sandbox API.""" + + name: str + + async def create(self, spec: SandboxSpec) -> SandboxHandle: + """Create a ready sandbox and return a provider-neutral handle. + + Providers must return only after the sandbox is healthy enough to run + commands and transfer files. If the sandbox cannot become ready before + the configured timeout, providers should raise ``SandboxCreateError`` + or a provider-specific subclass. + """ + ... + + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + """Run a command inside a sandbox.""" + ... + + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + """Upload one local file into a sandbox.""" + ... + + async def download_file(self, handle: SandboxHandle, source_path: str, target_path: Path) -> None: + """Download one sandbox file to the local filesystem.""" + ... + + async def status(self, handle: SandboxHandle) -> SandboxStatus: + """Return the current sandbox lifecycle status.""" + ... + + async def close(self, handle: SandboxHandle, *, delete: bool = False) -> None: + """Close provider resources and optionally delete the sandbox.""" + ... + + async def aclose(self) -> None: + """Close provider-scoped resources such as SDK clients.""" + ... diff --git a/nemo_gym/sandbox/providers/ecs_fargate/README.md b/nemo_gym/sandbox/providers/ecs_fargate/README.md new file mode 100644 index 0000000000..5010060028 --- /dev/null +++ b/nemo_gym/sandbox/providers/ecs_fargate/README.md @@ -0,0 +1,79 @@ +# ECS Fargate sandbox provider + +Runs each `nemo_gym.sandbox` sandbox as an AWS ECS Fargate task behind an SSH +sidecar. It implements the provider-neutral `SandboxProvider` contract, so any +sandbox-backed agent (not just mini-swe-agent) can use it by setting +`sandbox_provider.ecs_fargate` in its config. + +## Prerequisites + +- **Infrastructure** provisioned in the target account/region. The reference + Terraform stack publishes its outputs to SSM at + `//ecs-sandbox/config` (`ssm_project` defaults to `harbor`): + cluster, subnets, security groups, task/execution roles, ECR mirror, EFS, and + the SSH-sidecar key ARNs. A missing parameter raises an actionable error. +- **Credentials**: `AWS_ACCESS_KEY_ID`/`AWS_SECRET_ACCESS_KEY` (or an instance + role) plus `AWS_REGION`. +- **Network**: the host must reach each task's SSH sidecar port (`52222`) — run + inside the sandbox VPC/peered network, or allow the host IP on the sidecar + security group. Exec, file transfer, and model reverse-tunnels ride this SSH + connection. + +## Configuration + +Region-only is enough; everything else auto-discovers from SSM (explicit YAML +always wins): + +```yaml +sandbox_provider: + ecs_fargate: + region: ${oc.env:AWS_REGION} + cpu: "2048" + memory: "8192" + ephemeral_storage_gib: 50 +``` + +Common fields: + +| Field | Default | Purpose | +| --- | --- | --- | +| `region` | — | AWS region; enables SSM auto-discovery when `cluster` is omitted | +| `cpu` / `memory` | `"4096"` / `"8192"` | Fargate task size (CPU units / MiB) | +| `ephemeral_storage_gib` | task default | Task ephemeral disk | +| `auto_mirror` | `true` | Pull a missing public image into the ECR mirror on demand (see below) | +| `ssm_project` | `harbor` | SSM namespace for auto-discovery | +| `environment_dir` | — | Build a task image from a Dockerfile dir via CodeBuild instead of using a prebuilt image | + +Per-sandbox `ready_timeout_s`, `env`, `files`, `metadata`, and +`provider_options` (e.g. `platform`, `outside_endpoints`) come from the +`SandboxSpec`. + +## Images and on-demand mirroring + +ECS pulls task images from the account ECR mirror, not their origin registry. A +bare/public image (e.g. `docker.io/swebench/sweb.eval.x86_64.:latest`) +resolves to the mirror tag `:`. Resolution order: + +1. `environment_dir` set → build the image via CodeBuild and use it. +2. Image is already an ECR reference → use verbatim (never re-mirrored). +3. Bare/public name + `auto_mirror=true` → mirror into ECR on demand (CodeBuild + pull → retag → push) during `create`, then launch. Concurrent tasks for the + same image de-duplicate onto one build. + +Set `auto_mirror: false` to require a pre-populated mirror and fail fast on a +miss. The first task for a new image waits on a one-time build (~1–3 min for +typical SWE-bench images); later tasks hit the ECR cache. + +## Lifecycle + +`create` launches the task + SSH sidecar and returns once the exec server is +healthy. `exec`, `upload_file`, `download_file`, `status`, and `close` delegate +to the per-sandbox engine over the SSH tunnel. `outside_endpoints` (via +`spec.provider_options`) open reverse tunnels so an in-sandbox process can reach +a host-side endpoint (e.g. a model server). + +## Security + +The sidecar security group in the reference stack allows `0.0.0.0/0` on `52222` +for convenience. Restrict it to the orchestrator's egress IP (or move to a +private/Teleport path) before non-smoke use. diff --git a/nemo_gym/sandbox/providers/ecs_fargate/__init__.py b/nemo_gym/sandbox/providers/ecs_fargate/__init__.py new file mode 100644 index 0000000000..a37bc4d8b2 --- /dev/null +++ b/nemo_gym/sandbox/providers/ecs_fargate/__init__.py @@ -0,0 +1,32 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""ECS Fargate sandbox provider package.""" + +from nemo_gym.sandbox.providers.ecs_fargate.engine import ( + EcsFargateConfig, + SshSidecarConfig, +) +from nemo_gym.sandbox.providers.ecs_fargate.provider import ( + EcsFargateProvider, + engine_config_from_mapping, +) + + +__all__ = [ + "EcsFargateConfig", + "EcsFargateProvider", + "SshSidecarConfig", + "engine_config_from_mapping", +] diff --git a/nemo_gym/sandbox/providers/ecs_fargate/engine.py b/nemo_gym/sandbox/providers/ecs_fargate/engine.py new file mode 100644 index 0000000000..49d19609c1 --- /dev/null +++ b/nemo_gym/sandbox/providers/ecs_fargate/engine.py @@ -0,0 +1,2376 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""ECS Fargate sandbox engine. + +Lifted from nemo-evaluator-next (``nemo_evaluator/sandbox/ecs_fargate.py``). +The orchestration is unchanged; only the three host-protocol types it relied on +(``ExecResult``, ``OutsideEndpoint``, ``SandboxSpec``) and the SSM/port +constants are vendored here so the engine stands alone inside the provider. +The Gym-facing adapter lives in ``provider.py``. +""" + +from __future__ import annotations + +import asyncio +import atexit +import base64 +import hashlib +import io +import json +import logging +import os +import random +import re +import shlex +import socket +import subprocess +import tarfile +import tempfile +import threading +import time +import uuid +import zipfile +from dataclasses import dataclass, field +from dataclasses import replace as _dc_replace +from pathlib import Path +from typing import Any, Callable, Self, TypeVar +from urllib.parse import ParseResult, urlparse + +import aiohttp + + +# Matches the exec-server script's TB_EXEC_PORT fallback and the +# exec_server_port written by the reference Terraform. +DEFAULT_EXEC_SERVER_PORT = 19542 +# SSH sidecar port matching the reference Terraform ssh_tunnel_sshd_port. +DEFAULT_SSHD_PORT = 52222 +DEFAULT_SSM_PROJECT = "harbor" + + +@dataclass +class ExecResult: + """Result of executing a command inside a sandbox.""" + + stdout: str + stderr: str + return_code: int + + +@dataclass +class OutsideEndpoint: + """A host-side URL that must be reachable from inside the sandbox. + + The sandbox rewrites *url* for its network topology and exposes the + resolved address as the environment variable *env_var* inside the + container. + """ + + url: str + env_var: str + + +@dataclass +class VolumeMount: + """EFS mount (ECS Fargate). Host bind mounts are unused on Fargate.""" + + host_path: str = "" + container_path: str = "" + readonly: bool = False + efs_filesystem_id: str | None = None + efs_root_directory: str | None = None + efs_access_point_id: str | None = None + + @property + def is_efs(self) -> bool: + return self.efs_filesystem_id is not None + + +@dataclass +class SandboxSpec: + """Per-problem sandbox requirements.""" + + image: str + workdir: str = "/workspace" + env: dict[str, str] = field(default_factory=dict) + files: dict[str, str] = field(default_factory=dict) + entrypoint: str | None = None + volumes: list[VolumeMount] = field(default_factory=list) + environment_dir: str | None = None + + +logger = logging.getLogger(__name__) +T = TypeVar("T") + + +# ── Lazy AWS SDK import ────────────────────────────────────────────── + + +def _require_aws_sdks(): + try: + import importlib + + boto3 = importlib.import_module("boto3") + botocore_config = importlib.import_module("botocore.config") + botocore_exceptions = importlib.import_module("botocore.exceptions") + except ModuleNotFoundError as e: + raise RuntimeError( + "ECS Fargate sandbox requires boto3/botocore. " + "Install them (`pip install boto3`) or use a different sandbox backend." + ) from e + return boto3, getattr(botocore_config, "Config"), getattr(botocore_exceptions, "ClientError") + + +# ── SSM auto-discovery ─────────────────────────────────────────────── + +_ssm_config_cache: dict[str, dict[str, Any]] = {} + + +def resolve_ecs_config_from_ssm( + region: str, + project: str = DEFAULT_SSM_PROJECT, +) -> dict[str, Any]: + """Read ECS sandbox config from SSM Parameter Store. + + Returns a dict matching the JSON structure written by Terraform + (cluster, subnets, security_groups, roles, SSH ARNs, EFS, etc.). + Results are cached per (region, project) for the process lifetime. + """ + cache_key = f"{region}:{project}" + if cache_key in _ssm_config_cache: + return _ssm_config_cache[cache_key] + + boto3, _, ClientError = _require_aws_sdks() + ssm = boto3.client("ssm", region_name=region) + param_name = f"/{project}/ecs-sandbox/config" + try: + resp = _retry_with_backoff( + lambda: ssm.get_parameter(Name=param_name), + operation_name="ssm.get_parameter", + max_retries=5, + ) + except ClientError as exc: + code = (exc.response.get("Error") or {}).get("Code", "") + if code == "ParameterNotFound": + raise RuntimeError( + f"SSM parameter '{param_name}' not found in {region}. " + f"Run 'terraform apply' in the ecs-sandbox stack for this " + f"region, or specify all ECS fields explicitly in your YAML." + ) from exc + raise + + raw = resp["Parameter"]["Value"] + try: + config = json.loads(raw) + except json.JSONDecodeError as exc: + raise RuntimeError(f"SSM parameter '{param_name}' in {region} contains invalid JSON: {exc}") from exc + + _ssm_config_cache[cache_key] = config + logger.info( + "Resolved ECS config from SSM %s in %s (cluster=%s, %d subnets)", + param_name, + region, + config.get("cluster"), + len(config.get("subnets", [])), + ) + return config + + +# ── Config dataclasses ─────────────────────────────────────────────── + + +def _sanitize_id(value: str, max_len: int = 100) -> str: + cleaned = re.sub(r"[^a-zA-Z0-9-]+", "-", value).strip("-") + return cleaned[:max_len] or "task" + + +_ECR_IMAGE_REF_RE = re.compile(r"^[0-9]{12}\.dkr\.ecr\.[a-z0-9-]+\.amazonaws\.com/", re.IGNORECASE) + + +def _is_ecr_image_ref(image: str) -> bool: + """True for references already pointing at an ECR registry host. + + Such references (e.g. an image resolved against the configured ECR mirror) + must be used as-is; reapplying the ``ecr_repository`` + sanitize rewrite + would corrupt the tag. + """ + return bool(_ECR_IMAGE_REF_RE.match(image)) + + +@dataclass(frozen=True) +class SshSidecarConfig: + """SSH sidecar container configuration. + + exec_server_port set → exec-server mode (one-way tunnel). + exec_server_port None → agent-server mode (two-way tunnel). + """ + + sshd_port: int = 2222 + ssh_ready_timeout_sec: float = 300.0 + public_key_secret_arn: str = "" + private_key_secret_arn: str = "" + image: str | None = None + exec_server_port: int | None = None + + +@dataclass(frozen=True) +class EcsFargateConfig: + """Configuration for the ECS Fargate sandbox.""" + + region: str | None = None + cluster: str = "" + subnets: list[str] = field(default_factory=list) + security_groups: list[str] = field(default_factory=list) + assign_public_ip: bool = False + task_definition: str | None = None + task_definition_family_prefix: str = "ecs-sandbox" + image_template: str | None = None + container_name: str = "main" + container_port: int | None = None + cpu: str = "4096" + memory: str = "8192" + ephemeral_storage_gib: int | None = None + platform_version: str | None = None + execution_role_arn: str | None = None + task_role_arn: str | None = None + extra_env: dict[str, str] | None = None + log_group: str | None = None + log_stream_prefix: str | None = None + max_task_lifetime_sec: int = 14400 + startup_timeout_sec: float = 300.0 + poll_interval_sec: float = 2.0 + run_task_max_retries: int = 30 + ssh_sidecar: SshSidecarConfig | None = None + s3_bucket: str | None = None + s3_prefix: str | None = None + ecr_repository: str | None = None + environment_dir: str | None = None + codebuild_project: str | None = None + codebuild_service_role: str | None = None + codebuild_compute_type: str = "BUILD_GENERAL1_MEDIUM" + codebuild_build_timeout: int = 60 + dockerhub_secret_arn: str | None = None + build_parallelism: int = 50 + # When a public/bare image is routed to the ECR mirror and is not yet + # present, pull it into ECR on demand (via CodeBuild) during create. The + # mirror normally runs ahead of time, but this keeps create self-healing so + # callers never have to pre-stage images manually. Set False to require a + # pre-populated mirror and fail fast on a miss. + auto_mirror: bool = True + efs_filesystem_id: str | None = None + efs_access_point_id: str | None = None + ssm_project: str = DEFAULT_SSM_PROJECT + + +@dataclass(frozen=True) +class _OutsideEndpointRoute: + endpoint: OutsideEndpoint + source_netloc: str + host: str + target_port: int + remote_port: int + scheme: str + + @classmethod + def for_endpoint(cls, endpoint: OutsideEndpoint, *, remote_port: int | None = None) -> _OutsideEndpointRoute: + parsed = urlparse(endpoint.url) + host = parsed.hostname + if not host: + raise ValueError(f"Cannot resolve hostname from OutsideEndpoint: {endpoint.url}") + target_port = _port_from_url(parsed) + return cls( + endpoint=endpoint, + source_netloc=parsed.netloc, + host=host, + target_port=target_port, + remote_port=remote_port or target_port, + scheme=parsed.scheme or "http", + ) + + def resolved_endpoint_url(self) -> str: + return self._rewrite(urlparse(self.endpoint.url)) + + def resolve_url(self, url: str) -> str: + return self._rewrite(urlparse(url)) + + def _rewrite(self, parsed: ParseResult) -> str: + return parsed._replace( + scheme=self.scheme, + netloc=f"127.0.0.1:{self.remote_port}", + ).geturl() + + +def _port_from_url(parsed: ParseResult) -> int: + return parsed.port or (443 if parsed.scheme == "https" else 80) + + +@dataclass(frozen=True) +class _OutsideEndpointRouting: + endpoints: tuple[OutsideEndpoint, ...] = () + _routes_by_env: dict[str, _OutsideEndpointRoute] = field(default_factory=dict) + _reverse_specs: tuple[str, ...] = () + _agent_tunnel_port: int | None = None + + @classmethod + def empty(cls, endpoints: list[OutsideEndpoint] | None = None) -> _OutsideEndpointRouting: + return cls(endpoints=tuple(endpoints or [])) + + @classmethod + def for_exec_server( + cls, + endpoints: list[OutsideEndpoint], + sidecar: SshSidecarConfig, + ) -> _OutsideEndpointRouting: + reverse_specs: list[str] = [] + routes_by_env: dict[str, _OutsideEndpointRoute] = {} + used_ports = {sidecar.sshd_port} + if sidecar.exec_server_port is not None: + used_ports.add(sidecar.exec_server_port) + + target_port_map: dict[tuple[str, int], int] = {} + for ep in endpoints: + target = _OutsideEndpointRoute.for_endpoint(ep) + key = (target.host, target.target_port) + remote_port = target_port_map.get(key) + if remote_port is None: + remote_port = cls._allocate_reverse_port(target.target_port, used_ports) + target_port_map[key] = remote_port + reverse_specs.append(f"{remote_port}:{target.host}:{target.target_port}") + route = _OutsideEndpointRoute.for_endpoint(ep, remote_port=remote_port) + logger.info( + "Reverse tunnel: container :%d → host %s:%d (%s)", + route.remote_port, + route.host, + route.target_port, + ep.env_var, + ) + routes_by_env[ep.env_var] = route + + return cls(endpoints=tuple(endpoints), _routes_by_env=routes_by_env, _reverse_specs=tuple(reverse_specs)) + + @classmethod + def for_agent_server(cls, endpoints: list[OutsideEndpoint]) -> _OutsideEndpointRouting: + if len(endpoints) > 1: + raise ValueError("Agent-server mode supports only one OutsideEndpoint") + if not endpoints: + raise ValueError("Agent-server mode requires OutsideEndpoint passed to start()") + route = _OutsideEndpointRoute.for_endpoint(endpoints[0]) + return cls( + endpoints=tuple(endpoints), + _routes_by_env={route.endpoint.env_var: route}, + _agent_tunnel_port=route.target_port, + ) + + @property + def reverse_specs(self) -> list[str]: + return list(self._reverse_specs) + + @property + def agent_tunnel_port(self) -> int | None: + return self._agent_tunnel_port + + def agent_tunnel_target(self) -> tuple[str, int]: + if not self.endpoints: + raise ValueError("Agent-server mode requires OutsideEndpoint passed to start()") + route = self._routes_by_env[self.endpoints[0].env_var] + return route.host, route.target_port + + def env_overrides(self) -> dict[str, str]: + return {env_var: route.resolved_endpoint_url() for env_var, route in self._routes_by_env.items()} + + def resolved_endpoint_url(self, env_var: str) -> str | None: + route = self._routes_by_env.get(env_var) + if route is None: + return None + return route.resolved_endpoint_url() + + def resolve_url(self, url: str) -> str: + parsed = urlparse(url) + for route in self._routes_by_env.values(): + if route.source_netloc == parsed.netloc: + return route.resolve_url(url) + if self._agent_tunnel_port is not None: + return parsed._replace(netloc=f"127.0.0.1:{self._agent_tunnel_port}").geturl() + raise RuntimeError("resolve_outside_endpoint() requires SSH reverse tunnel") + + @staticmethod + def _allocate_reverse_port(preferred: int, used_ports: set[int]) -> int: + if 0 < preferred <= 65535 and preferred not in used_ports: + used_ports.add(preferred) + return preferred + for candidate in range(20000, 61000): + if candidate not in used_ports: + used_ports.add(candidate) + return candidate + raise RuntimeError("No available local port for ECS reverse tunnel") + + +# ── Retry utilities ────────────────────────────────────────────────── + +_RETRYABLE_CODES = frozenset( + { + "ThrottlingException", + "TooManyRequestsException", + "ServiceUnavailable", + "RequestLimitExceeded", + } +) +_RETRYABLE_MESSAGES = ( + "capacity is unavailable", + "rate exceeded", + "too many concurrent", + "throttl", + "connect timeout", + "read timeout", + "connection reset", + "endpointconnectionerror", +) + + +def _is_retryable_error(exc: Exception) -> bool: + msg = str(exc).lower() + code = "" + if hasattr(exc, "response"): + code = (exc.response.get("Error") or {}).get("Code", "") # type: ignore[union-attr] + return code in _RETRYABLE_CODES or any(m in msg for m in _RETRYABLE_MESSAGES) + + +def _retry_with_backoff( + func: Callable[[], T], + *, + operation_name: str, + max_retries: int | None = None, + base_delay: float = 1.0, + max_delay: float = 60.0, + jitter: float = 0.5, +) -> T: + attempt = 0 + while True: + try: + return func() + except Exception as exc: + if not _is_retryable_error(exc): + raise + attempt += 1 + if max_retries is not None and attempt > max_retries: + logger.error("%s failed after %d retries: %s", operation_name, attempt - 1, exc) + raise + delay = min(base_delay * (2 ** (attempt - 1)), max_delay) + delay *= 1 + random.uniform(-jitter, jitter) + logger.warning("%s throttled (attempt %d), retrying in %.1fs: %s", operation_name, attempt, delay, exc) + time.sleep(delay) + + +# ── SSH helpers ────────────────────────────────────────────────────── + + +def _free_port() -> int: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def download_secret_to_file(secret_arn: str, region: str | None = None) -> str: + """Fetch a Secrets Manager secret → temp file (mode 0600).""" + key_material = download_secret_to_string(secret_arn, region=region) + fd, path = tempfile.mkstemp(prefix="ecs-ssh-", suffix=".key") + try: + os.write(fd, key_material.encode()) + finally: + os.close(fd) + os.chmod(path, 0o600) + return path + + +def download_secret_to_string(secret_arn: str, region: str | None = None) -> str: + boto3, *_ = _require_aws_sdks() + sm = boto3.client("secretsmanager", region_name=region) + return _retry_with_backoff( + lambda: sm.get_secret_value(SecretId=secret_arn)["SecretString"], + operation_name="secretsmanager.get_secret_value", + max_retries=5, + ) + + +# ── SSH tunnel ─────────────────────────────────────────────────────── + + +class SshTunnel: + """Manages an ``ssh -N`` subprocess with ``-L`` / ``-R`` tunnels.""" + + def __init__( + self, + *, + host: str, + port: int = 2222, + user: str = "root", + key_file: str, + forward_port: int | None = None, + forwards: list[str] | None = None, + reverses: list[str] | None = None, + local_port_override: int | None = None, + ) -> None: + self._host = host + self._port = port + self._user = user + self._key_file = key_file + self._simple_forward_port = forward_port + self._forwards = list(forwards or []) + self._reverses = list(reverses or []) + self._local_port: int | None = local_port_override + self._proc: subprocess.Popen[bytes] | None = None + + @property + def local_port(self) -> int: + if self._local_port is None: + raise RuntimeError("Tunnel not open yet — call open() first") + return self._local_port + + @property + def is_open(self) -> bool: + return self._proc is not None and self._proc.poll() is None + + def open(self, *, max_retries: int = 15, initial_backoff: float = 5.0) -> None: + if self.is_open: + return + use_simple = self._simple_forward_port is not None + last_err = "" + backoff = initial_backoff + for attempt in range(1, max_retries + 1): + if use_simple: + self._local_port = _free_port() + cmd = self._build_ssh_cmd() + logger.info("SSH tunnel attempt %d/%d: %s", attempt, max_retries, " ".join(cmd)) + self._proc = subprocess.Popen(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.PIPE) + time.sleep(3) + if self._proc.poll() is None: + if self._local_port: + try: + self._wait_for_local_port(self._local_port, timeout=15.0) + except Exception as port_exc: + logger.warning("SSH alive but forward port %d not open: %s", self._local_port, port_exc) + self._kill() + last_err = str(port_exc) + time.sleep(min(5.0, attempt * 1.5)) + continue + logger.info("SSH tunnel started (pid=%d, attempt %d/%d)", self._proc.pid, attempt, max_retries) + return + stderr = self._proc.stderr.read().decode(errors="replace") if self._proc.stderr else "" + last_err = stderr.strip() + self._proc = None + if not any( + m in last_err + for m in ( + "Connection refused", + "Connection timed out", + "No route to host", + "Connection reset", + ) + ): + raise RuntimeError(f"SSH tunnel exited immediately (attempt {attempt}): {last_err}") + logger.warning( + "SSH tunnel attempt %d/%d failed: %s — retrying in %.0fs", attempt, max_retries, last_err, backoff + ) + time.sleep(backoff) + backoff = min(30.0, backoff * 1.5) + raise RuntimeError(f"SSH tunnel failed after {max_retries} attempts: {last_err}") + + def close(self) -> None: + self._kill() + + def wait_ready(self, *, health_url: str | None = None, timeout: float = 300.0) -> None: + if health_url: + self._poll_health(health_url, timeout) + elif self._local_port: + self._wait_for_local_port(self._local_port, timeout) + + def check_health(self) -> bool: + return self.is_open + + def __enter__(self) -> SshTunnel: + self.open() + return self + + def __exit__(self, *exc: object) -> None: + self.close() + + def _build_ssh_cmd(self) -> list[str]: + cmd = [ + "ssh", + "-N", + "-o", + "StrictHostKeyChecking=no", + "-o", + "UserKnownHostsFile=/dev/null", + "-o", + "ServerAliveInterval=30", + "-o", + "ServerAliveCountMax=20", + "-o", + "ConnectTimeout=15", + "-o", + "ExitOnForwardFailure=yes", + "-o", + "LogLevel=ERROR", + "-i", + self._key_file, + "-p", + str(self._port), + ] + if self._simple_forward_port is not None: + cmd += ["-L", f"127.0.0.1:{self._local_port}:127.0.0.1:{self._simple_forward_port}"] + for spec in self._forwards: + cmd += ["-L", spec] + for spec in self._reverses: + cmd += ["-R", spec] + cmd.append(f"{self._user}@{self._host}") + return cmd + + def _kill(self) -> None: + if self._proc is None: + return + try: + self._proc.terminate() + try: + self._proc.wait(timeout=5) + except subprocess.TimeoutExpired: + self._proc.kill() + logger.info("SSH tunnel closed (pid=%d)", self._proc.pid) + except ProcessLookupError: + pass + finally: + self._proc = None + + def _wait_for_local_port(self, port: int, timeout: float = 30.0) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if self._proc and self._proc.poll() is not None: + raise RuntimeError("SSH tunnel process exited while waiting for port") + try: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.settimeout(1.0) + s.connect(("127.0.0.1", port)) + return + except OSError: + time.sleep(0.3) + raise TimeoutError(f"Local port 127.0.0.1:{port} not open after {timeout:.0f}s") + + def _poll_health(self, url: str, timeout: float) -> None: + import urllib.error + import urllib.request + + deadline = time.monotonic() + timeout + attempt = 0 + while time.monotonic() < deadline: + attempt += 1 + if not self.is_open: + raise RuntimeError("SSH tunnel died while waiting for health endpoint") + try: + with urllib.request.urlopen(urllib.request.Request(url, method="GET"), timeout=5) as resp: + if resp.status == 200: + logger.info("Health endpoint ready (attempt %d): %s", attempt, url) + return + except (urllib.error.URLError, OSError, TimeoutError): + pass + time.sleep(min(3.0, 1.0 + attempt * 0.5)) + raise TimeoutError(f"Health endpoint not reachable after {timeout:.0f}s: {url}") + + +# ── SSH sidecar container builder ──────────────────────────────────── + + +def build_ssh_sidecar_container( + sidecar_cfg: SshSidecarConfig, + *, + public_key_value: str, + max_lifetime_sec: int, + log_group: str | None = None, + log_region: str = "us-east-1", + log_stream_prefix: str = "ecs-sandbox", +) -> dict[str, Any]: + port = sidecar_cfg.sshd_port + image = sidecar_cfg.image or "alpine:latest" + sshd_cfg = ( + f"Port {port}\\nPermitRootLogin prohibit-password\\n" + "PasswordAuthentication no\\nAllowTcpForwarding yes\\n" + "PermitListen any\\nGatewayPorts clientspecified\\n" + "X11Forwarding no\\nPrintMotd no\\nLogLevel ERROR\\n" + "ClientAliveInterval 30\\nClientAliveCountMax 20\\n" + "TCPKeepAlive yes\\nUseDNS no\\nMaxSessions 50\\n" + ) + watchdog = "" + if max_lifetime_sec > 0: + watchdog = ( + f"( sleep {max_lifetime_sec}; " + f"echo 'sidecar watchdog: TTL ({max_lifetime_sec}s) reached'; " + "kill 1 2>/dev/null; sleep 3; kill -9 1 2>/dev/null ) & " + ) + sshd_cmd = ( + "set -e; apk add --no-cache openssh-server netcat-openbsd; " + "mkdir -p /root/.ssh; chmod 700 /root/.ssh; " + 'printf "%s\\n" "$SSH_PUBLIC_KEY" > /root/.ssh/authorized_keys; ' + "chmod 600 /root/.ssh/authorized_keys; ssh-keygen -A; " + f"printf '{sshd_cfg}' > /etc/ssh/sshd_config; " + f"{watchdog}exec /usr/sbin/sshd -D -e -p {port}" + ) + container: dict[str, Any] = { + "name": "ssh-tunnel", + "image": image, + "essential": True, + "entryPoint": ["sh", "-c"], + "command": [sshd_cmd], + "environment": [{"name": "SSH_PUBLIC_KEY", "value": public_key_value}], + "healthCheck": { + "command": ["CMD-SHELL", f"nc -z localhost {port} || exit 1"], + "interval": 5, + "timeout": 3, + "retries": 10, + "startPeriod": 30, + }, + } + if log_group: + container["logConfiguration"] = { + "logDriver": "awslogs", + "options": { + "awslogs-group": log_group, + "awslogs-region": log_region, + "awslogs-stream-prefix": f"{log_stream_prefix}-tunnel", + "awslogs-create-group": "true", + }, + } + return container + + +# ── Exec server — embedded script + HTTP client ───────────────────── + +EXEC_SERVER_SCRIPT = r'''#!/usr/bin/env python3 +"""Zero-dependency HTTP exec server for sandbox containers.""" +from __future__ import annotations +import base64, json, os, shutil, subprocess +from http.server import BaseHTTPRequestHandler, HTTPServer, ThreadingHTTPServer +from urllib.parse import parse_qs, urlparse +_BASH = shutil.which("bash") +_PORT = int(os.environ.get("TB_EXEC_PORT", "19542")) +_BIND = os.environ.get("TB_EXEC_BIND", "127.0.0.1") +class _H(BaseHTTPRequestHandler): + def log_message(self, fmt, *a): pass + def do_GET(self): + p = urlparse(self.path) + if p.path == "/health": self._ok({"ok": True}) + elif p.path == "/download": + qs = parse_qs(p.query) + paths = qs.get("path", []) + if not paths: self._err(400, "missing ?path=") + else: self._dl(paths[0]) + else: self._err(404, f"not found: {p.path}") + def do_POST(self): + p = urlparse(self.path) + body = self._body() + if p.path == "/exec": self._exec(body) + elif p.path == "/upload": self._up(body) + else: self._err(404, f"not found: {p.path}") + def _exec(self, b): + cmd = b.get("cmd") + if not cmd: self._err(400, "missing 'cmd'"); return + t = b.get("timeout", 300) + try: + cp = subprocess.run(cmd, shell=True, executable=_BASH, capture_output=True, timeout=t) + self._ok({"stdout": cp.stdout.decode("utf-8", errors="replace"), + "stderr": cp.stderr.decode("utf-8", errors="replace"), + "rc": cp.returncode}) + except subprocess.TimeoutExpired: + self._ok({"stdout":"","stderr":f"timed out after {t}s","rc":124}) + except Exception as e: + self._ok({"stdout":"","stderr":str(e),"rc":-1}) + def _up(self, b): + path, c = b.get("path"), b.get("content") + if not path or c is None: self._err(400, "missing path/content"); return + try: + data = base64.b64decode(c) + os.makedirs(os.path.dirname(path) or ".", exist_ok=True) + with open(path, "wb") as f: f.write(data) + m = b.get("mode") + if m: os.chmod(path, int(m, 8)) + self._ok({"ok": True}) + except Exception as e: self._err(500, str(e)) + def _dl(self, path): + if not os.path.isfile(path): self._err(404, f"not found: {path}"); return + try: + with open(path, "rb") as f: data = f.read() + self.send_response(200) + self.send_header("Content-Type", "application/octet-stream") + self.send_header("Content-Length", str(len(data))) + self.end_headers(); self.wfile.write(data) + except Exception as e: self._err(500, str(e)) + def _body(self): + n = int(self.headers.get("Content-Length", 0)) + if n == 0: return {} + try: return json.loads(self.rfile.read(n)) + except Exception: return {} + def _ok(self, obj): + p = json.dumps(obj).encode() + self.send_response(200) + self.send_header("Content-Type","application/json") + self.send_header("Content-Length",str(len(p))) + self.end_headers(); self.wfile.write(p) + def _err(self, code, msg): + p = json.dumps({"error": msg}).encode() + self.send_response(code) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(p))) + self.end_headers(); self.wfile.write(p) +if __name__ == "__main__": + s = ThreadingHTTPServer((_BIND, _PORT), _H) + print(f"exec_server on {_BIND}:{_PORT}", flush=True) + try: s.serve_forever() + except KeyboardInterrupt: pass + finally: s.server_close() +''' + +_EXEC_SERVER_B64 = base64.b64encode(EXEC_SERVER_SCRIPT.encode()).decode() + +_TRANSIENT_ERRORS = ( + ConnectionResetError, + ConnectionRefusedError, + ConnectionAbortedError, + BrokenPipeError, + TimeoutError, + OSError, +) + + +class ExecClient: + """Async HTTP client for the exec server (through the SSH tunnel). + + Methods are coroutines so each in-flight request occupies an event-loop + slot, not an executor thread — required to scale past asyncio's default + thread-pool cap when many long agent commands run concurrently (FEP-886). + """ + + def __init__(self, *, port: int, connect_timeout: float = 30.0) -> None: + self._base = f"http://127.0.0.1:{port}" + self._timeout = connect_timeout + # Lazy: aiohttp.ClientSession must be created inside a running event + # loop, but ExecClient is constructed from `_do_start` (which runs in + # asyncio.to_thread). Defer creation to the first request. + self._session: aiohttp.ClientSession | None = None + + async def _ensure_session(self) -> aiohttp.ClientSession: + if self._session is None or self._session.closed: + self._session = aiohttp.ClientSession() + return self._session + + async def close(self) -> None: + if self._session is not None and not self._session.closed: + await self._session.close() + self._session = None + + async def exec(self, cmd: str, *, timeout: int = 300) -> ExecResult: + resp = await self._post("/exec", {"cmd": cmd, "timeout": timeout}) + return ExecResult( + stdout=resp.get("stdout", ""), + stderr=resp.get("stderr", ""), + return_code=resp.get("rc", -1), + ) + + async def upload( + self, remote_path: str, data: bytes | Path, *, mode: str | None = None, max_retries: int = 3 + ) -> None: + if isinstance(data, Path): + data = data.read_bytes() + body: dict[str, Any] = {"path": remote_path, "content": base64.b64encode(data).decode()} + if mode is not None: + body["mode"] = mode + payload_mb = len(body["content"]) / (1024 * 1024) + upload_timeout = max(self._timeout, 60.0 + payload_mb * 2.0) + last_err: Exception | None = None + for attempt in range(1, max_retries + 1): + try: + resp = await self._post("/upload", body, timeout_override=upload_timeout) + if not resp.get("ok"): + raise RuntimeError(f"upload to {remote_path} failed: {resp}") + return + except (TimeoutError, OSError, RuntimeError) as exc: + last_err = exc + if attempt < max_retries: + logger.warning("upload %s attempt %d/%d: %s", remote_path, attempt, max_retries, exc) + await asyncio.sleep(2.0 * attempt) + raise RuntimeError(f"upload to {remote_path} failed after {max_retries} attempts: {last_err}") + + async def download(self, remote_path: str, *, max_retries: int = 3) -> bytes: + import urllib.parse + + url = f"{self._base}/download?path={urllib.parse.quote(remote_path)}" + return await self._request( + label=f"download {remote_path}", url=url, method="GET", timeout=self._timeout, max_retries=max_retries + ) + + async def health(self) -> bool: + try: + await self._request(label="health", url=f"{self._base}/health", method="GET", timeout=5, max_retries=1) + return True + except (ConnectionError, OSError, TimeoutError, RuntimeError): + return False + + async def _post( + self, path: str, body: dict[str, Any], *, timeout_override: float | None = None, max_retries: int = 4 + ) -> dict[str, Any]: + url = f"{self._base}{path}" + payload = json.dumps(body).encode() + if timeout_override is not None: + http_timeout = timeout_override + else: + cmd_timeout = body.get("timeout") + http_timeout = ( + max(self._timeout, cmd_timeout + 30) if isinstance(cmd_timeout, (int, float)) else self._timeout + ) + raw = await self._request( + label=f"POST {path}", + url=url, + method="POST", + data=payload, + headers={"Content-Type": "application/json"}, + timeout=http_timeout, + max_retries=max_retries, + ) + return json.loads(raw) + + async def _request( + self, + *, + label: str, + url: str, + method: str, + data: bytes | None = None, + headers: dict[str, str] | None = None, + timeout: float, + max_retries: int, + ) -> bytes: + session = await self._ensure_session() + client_timeout = aiohttp.ClientTimeout(total=timeout) + last_err: Exception | None = None + for attempt in range(1, max_retries + 1): + try: + async with session.request(method, url, data=data, headers=headers, timeout=client_timeout) as resp: + body = await resp.read() + if resp.status >= 400: + raise RuntimeError(f"{label} failed (HTTP {resp.status}): {body.decode(errors='replace')}") + return body + except RuntimeError: + raise + except (aiohttp.ClientError, TimeoutError, OSError) as exc: + last_err = exc + if attempt < max_retries: + wait = min(15.0, 2.0 ** (attempt - 1)) + logger.warning("%s attempt %d/%d: %s — retry in %.1fs", label, attempt, max_retries, exc, wait) + await asyncio.sleep(wait) + continue + raise ConnectionError(f"{label} failed after {max_retries} attempts: {last_err}") from last_err + raise ConnectionError(f"{label} unreachable") + + +# ── Image builder — AWS CodeBuild + ECR caching ───────────────────── + + +class ImageBuilder: + """Build Docker images via CodeBuild → ECR with content-hash caching.""" + + _lock = threading.Lock() + _inflight_builds: dict[str, threading.Event] = {} + _build_semaphore: threading.Semaphore | None = None + _build_semaphore_size: int = 0 + + @staticmethod + def get_ecr_image_tag(environment_dir: str | Path, environment_name: str) -> str: + h = hashlib.sha256() + root = Path(environment_dir) + for p in sorted(root.rglob("*")): + if p.is_file(): + h.update(str(p.relative_to(root)).encode()) + h.update(p.read_bytes()) + return f"{environment_name}__{h.hexdigest()[:8]}" + + @staticmethod + def image_exists_in_ecr(ecr_repository: str, tag: str, region: str | None = None) -> bool: + boto3, _, ClientError = _require_aws_sdks() + ecr_region = ImageBuilder._ecr_region(ecr_repository, fallback=region) + ecr = boto3.client("ecr", region_name=ecr_region) + repo_name = ecr_repository.split("/", 1)[1] if "/" in ecr_repository else ecr_repository + try: + _retry_with_backoff( + lambda: ecr.describe_images(repositoryName=repo_name, imageIds=[{"imageTag": tag}]), + operation_name="ecr.describe_images", + max_retries=5, + ) + return True + except ClientError as exc: + code = exc.response.get("Error", {}).get("Code", "") + if code in ("ImageNotFoundException", "RepositoryNotFoundException"): + return False + raise + + @staticmethod + def _ecr_region(ecr_repository: str, fallback: str | None = None) -> str | None: + """Extract the region from an ECR repo URL like '123.dkr.ecr.us-west-2.amazonaws.com/repo'.""" + parts = ecr_repository.split(".") + if len(parts) >= 4 and parts[1] == "dkr" and parts[2] == "ecr": + return parts[3] + return fallback + + @staticmethod + def list_ecr_tags(ecr_repository: str, region: str | None = None) -> set[str]: + """Return all image tags present in an ECR repository. + + Uses paginated ``list_images`` to fetch every tagged image in + a handful of API calls, rather than one ``describe_images`` + call per tag. + """ + boto3, _, ClientError = _require_aws_sdks() + ecr_region = ImageBuilder._ecr_region(ecr_repository, fallback=region) + ecr = boto3.client("ecr", region_name=ecr_region) + repo_name = ecr_repository.split("/", 1)[1] if "/" in ecr_repository else ecr_repository + + def _fetch_all_tags() -> set[str]: + tags: set[str] = set() + paginator = ecr.get_paginator("list_images") + for page in paginator.paginate( + repositoryName=repo_name, + filter={"tagStatus": "TAGGED"}, + ): + for img_id in page.get("imageIds", []): + if tag := img_id.get("imageTag"): + tags.add(tag) + return tags + + try: + return _retry_with_backoff(_fetch_all_tags, operation_name="ecr.list_images", max_retries=5) + except ClientError as exc: + if exc.response.get("Error", {}).get("Code") == "RepositoryNotFoundException": + return set() + raise + + @staticmethod + def ecr_docker_login(ecr_repository: str, region: str | None = None) -> None: + """Authenticate the local Docker daemon against an ECR registry.""" + ecr_region = ImageBuilder._ecr_region(ecr_repository, fallback=region) + registry = ecr_repository.split("/")[0] + region_flag = f" --region {ecr_region}" if ecr_region else "" + cmd = f"aws ecr get-login-password{region_flag} | docker login --username AWS --password-stdin {registry}" + result = subprocess.run(cmd, shell=True, capture_output=True, text=True) + if result.returncode != 0: + raise RuntimeError(f"ECR docker login failed: {result.stderr.strip()}") + logger.info("ECR docker login succeeded for %s", registry) + + @staticmethod + def docker_push_to_ecr(local_image: str, ecr_repository: str, tag: str) -> str: + """Tag a local Docker image and push it to ECR. Returns the ECR URL.""" + ecr_url = f"{ecr_repository}:{tag}" + subprocess.run(["docker", "tag", local_image, ecr_url], check=True, capture_output=True) + result = subprocess.run(["docker", "push", ecr_url], capture_output=True, text=True) + if result.returncode != 0: + raise RuntimeError(f"docker push {ecr_url} failed: {result.stderr.strip()}") + logger.info("Pushed %s -> %s", local_image, ecr_url) + return ecr_url + + @classmethod + def ensure_image_built(cls, *, cfg: EcsFargateConfig, environment_name: str, force_build: bool = False) -> str: + ecr_repo = cfg.ecr_repository + env_dir = cfg.environment_dir + if not ecr_repo or not env_dir: + raise ValueError("ecr_repository and environment_dir are required for image building") + tag = cls.get_ecr_image_tag(env_dir, environment_name) + image_url = f"{ecr_repo}:{tag}" + + # Dedup: check if another thread is already building this tag + with cls._lock: + if tag in cls._inflight_builds: + event = cls._inflight_builds[tag] + event.wait() + return image_url + if not force_build and cls.image_exists_in_ecr(ecr_repo, tag, cfg.region): + logger.info("ECR cache hit — skipping build: %s", image_url) + return image_url + + # Register as builder + event = threading.Event() + with cls._lock: + if tag in cls._inflight_builds: + cls._inflight_builds[tag].wait() + return image_url + cls._inflight_builds[tag] = event + with cls._lock: + if cls._build_semaphore is None or cls._build_semaphore_size != cfg.build_parallelism: + cls._build_semaphore = threading.Semaphore(cfg.build_parallelism) + cls._build_semaphore_size = cfg.build_parallelism + try: + cls._build_semaphore.acquire() # type: ignore[union-attr] + try: + if not force_build and cls.image_exists_in_ecr(ecr_repo, tag, cfg.region): + return image_url + cls._build_and_push(cfg=cfg, environment_name=environment_name, tag=tag, image_url=image_url) + finally: + cls._build_semaphore.release() # type: ignore[union-attr] + finally: + event.set() + with cls._lock: + cls._inflight_builds.pop(tag, None) + return image_url + + @classmethod + def ensure_mirrored(cls, *, cfg: EcsFargateConfig, src_image: str, force: bool = False) -> str: + """Ensure ``src_image`` is present in the ECR mirror, pulling it if not. + + Public/bare image references are served from the ECR mirror tag + (``{ecr_repository}:{sanitize(src_image)}``) rather than pulled from + their origin registry at task-launch time. This mirrors the image into + ECR on demand via a self-contained CodeBuild job (privileged DinD that + logs into Docker Hub + ECR, pulls, retags, and pushes), so callers never + have to pre-stage images. Concurrent callers for the same tag dedup on a + shared in-flight event, exactly like :meth:`ensure_image_built`. + """ + ecr_repo = cfg.ecr_repository + if not ecr_repo: + raise ValueError("ecr_repository is required to mirror images") + tag = _sanitize_id(src_image) + image_url = f"{ecr_repo}:{tag}" + + with cls._lock: + if tag in cls._inflight_builds: + cls._inflight_builds[tag].wait() + return image_url + if not force and cls.image_exists_in_ecr(ecr_repo, tag, cfg.region): + logger.info("ECR mirror hit — skipping pull: %s", image_url) + return image_url + + event = threading.Event() + with cls._lock: + if tag in cls._inflight_builds: + cls._inflight_builds[tag].wait() + return image_url + cls._inflight_builds[tag] = event + with cls._lock: + if cls._build_semaphore is None or cls._build_semaphore_size != cfg.build_parallelism: + cls._build_semaphore = threading.Semaphore(cfg.build_parallelism) + cls._build_semaphore_size = cfg.build_parallelism + try: + cls._build_semaphore.acquire() # type: ignore[union-attr] + try: + if not force and cls.image_exists_in_ecr(ecr_repo, tag, cfg.region): + return image_url + logger.info("Mirroring %s -> %s via CodeBuild ...", src_image, image_url) + buildspec = cls._generate_mirror_buildspec(cfg, src_image, image_url) + cls.run_buildspec_via_codebuild( + cfg=cfg, + buildspec=buildspec, + job_label=f"mirror::{tag}", + timeout_minutes=cfg.codebuild_build_timeout, + ) + logger.info("Mirrored OK: %s -> %s", src_image, image_url) + finally: + cls._build_semaphore.release() # type: ignore[union-attr] + finally: + event.set() + with cls._lock: + cls._inflight_builds.pop(tag, None) + return image_url + + @staticmethod + def _generate_mirror_buildspec(cfg: EcsFargateConfig, src_image: str, ecr_url: str) -> str: + ecr_registry = (cfg.ecr_repository or "").split("/")[0] + ecr_region = ImageBuilder._ecr_region(cfg.ecr_repository or "", fallback="$AWS_DEFAULT_REGION") + pre_build_cmds = [ + f"aws ecr get-login-password --region {ecr_region}" + f" | docker login --username AWS --password-stdin {ecr_registry}", + ] + if cfg.dockerhub_secret_arn: + pre_build_cmds.append( + f"DOCKERHUB_CREDS=$(aws secretsmanager get-secret-value" + f" --secret-id {cfg.dockerhub_secret_arn}" + f" --query SecretString --output text --region $AWS_DEFAULT_REGION)" + f' && DH_USER=$(echo "$DOCKERHUB_CREDS" | python3 -c' + """ "import sys,json;print(json.load(sys.stdin)['username'])")""" + f' && if [ -n "$DH_USER" ]; then echo "$DOCKERHUB_CREDS" | python3 -c' + """ "import sys,json;print(json.load(sys.stdin)['password'])" """ + f'| docker login -u "$DH_USER" --password-stdin; fi' + f' || echo "Docker Hub login failed — continuing without auth"' + ) + pull_cmd = ( + f"for i in 1 2 3; do docker pull --platform linux/amd64 {src_image} && break; " + f'echo "pull failed ($i/3), retry in 30s"; sleep 30; done' + ) + pre_yaml = "\n".join(f" - {c}" for c in pre_build_cmds) + return ( + "version: 0.2\nphases:\n pre_build:\n commands:\n" + f"{pre_yaml}\n build:\n commands:\n" + f" - {pull_cmd}\n" + f" - docker tag {src_image} {ecr_url}\n" + f" post_build:\n commands:\n - docker push {ecr_url}\n" + ) + + @staticmethod + def _upload_build_context(cfg: EcsFargateConfig, environment_name: str, nonce: str) -> str: + boto3, *_ = _require_aws_sdks() + env_dir = Path(cfg.environment_dir or ".") + buf = io.BytesIO() + with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as zf: + for item in env_dir.rglob("*"): + if item.is_file(): + zf.write(item, arcname=str(item.relative_to(env_dir))) + buf.seek(0) + s3 = boto3.client("s3", region_name=cfg.region) + s3_prefix = cfg.s3_prefix or "ecs-sandbox" + s3_key = f"{s3_prefix}/codebuild/{environment_name}-{nonce}.zip" + body = buf.read() + _retry_with_backoff( + lambda: s3.put_object(Bucket=cfg.s3_bucket, Key=s3_key, Body=body), + operation_name="s3.put_object(build_context)", + max_retries=5, + ) + return s3_key + + @staticmethod + def _resolve_codebuild_project(cfg: EcsFargateConfig, cb: Any, nonce: str) -> str: + _, _, ClientError = _require_aws_sdks() + if cfg.codebuild_project: + return cfg.codebuild_project + if not cfg.codebuild_service_role: + raise RuntimeError("codebuild_project or codebuild_service_role is required") + project_name = f"ecs-sandbox-build-{nonce}" + try: + _retry_with_backoff( + lambda: cb.create_project( + name=project_name, + source={"type": "NO_SOURCE", "buildspec": "version: 0.2"}, + artifacts={"type": "NO_ARTIFACTS"}, + environment={ + "type": "LINUX_CONTAINER", + "image": "aws/codebuild/amazonlinux-x86_64-standard:5.0", + "computeType": cfg.codebuild_compute_type, + "privilegedMode": True, + }, + serviceRole=cfg.codebuild_service_role, + timeoutInMinutes=cfg.codebuild_build_timeout, + ), + operation_name="codebuild.create_project", + max_retries=5, + ) + except ClientError as e: + if "already exists" not in str(e).lower(): + raise + return project_name + + @staticmethod + def _generate_buildspec(cfg: EcsFargateConfig, repo_name: str, tag: str, image_url: str) -> str: + ecr_registry = (cfg.ecr_repository or "").split("/")[0] + ecr_region = ImageBuilder._ecr_region(cfg.ecr_repository or "", fallback="$AWS_DEFAULT_REGION") + pre_build_cmds = [ + f"aws ecr get-login-password --region {ecr_region}" + f" | docker login --username AWS --password-stdin {ecr_registry}", + ] + if cfg.dockerhub_secret_arn: + pre_build_cmds.append( + f"DOCKERHUB_CREDS=$(aws secretsmanager get-secret-value" + f" --secret-id {cfg.dockerhub_secret_arn}" + f" --query SecretString --output text --region $AWS_DEFAULT_REGION)" + f' && DH_USER=$(echo "$DOCKERHUB_CREDS" | python3 -c' + """ "import sys,json;print(json.load(sys.stdin)['username'])")""" + f' && if [ -n "$DH_USER" ]; then echo "$DOCKERHUB_CREDS" | python3 -c' + """ "import sys,json;print(json.load(sys.stdin)['password'])" """ + f'| docker login -u "$DH_USER" --password-stdin; fi' + f' || echo "Docker Hub login failed — continuing without auth"' + ) + pre_yaml = "\n".join(f" - {c}" for c in pre_build_cmds) + build_cmd = ( + f"for i in 1 2 3; do docker build -t {repo_name}:{tag} . && break; " + f'echo "build failed ($i/3), retry in 30s"; sleep 30; done' + ) + return ( + "version: 0.2\nphases:\n pre_build:\n commands:\n" + f"{pre_yaml}\n build:\n commands:\n" + f" - {build_cmd}\n - docker tag {repo_name}:{tag} {image_url}\n" + f" post_build:\n commands:\n - docker push {image_url}\n" + ) + + @staticmethod + def _poll_codebuild(cb: Any, build_id: str, image_url: str) -> None: + while True: + time.sleep(10 + random.uniform(0, 5)) + build = _retry_with_backoff( + lambda: cb.batch_get_builds(ids=[build_id])["builds"][0], + operation_name=f"BatchGetBuilds({build_id})", + max_retries=8, + base_delay=2.0, + max_delay=120.0, + ) + status = build["buildStatus"] + if status == "SUCCEEDED": + logger.info("CodeBuild succeeded: %s", build_id) + return + if status in ("FAILED", "FAULT", "STOPPED", "TIMED_OUT"): + phases = build.get("phases", []) + failed = [p for p in phases if p.get("phaseStatus") not in (None, "SUCCEEDED")] + ctx = "; ".join(f"{p['phaseType']}: {p.get('phaseStatus')}" for p in failed) or status + raise RuntimeError(f"CodeBuild failed for {image_url}: {ctx} (build: {build_id})") + logger.debug("CodeBuild %s — phase=%s status=%s", build_id, build.get("currentPhase"), status) + + @classmethod + def _build_and_push(cls, *, cfg: EcsFargateConfig, environment_name: str, tag: str, image_url: str) -> None: + boto3, *_ = _require_aws_sdks() + ecr_repo = cfg.ecr_repository or "" + repo_name = ecr_repo.split("/", 1)[1] if "/" in ecr_repo else ecr_repo + nonce = uuid.uuid4().hex[:8] + logger.info("Building image via CodeBuild: %s", image_url) + s3_key = cls._upload_build_context(cfg, environment_name, nonce) + cb = boto3.client("codebuild", region_name=cfg.region) + project_name = cls._resolve_codebuild_project(cfg, cb, nonce) + buildspec = cls._generate_buildspec(cfg, repo_name, tag, image_url) + resp = _retry_with_backoff( + lambda: cb.start_build( + projectName=project_name, + sourceTypeOverride="S3", + sourceLocationOverride=f"{cfg.s3_bucket}/{s3_key}", + buildspecOverride=buildspec, + timeoutInMinutesOverride=cfg.codebuild_build_timeout, + privilegedModeOverride=True, + environmentTypeOverride="LINUX_CONTAINER", + imageOverride="aws/codebuild/amazonlinux-x86_64-standard:5.0", + computeTypeOverride=cfg.codebuild_compute_type, + ), + operation_name="codebuild.start_build", + max_retries=5, + ) + build_id = resp["build"]["id"] + logger.info("CodeBuild started: %s", build_id) + cls._poll_codebuild(cb, build_id, image_url) + + @classmethod + def run_buildspec_via_codebuild( + cls, + *, + cfg: EcsFargateConfig, + buildspec: str, + job_label: str = "harness-build", + timeout_minutes: int | None = None, + ) -> None: + """Run an arbitrary buildspec via CodeBuild (privileged mode for DinD). + + Unlike :meth:`_build_and_push` which uploads a Dockerfile context to S3, + this method uses ``NO_SOURCE`` — the buildspec is fully self-contained + (e.g. it installs packages and runs a harness that builds Docker images + internally). + """ + boto3, *_ = _require_aws_sdks() + nonce = uuid.uuid4().hex[:8] + cb = boto3.client("codebuild", region_name=cfg.region) + project_name = cls._resolve_codebuild_project(cfg, cb, nonce) + timeout = timeout_minutes or cfg.codebuild_build_timeout + + logger.info("Starting CodeBuild harness build: %s (timeout=%dm)", job_label, timeout) + resp = _retry_with_backoff( + lambda: cb.start_build( + projectName=project_name, + sourceTypeOverride="NO_SOURCE", + buildspecOverride=buildspec, + timeoutInMinutesOverride=timeout, + privilegedModeOverride=True, + environmentTypeOverride="LINUX_CONTAINER", + imageOverride="aws/codebuild/amazonlinux-x86_64-standard:5.0", + computeTypeOverride=cfg.codebuild_compute_type, + ), + operation_name="codebuild.start_build(harness)", + max_retries=5, + ) + build_id = resp["build"]["id"] + logger.info("CodeBuild harness build started: %s (build=%s)", job_label, build_id) + cls._poll_codebuild(cb, build_id, job_label) + + +# ── Core sandbox ───────────────────────────────────────────────────── + +_active_sandboxes: dict[int, Any] = {} +_cleanup_lock = threading.RLock() +_atexit_registered = False +_PROCESS_NONCE = f"{int(time.time())}-{uuid.uuid4().hex[:8]}" +_exec_server_url_cache: dict[str, str] = {} + +_task_def_cache: dict[str, str] = {} +_task_def_cache_lock = threading.Lock() +_task_def_inflight: dict[str, threading.Event] = {} + +# Env vars whose values vary per sandbox invocation (workspace session id, etc). +# Routed via RunTask containerOverrides so the underlying task definition stays +# content-stable and the FEP-866 hash cache hits across invocations of the same +# task. Per-invocation env keys discovered dynamically from OutsideEndpoint +# routing (e.g. MODEL_BASE_URL with session-scoped URLs) are merged in at call +# time; this set holds the keys that are NOT visible to OutsideEndpoint routing. +_PER_INVOCATION_ENV_KEYS: frozenset[str] = frozenset({"_NEL_EFS_SESSION"}) + + +def _compute_task_def_hash(payload: dict[str, Any]) -> str: + # Strip logConfiguration from every container definition before hashing. + # Log config (group, stream-prefix, region) is a visibility annotation — it has + # no effect on what the sandbox does. Two task defs that differ only in + # log_stream_prefix or log_group are functionally identical and should share + # a cache entry so cross-run SSM cache hits work across differently-named runs. + def _strip_log_cfg(containers: list) -> list: + return [{k: v for k, v in c.items() if k != "logConfiguration"} for c in containers] + + canonical = { + k: (_strip_log_cfg(v) if k == "containerDefinitions" else v) for k, v in payload.items() if k != "family" + } + blob = json.dumps(canonical, sort_keys=True, default=str).encode() + return hashlib.sha256(blob).hexdigest()[:24] + + +def _emergency_cleanup() -> None: + with _cleanup_lock: + for sb in list(_active_sandboxes.values()): + try: + sb._sync_stop() + except Exception: + logger.debug("Emergency cleanup failed for sandbox %s", id(sb), exc_info=True) + + +class EcsFargateSandbox: + """ECS Fargate sandbox — async :class:`Sandbox` protocol.""" + + def __init__(self, spec: SandboxSpec, *, ecs_config: EcsFargateConfig) -> None: + self._spec = spec + self._cfg = ecs_config + self._task_arn: str | None = None + self._task_def_arn: str | None = None + self._task_ip: str | None = None + self._ssh_key_file: str | None = None + self._ssh_tunnel: SshTunnel | None = None + self._exec_client: ExecClient | None = None + self._started = False + self._stopped = False + self._ecs: Any = None + self._ec2: Any = None + self._ssm: Any = None + self._runtime_container_env: dict[str, str] = {} + self._ssh_tunnel_port: int | None = None + self._agent_forward_port: int | None = None + self._outside_endpoints: list[OutsideEndpoint] = [] + self._outside_endpoint_routing = _OutsideEndpointRouting.empty() + self._run_id = uuid.uuid4().hex[:12] + + # ── Protocol properties ────────────────────────────────────────── + + @property + def spec(self) -> SandboxSpec: + return self._spec + + def resolved_endpoint_url(self, env_var: str) -> str | None: + return self._outside_endpoint_routing.resolved_endpoint_url(env_var) + + @property + def is_running(self) -> bool: + return self._started and not self._stopped + + @property + def container_ip(self) -> str | None: + return self._task_ip + + # ── Extra properties ───────────────────────────────────────────── + + @property + def task_arn(self) -> str | None: + return self._task_arn + + @property + def local_port(self) -> int | None: + if self._ssh_tunnel: + try: + return self._ssh_tunnel.local_port + except RuntimeError: + pass + return None + + @property + def ssh_tunnel(self) -> SshTunnel | None: + return self._ssh_tunnel + + @property + def exec_client(self) -> ExecClient | None: + return self._exec_client + + @property + def model_tunnel_port(self) -> int | None: + return self._ssh_tunnel_port + + # ── Protocol async lifecycle ───────────────────────────────────── + + async def start(self, *, outside_endpoints: list[OutsideEndpoint] | None = None) -> None: + if self._started: + return + self._outside_endpoints = outside_endpoints or [] + self._outside_endpoint_routing = _OutsideEndpointRouting.empty(self._outside_endpoints) + sidecar = self._cfg.ssh_sidecar + if sidecar and sidecar.exec_server_port is None: + _OutsideEndpointRouting.for_agent_server(self._outside_endpoints) + try: + await asyncio.to_thread(self._do_start) + self._started = True + except Exception: + if self._exec_client is not None: + await self._exec_client.close() + await asyncio.to_thread(self._cleanup) + raise + + async def stop(self) -> None: + if self._stopped: + return + self._stopped = True + if self._exec_client is not None: + await self._exec_client.close() + await asyncio.to_thread(self._cleanup) + self._unregister_from_cleanup() + + async def exec( + self, + command: str, + timeout_sec: float = 180, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + user: str | int | None = None, + ) -> ExecResult: + self._require_exec_client() + shell_cmd = command + if env: + exports = " ".join(f"{k}={v}" for k, v in env.items()) + shell_cmd = f"export {exports} && {shell_cmd}" + if cwd: + shell_cmd = f"cd {cwd} && {shell_cmd}" + if user is not None: + if isinstance(user, int): + shell_cmd = f'su -s /bin/bash "$(getent passwd {user} | cut -d: -f1)" -c {shlex.quote(shell_cmd)}' + else: + shell_cmd = f"su -s /bin/bash {shlex.quote(str(user))} -c {shlex.quote(shell_cmd)}" + try: + return await self._exec_client.exec(shell_cmd, timeout=int(timeout_sec)) # type: ignore[union-attr] + except ConnectionError: + if self._ssh_tunnel and not self._ssh_tunnel.is_open: + logger.warning("SSH tunnel dead — attempting reconnect before re-raising") + try: + await self.reconnect_tunnel() + sidecar = self._cfg.ssh_sidecar + if sidecar and sidecar.exec_server_port is not None: + health_url = f"http://127.0.0.1:{self._ssh_tunnel.local_port}/health" # type: ignore[union-attr] + self._ssh_tunnel.wait_ready(health_url=health_url, timeout=60.0) # type: ignore[union-attr] + old_client = self._exec_client + self._exec_client = ExecClient(port=self._ssh_tunnel.local_port) # type: ignore[union-attr] + if old_client is not None: + await old_client.close() + return await self._exec_client.exec(shell_cmd, timeout=int(timeout_sec)) + except Exception as reconnect_err: + logger.warning("Tunnel reconnect failed: %s", reconnect_err) + raise + + async def upload(self, local_path: Path, remote_path: str) -> None: + self._require_exec_client() + local = Path(local_path) + if local.is_dir(): + for child in local.rglob("*"): + if child.is_file(): + await self.upload(child, f"{remote_path}/{child.relative_to(local)}") + return + if local.stat().st_size > 512 * 1024 and self._cfg.s3_bucket: + await self._upload_via_s3([local], os.path.dirname(remote_path) or "/tmp") + else: + await self._exec_client.upload(remote_path, local) # type: ignore[union-attr] + + async def download(self, remote_path: str, local_path: Path) -> None: + self._require_exec_client() + data = await self._exec_client.download(remote_path) # type: ignore[union-attr] + dest = Path(local_path) + dest.parent.mkdir(parents=True, exist_ok=True) + dest.write_bytes(data) + + def resolve_outside_endpoint(self, url: str) -> str: + return self._outside_endpoint_routing.resolve_url(url) + + async def __aenter__(self) -> Self: + await self.start() + return self + + async def __aexit__(self, *exc: object) -> None: + await self.stop() + + # ── Extra public methods ───────────────────────────────────────── + + async def reconnect_tunnel(self) -> None: + if self._stopped or not self._started: + raise RuntimeError("Cannot reconnect tunnel on a stopped/unstarted sandbox") + sidecar = self._cfg.ssh_sidecar + if sidecar is None: + return + if self._ssh_tunnel: + self._ssh_tunnel.close() + self._ssh_tunnel = None + await asyncio.to_thread(self._open_tunnel, sidecar) + + # ── Sync start (runs via asyncio.to_thread) ────────────────────── + + def _do_start(self) -> None: + cfg = self._cfg + sidecar = cfg.ssh_sidecar + if sidecar is None: + raise ValueError("ssh_sidecar must be configured") + self._init_aws_clients() + + built_image: str | None = None + env_dir = cfg.environment_dir or self._spec.environment_dir + if cfg.ecr_repository and env_dir: + per_task_cfg = _dc_replace(cfg, environment_dir=env_dir) + built_image = ImageBuilder.ensure_image_built( + cfg=per_task_cfg, environment_name=_sanitize_id(self._spec.image or "sandbox") + ) + elif ( + cfg.auto_mirror + and cfg.ecr_repository + and self._spec.image + and not cfg.image_template + and not _is_ecr_image_ref(self._spec.image) + ): + # The bare/public image is served from the ECR mirror tag; pull it + # into ECR on demand so a missing mirror entry self-heals instead of + # failing the task's image pull. + ImageBuilder.ensure_mirrored(cfg=cfg, src_image=self._spec.image) + image = self._resolve_image(built_image) + + if not sidecar.private_key_secret_arn or not sidecar.public_key_secret_arn: + raise ValueError("ssh_sidecar private_key_secret_arn and public_key_secret_arn are required") + self._ssh_key_file = download_secret_to_file(sidecar.private_key_secret_arn, cfg.region) + ssh_public_key_value = download_secret_to_string(sidecar.public_key_secret_arn, cfg.region) + + has_exec_server = sidecar.exec_server_port is not None + if not has_exec_server: + self._outside_endpoint_routing = _OutsideEndpointRouting.for_agent_server(self._outside_endpoints) + self._ssh_tunnel_port = self._outside_endpoint_routing.agent_tunnel_port + else: + self._outside_endpoint_routing = _OutsideEndpointRouting.for_exec_server(self._outside_endpoints, sidecar) + + command = self._build_container_command(sidecar) + env = self._build_env_vars() + stable_env, self._runtime_container_env = self._split_env(env) + log_region = cfg.region or os.environ.get("AWS_DEFAULT_REGION", "us-east-1") + sidecar_def = build_ssh_sidecar_container( + sidecar, + public_key_value=ssh_public_key_value, + max_lifetime_sec=cfg.max_task_lifetime_sec, + log_group=cfg.log_group, + log_region=log_region, + log_stream_prefix=cfg.log_stream_prefix or "ecs-sandbox", + ) + self._task_def_arn = self._register_task_definition( + image=image, command=command, env=stable_env, sidecar_def=sidecar_def + ) + self._task_arn = self._run_task(self._task_def_arn) + self._register_for_cleanup() + self._wait_for_running() + self._task_ip = self._get_task_public_ip() + self._wait_for_ssh_ready(self._task_ip, sidecar.sshd_port, sidecar.ssh_ready_timeout_sec) + self._open_tunnel(sidecar) + + if has_exec_server: + health_url = f"http://127.0.0.1:{self._ssh_tunnel.local_port}/health" # type: ignore[union-attr] + self._ssh_tunnel.wait_ready(health_url=health_url, timeout=sidecar.ssh_ready_timeout_sec) # type: ignore[union-attr] + self._exec_client = ExecClient(port=self._ssh_tunnel.local_port) # type: ignore[union-attr] + + # ── Internal helpers ───────────────────────────────────────────── + + def _init_aws_clients(self) -> None: + boto3, Config, _ = _require_aws_sdks() + boto_cfg = Config(connect_timeout=30, read_timeout=60, retries={"max_attempts": 8, "mode": "adaptive"}) + self._ecs = boto3.client("ecs", region_name=self._cfg.region, config=boto_cfg) + self._ec2 = boto3.client("ec2", region_name=self._cfg.region, config=boto_cfg) + self._ssm = boto3.client("ssm", region_name=self._cfg.region, config=boto_cfg) + + def _resolve_image(self, built_image: str | None = None) -> str: + if built_image: + return built_image + cfg = self._cfg + if cfg.image_template: + sanitized = _sanitize_id(self._spec.image or "sandbox") + fmt_keys = { + "task_id": self._spec.image or "", + "task_id_sanitized": sanitized, + **(self._spec.env or {}), + } + try: + return cfg.image_template.format_map(fmt_keys) + except KeyError as exc: + raise ValueError( + f"ecs.image_template placeholder {exc} not found in " + f"available keys: {sorted(fmt_keys)}. " + f"Hint: use sandbox.image_template (resolved via seed " + f"metadata) instead of ecs.image_template for task-specific " + f"placeholders like {{task_id}}." + ) from exc + if self._spec.image: + # A reference that already points at an ECR registry (e.g. the + # configured mirror) is used verbatim — re-prefixing it under + # ecr_repository and re-sanitizing would corrupt the tag and make + # an existing image unresolvable. + if _is_ecr_image_ref(self._spec.image): + return self._spec.image + # Bare / public names are routed to the ECR mirror tag rather than + # pulled directly from their origin registry (avoids Docker Hub + # rate limits when many tasks start concurrently). + if cfg.ecr_repository: + return f"{cfg.ecr_repository}:{_sanitize_id(self._spec.image)}" + return self._spec.image + if not cfg.task_definition: + raise ValueError( + "No image available: set image on SandboxSpec, image_template, " + "ecr_repository + environment_dir, or task_definition" + ) + return "" + + def _upload_exec_server(self) -> str: + cfg = self._cfg + if not cfg.s3_bucket: + raise ValueError("s3_bucket is required for exec server upload") + cache_key = f"{cfg.s3_bucket}/{self._run_id}" + if cache_key in _exec_server_url_cache: + return _exec_server_url_cache[cache_key] + boto3, *_ = _require_aws_sdks() + s3 = boto3.client("s3", region_name=cfg.region) + prefix = cfg.s3_prefix or "ecs-sandbox" + key = f"{prefix}/{self._run_id}-{_PROCESS_NONCE}/_exec_server/exec_server.py" + _retry_with_backoff( + lambda: s3.put_object(Bucket=cfg.s3_bucket, Key=key, Body=EXEC_SERVER_SCRIPT.encode()), + operation_name="s3.put_object(exec_server)", + max_retries=5, + ) + url = s3.generate_presigned_url("get_object", Params={"Bucket": cfg.s3_bucket, "Key": key}, ExpiresIn=21600) + _exec_server_url_cache[cache_key] = url + logger.info("Uploaded exec server → s3://%s/%s", cfg.s3_bucket, key) + return url + + def _build_container_command(self, sidecar: SshSidecarConfig) -> list[str] | None: + if sidecar.exec_server_port is None: + return None + exec_port = sidecar.exec_server_port or 19542 + hostname = re.sub(r"[^A-Za-z0-9._-]", "-", self._spec.image or "sandbox")[:63] + setup = ( + f"hostname {shlex.quote(hostname)} 2>/dev/null || true; " + f"echo '{_EXEC_SERVER_B64}' | base64 -d > /tmp/_exec_server.py; " + "if ! command -v python3 >/dev/null 2>&1; then " + " if command -v apt-get >/dev/null 2>&1; then " + " apt-get update -qq && apt-get install -y -qq --no-install-recommends python3 bash; " + " elif command -v apk >/dev/null 2>&1; then " + " apk add --no-cache python3 bash; " + " elif command -v yum >/dev/null 2>&1; then " + " yum install -y python3 bash; " + " elif command -v dnf >/dev/null 2>&1; then " + " dnf install -y python3 bash; " + " fi; " + "fi; " + "if ! command -v python3 >/dev/null 2>&1; then " + " echo 'FATAL: exec_server bootstrap failed — python3 not available' >&2; " + " exit 1; " + "fi; " + f"TB_EXEC_PORT={exec_port} TB_EXEC_BIND=127.0.0.1 " + "exec python3 /tmp/_exec_server.py" + ) + return ["sh", "-lc", setup] + + def _build_env_vars(self) -> dict[str, str]: + env: dict[str, str] = dict(self._spec.env) + cfg = self._cfg + if cfg.extra_env: + for k, v in cfg.extra_env.items(): + env[k] = self._render_env_value(v) + env.update(self._outside_endpoint_routing.env_overrides()) + return env + + def _split_env(self, env: dict[str, str]) -> tuple[dict[str, str], dict[str, str]]: + runtime_keys = _PER_INVOCATION_ENV_KEYS | set(self._outside_endpoint_routing.env_overrides().keys()) + stable = {k: v for k, v in env.items() if k not in runtime_keys} + runtime = {k: v for k, v in env.items() if k in runtime_keys} + return stable, runtime + + def _render_env_value(self, value: str) -> str: + if self._ssh_tunnel_port is not None: + value = value.replace("{ssh_tunnel_port}", str(self._ssh_tunnel_port)) + if self._task_ip: + value = value.replace("{task_ip}", self._task_ip) + value = value.replace("{image}", self._spec.image or "") + return value + + # ── Task definition registration ───────────────────────────────── + + def _register_task_definition( + self, *, image: str, command: list[str] | None, env: dict[str, str], sidecar_def: dict[str, Any] + ) -> str: + cfg = self._cfg + log_region = cfg.region or os.environ.get("AWS_DEFAULT_REGION", "us-east-1") + log_cfg: dict[str, Any] | None = None + if cfg.log_group: + log_cfg = { + "logDriver": "awslogs", + "options": { + "awslogs-group": cfg.log_group, + "awslogs-region": log_region, + "awslogs-stream-prefix": cfg.log_stream_prefix or "ecs-sandbox", + "awslogs-create-group": "true", + }, + } + + _, _, ClientError = _require_aws_sdks() + base: dict[str, Any] | None = None + if cfg.task_definition: + try: + base = _retry_with_backoff( + lambda: self._ecs.describe_task_definition(taskDefinition=cfg.task_definition)["taskDefinition"], + operation_name="ecs.describe_task_definition", + max_retries=5, + ) + except ClientError as exc: + if exc.response.get("Error", {}).get("Code") == "ClientException": + logger.warning("Base task definition %s not found, registering from scratch", cfg.task_definition) + else: + raise + + if base is not None: + return self._register_from_base( + base=base, image=image, command=command, env=env, sidecar_def=sidecar_def, log_cfg=log_cfg + ) + return self._register_from_scratch( + image=image, command=command, env=env, sidecar_def=sidecar_def, log_cfg=log_cfg + ) + + def _register_from_base( + self, + *, + base: dict, + image: str, + command: list[str] | None, + env: dict[str, str], + sidecar_def: dict, + log_cfg: dict | None, + ) -> str: + cfg = self._cfg + containers = list(base.get("containerDefinitions") or []) + target = next((cd for cd in containers if cd.get("name") == cfg.container_name), None) + if target is None: + raise RuntimeError( + f"Base task-def has no container '{cfg.container_name}'. " + f"Available: {[c.get('name') for c in containers]}" + ) + if image: + target["image"] = image + if command is not None: + target["command"] = command + target.pop("entryPoint", None) + if env: + existing = {e["name"]: e["value"] for e in target.get("environment", [])} + existing.update(env) + target["environment"] = [{"name": k, "value": v} for k, v in sorted(existing.items())] + if log_cfg: + target["logConfiguration"] = log_cfg + target["dependsOn"] = [{"containerName": "ssh-tunnel", "condition": "HEALTHY"}] + containers = [c for c in containers if c.get("name") != "ssh-tunnel"] + containers.append(sidecar_def) + + task_volumes, mount_points = self._build_efs_volumes() + if mount_points: + existing_mounts = target.get("mountPoints") or [] + target["mountPoints"] = existing_mounts + mount_points + + family = self._make_family_name() + payload: dict[str, Any] = { + "family": family, + "networkMode": base.get("networkMode", "awsvpc"), + "requiresCompatibilities": base.get("requiresCompatibilities", ["FARGATE"]), + "cpu": str(max(int(base.get("cpu") or "256"), int(cfg.cpu))), + "memory": str(max(int(base.get("memory") or "512"), int(cfg.memory))), + "containerDefinitions": containers, + "ephemeralStorage": { + "sizeInGiB": max( + (base.get("ephemeralStorage") or {}).get("sizeInGiB", 20), cfg.ephemeral_storage_gib or 20 + ) + }, + } + for k in ("taskRoleArn", "executionRoleArn", "runtimePlatform", "volumes"): + if base.get(k) is not None: + payload[k] = base[k] + if task_volumes: + existing_vols = payload.get("volumes") or [] + payload["volumes"] = existing_vols + task_volumes + if cfg.execution_role_arn: + payload["executionRoleArn"] = cfg.execution_role_arn + if cfg.task_role_arn: + payload["taskRoleArn"] = cfg.task_role_arn + return self._do_register(payload) + + def _register_from_scratch( + self, *, image: str, command: list[str] | None, env: dict[str, str], sidecar_def: dict, log_cfg: dict | None + ) -> str: + cfg = self._cfg + if not cfg.execution_role_arn: + raise RuntimeError("execution_role_arn required when no base task definition provided") + container_def: dict[str, Any] = { + "name": cfg.container_name, + "essential": True, + "dependsOn": [{"containerName": "ssh-tunnel", "condition": "HEALTHY"}], + } + if image: + container_def["image"] = image + if command is not None: + container_def["command"] = command + if cfg.container_port: + container_def["portMappings"] = [{"containerPort": cfg.container_port, "protocol": "tcp"}] + if env: + container_def["environment"] = [{"name": k, "value": v} for k, v in sorted(env.items())] + if log_cfg: + container_def["logConfiguration"] = log_cfg + + task_volumes, mount_points = self._build_efs_volumes() + if mount_points: + container_def["mountPoints"] = mount_points + + payload: dict[str, Any] = { + "family": self._make_family_name(), + "networkMode": "awsvpc", + "requiresCompatibilities": ["FARGATE"], + "cpu": cfg.cpu, + "memory": cfg.memory, + "executionRoleArn": cfg.execution_role_arn, + "containerDefinitions": [container_def, sidecar_def], + } + if task_volumes: + payload["volumes"] = task_volumes + if cfg.task_role_arn: + payload["taskRoleArn"] = cfg.task_role_arn + if cfg.ephemeral_storage_gib: + payload["ephemeralStorage"] = {"sizeInGiB": cfg.ephemeral_storage_gib} + return self._do_register(payload) + + def _build_efs_volumes(self) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + """Build EFS volume definitions and mount points from spec.volumes.""" + task_volumes: list[dict[str, Any]] = [] + mount_points: list[dict[str, Any]] = [] + for i, vol in enumerate(self._spec.volumes): + if not vol.is_efs: + continue + vol_name = f"efs-{i}" + efs_cfg: dict[str, Any] = { + "fileSystemId": vol.efs_filesystem_id, + "transitEncryption": "ENABLED", + } + if vol.efs_access_point_id: + efs_cfg["authorizationConfig"] = { + "accessPointId": vol.efs_access_point_id, + "iam": "ENABLED", + } + elif vol.efs_root_directory: + efs_cfg["rootDirectory"] = vol.efs_root_directory + task_volumes.append({"name": vol_name, "efsVolumeConfiguration": efs_cfg}) + mount_points.append( + { + "sourceVolume": vol_name, + "containerPath": vol.container_path, + "readOnly": vol.readonly, + } + ) + return task_volumes, mount_points + + def _ssm_lookup_task_def(self, h: str) -> str | None: + _, _, ClientError = _require_aws_sdks() + param_name = f"/{self._cfg.ssm_project}/task-defs/{h}" + try: + resp = self._ssm.get_parameter(Name=param_name) + except ClientError as exc: + code = exc.response["Error"]["Code"] + if code == "ParameterNotFound": + return None + logger.warning("SSM GetParameter %s failed (%s); falling through to register", param_name, code) + return None + + arn = resp["Parameter"]["Value"] + try: + desc = self._ecs.describe_task_definition(taskDefinition=arn) + except ClientError as exc: + logger.warning( + "Cached task def %s no longer describable (%s); re-registering", + arn, + exc.response["Error"]["Code"], + ) + return None + if desc["taskDefinition"]["status"] != "ACTIVE": + logger.info("SSM cache entry %s is not ACTIVE; re-registering (hash %s)", arn, h) + return None + logger.info("Reusing task def from SSM cache: %s (hash %s)", arn, h) + return arn + + def _ssm_write_task_def(self, h: str, arn: str) -> None: + _, _, ClientError = _require_aws_sdks() + param_name = f"/{self._cfg.ssm_project}/task-defs/{h}" + try: + self._ssm.put_parameter(Name=param_name, Value=arn, Type="String", Overwrite=True) + logger.info("Wrote SSM task-def cache entry: %s -> %s", param_name, arn) + except ClientError as exc: + code = exc.response["Error"]["Code"] + logger.warning("SSM PutParameter %s failed (%s); cache entry not written", param_name, code) + + def _do_register(self, payload: dict[str, Any]) -> str: + h = _compute_task_def_hash(payload) + + while True: + with _task_def_cache_lock: + if h in _task_def_cache: + arn = _task_def_cache[h] + logger.info("Reusing cached task def %s (hash %s)", arn, h) + return arn + if h in _task_def_inflight: + event = _task_def_inflight[h] + else: + event = threading.Event() + _task_def_inflight[h] = event + break + event.wait() + + try: + arn = self._ssm_lookup_task_def(h) or self._register_task_def_fresh(payload, h) + with _task_def_cache_lock: + _task_def_cache[h] = arn + return arn + finally: + self._release_inflight(h, event) + + def _register_task_def_fresh(self, payload: dict[str, Any], h: str) -> str: + resp = _retry_with_backoff( + lambda: self._ecs.register_task_definition(**payload), + operation_name="register_task_definition", + max_retries=25, + ) + arn = resp["taskDefinition"]["taskDefinitionArn"] + logger.info("Registered task def: %s (hash %s)", arn, h) + self._ssm_write_task_def(h, arn) + return arn + + def _release_inflight(self, h: str, event: threading.Event) -> None: + with _task_def_cache_lock: + _task_def_inflight.pop(h, None) + event.set() + + def _make_family_name(self) -> str: + nonce = uuid.uuid4().hex[:12] + raw = f"{self._cfg.task_definition_family_prefix}-{_sanitize_id(self._spec.image or 'sandbox')}-{nonce}" + family = re.sub(r"[^A-Za-z0-9_-]", "_", raw)[:255] + if not family or not re.match(r"^[A-Za-z0-9]", family): + family = f"ecs_{family}" + return family + + # ── Run task + wait ────────────────────────────────────────────── + + def _run_task(self, task_def_arn: str) -> str: + cfg = self._cfg + run_kwargs: dict[str, Any] = { + "cluster": cfg.cluster, + "taskDefinition": task_def_arn, + "launchType": "FARGATE", + "networkConfiguration": { + "awsvpcConfiguration": { + "subnets": cfg.subnets, + "securityGroups": cfg.security_groups, + "assignPublicIp": "ENABLED" if cfg.assign_public_ip else "DISABLED", + } + }, + } + if self._runtime_container_env: + run_kwargs["overrides"] = { + "containerOverrides": [ + { + "name": cfg.container_name, + "environment": [ + {"name": k, "value": v} for k, v in sorted(self._runtime_container_env.items()) + ], + } + ] + } + has_efs = any(v.is_efs for v in self._spec.volumes) + if cfg.platform_version: + run_kwargs["platformVersion"] = cfg.platform_version + elif has_efs: + run_kwargs["platformVersion"] = "1.4.0" + + last_failures: Any = None + for attempt in range(1, cfg.run_task_max_retries + 1): + try: + resp = _retry_with_backoff( + lambda: self._ecs.run_task(**run_kwargs), operation_name="run_task", max_retries=3 + ) + except Exception as exc: + if not _is_retryable_error(exc) or attempt >= cfg.run_task_max_retries: + raise + delay = min(60.0, 2.0 ** min(6, attempt - 1)) + random.random() * 2 + logger.warning( + "run_task failed (%d/%d): %s — retry in %.1fs", attempt, cfg.run_task_max_retries, exc, delay + ) + time.sleep(delay) + continue + failures = resp.get("failures") or [] + if not failures: + tasks = resp.get("tasks") or [] + if not tasks: + raise RuntimeError("run_task returned no tasks") + task_arn = tasks[0]["taskArn"] + logger.info("Started ECS task: %s", task_arn) + return task_arn + last_failures = failures + reasons = " | ".join(str(f.get("reason", "")) for f in failures) + if not any(m in reasons.lower() for m in _RETRYABLE_MESSAGES) or attempt >= cfg.run_task_max_retries: + raise RuntimeError(f"run_task failures: {failures}") + delay = min(60.0, 2.0 ** min(6, attempt - 1)) + random.random() * 2 + logger.warning( + "run_task capacity issue (%d/%d): %s — retry in %.1fs", + attempt, + cfg.run_task_max_retries, + reasons, + delay, + ) + time.sleep(delay) + raise RuntimeError(f"run_task failed after {cfg.run_task_max_retries} retries: {last_failures}") + + def _wait_for_running(self) -> None: + cfg = self._cfg + start = time.monotonic() + poll = 5.0 + last_status = "" + while True: + elapsed = time.monotonic() - start + if elapsed > cfg.startup_timeout_sec: + raise TimeoutError(f"ECS task not RUNNING after {elapsed:.0f}s (last: {last_status})") + try: + resp = self._ecs.describe_tasks(cluster=cfg.cluster, tasks=[self._task_arn]) + except Exception as exc: + if _is_retryable_error(exc): + time.sleep(poll + random.random() * 3) + continue + raise + tasks = resp.get("tasks") or [] + if not tasks: + raise RuntimeError("ECS task disappeared") + status = tasks[0].get("lastStatus", "UNKNOWN") + if status == "RUNNING": + logger.info("ECS task RUNNING after %.0fs", elapsed) + return + if status == "STOPPED": + raise RuntimeError(f"ECS task stopped: {tasks[0].get('stoppedReason')}") + if status != last_status: + logger.info("ECS task %s (%.0fs)", status, elapsed) + last_status = status + time.sleep(poll + random.random() * 3) + poll = min(15.0, poll + 0.5) + + def _get_task_public_ip(self) -> str: + max_retries = 10 + for attempt in range(1, max_retries + 1): + try: + resp = self._ecs.describe_tasks(cluster=self._cfg.cluster, tasks=[self._task_arn]) + tasks = resp.get("tasks") or [] + if not tasks: + raise RuntimeError("Task not found") + eni_id = None + for att in tasks[0].get("attachments") or []: + if att.get("type") == "ElasticNetworkInterface": + for d in att.get("details") or []: + if d.get("name") == "networkInterfaceId": + eni_id = d["value"] + break + if eni_id: + break + if not eni_id: + for att in tasks[0].get("attachments") or []: + for d in att.get("details") or []: + if d.get("name") == "privateIPv4Address" and d.get("value"): + return d["value"] + raise RuntimeError("No ENI/IP yet") + iface = self._ec2.describe_network_interfaces(NetworkInterfaceIds=[eni_id])["NetworkInterfaces"][0] + pub = (iface.get("Association") or {}).get("PublicIp") + if pub: + logger.info("Container public IP: %s", pub) + return pub + priv = iface.get("PrivateIpAddress") + if priv: + logger.info("Container private IP: %s", priv) + return priv + raise RuntimeError(f"ENI {eni_id} has no IP") + except Exception as exc: + if attempt >= max_retries: + raise + if _is_retryable_error(exc): + time.sleep(min(15.0, 2.0**attempt + random.random())) + else: + logger.warning("get_task_ip attempt %d/%d: %s", attempt, max_retries, exc) + time.sleep(min(15.0, 3.0 + attempt * 2)) + raise RuntimeError("get_task_ip exhausted retries") + + @staticmethod + def _wait_for_ssh_ready(host: str, port: int, timeout: float) -> None: + deadline = time.monotonic() + timeout + logger.info("Waiting for SSH at %s:%d", host, port) + while time.monotonic() < deadline: + try: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.settimeout(5.0) + s.connect((host, port)) + s.settimeout(5.0) + if b"SSH" in s.recv(256): + logger.info("SSH ready at %s:%d", host, port) + return + except OSError: + pass + time.sleep(2.0) + raise TimeoutError(f"SSH not ready at {host}:{port} after {timeout:.0f}s") + + def _open_tunnel(self, sidecar: SshSidecarConfig) -> None: + assert self._task_ip is not None and self._ssh_key_file is not None + if sidecar.exec_server_port is not None: + self._ssh_tunnel = SshTunnel( + host=self._task_ip, + port=sidecar.sshd_port, + user="root", + key_file=self._ssh_key_file, + forward_port=sidecar.exec_server_port, + reverses=self._outside_endpoint_routing.reverse_specs, + ) + self._ssh_tunnel.open() + else: + remote_host, remote_port = self._outside_endpoint_routing.agent_tunnel_target() + assert self._ssh_tunnel_port is not None + self._agent_forward_port = _free_port() + container_port = self._cfg.container_port + if not container_port: + raise ValueError("container_port is required in agent-server mode") + self._ssh_tunnel = SshTunnel( + host=self._task_ip, + port=sidecar.sshd_port, + user="root", + key_file=self._ssh_key_file, + forwards=[f"{self._agent_forward_port}:localhost:{container_port}"], + reverses=[f"{self._ssh_tunnel_port}:{remote_host}:{remote_port}"], + local_port_override=self._agent_forward_port, + ) + self._ssh_tunnel.open() + + # ── Cleanup ────────────────────────────────────────────────────── + + def _cleanup(self) -> None: + if self._ssh_tunnel: + try: + self._ssh_tunnel.close() + except Exception: + logger.debug("Failed to close SSH tunnel", exc_info=True) + self._ssh_tunnel = None + if self._task_arn and self._ecs: + try: + _retry_with_backoff( + lambda: self._ecs.stop_task( + cluster=self._cfg.cluster, task=self._task_arn, reason="sandbox cleanup" + ), + operation_name="stop_task", + max_retries=10, + ) + logger.info("Stopped ECS task: %s", self._task_arn) + except Exception as exc: + logger.warning("Failed to stop task %s: %s", self._task_arn, exc) + if self._ssh_key_file: + try: + os.remove(self._ssh_key_file) + except Exception: + logger.debug("Failed to remove SSH key file %s", self._ssh_key_file, exc_info=True) + self._ssh_key_file = None + + def _sync_stop(self) -> None: + """Synchronous stop for emergency cleanup (atexit handler).""" + if self._stopped: + return + self._stopped = True + self._cleanup() + + def _require_exec_client(self) -> None: + if self._exec_client is None: + raise RuntimeError( + "exec()/upload()/download() require exec-server mode " + "(ssh_sidecar.exec_server_port). In agent-server mode use sandbox.ssh_tunnel." + ) + + async def _upload_via_s3(self, paths: list[Path], dest_dir: str) -> None: + cfg = self._cfg + if not cfg.s3_bucket: + raise ValueError("s3_bucket is required for S3 staging") + + def _pack() -> bytes: + buf = io.BytesIO() + with tarfile.open(fileobj=buf, mode="w:gz") as tar: + for p in paths: + if p.is_file(): + tar.add(str(p), arcname=p.name) + elif p.is_dir(): + for child in p.rglob("*"): + if child.is_file(): + tar.add(str(child), arcname=str(child.relative_to(p))) + buf.seek(0) + return buf.read() + + body = await asyncio.to_thread(_pack) + boto3, *_ = _require_aws_sdks() + s3 = boto3.client("s3", region_name=cfg.region) + prefix = cfg.s3_prefix or "ecs-sandbox" + nonce = uuid.uuid4().hex[:12] + key = f"{prefix}/{self._run_id}/upload-{nonce}.tar.gz" + await asyncio.to_thread( + _retry_with_backoff, + lambda: s3.put_object(Bucket=cfg.s3_bucket, Key=key, Body=body), + operation_name="s3.put_object(upload)", + max_retries=5, + ) + url = await asyncio.to_thread( + s3.generate_presigned_url, + "get_object", + Params={"Bucket": cfg.s3_bucket, "Key": key}, + ExpiresIn=21600, + ) + dl_cmd = ( + f"mkdir -p {shlex.quote(dest_dir)} && TGZ=/tmp/_upload_$$.tar.gz && " + f"( curl -sf -L --max-time 300 -o $TGZ {shlex.quote(url)} 2>/dev/null || " + f"python3 -c 'import urllib.request as u,sys;u.urlretrieve(sys.argv[1],sys.argv[2])' " + f"{shlex.quote(url)} $TGZ ) && " + f"tar xzf $TGZ -C {shlex.quote(dest_dir)} && rm -f $TGZ && echo ok" + ) + result = await self._exec_client.exec(dl_cmd, timeout=360) # type: ignore[union-attr] + if "ok" not in result.stdout: + raise RuntimeError( + f"S3 upload extraction failed (rc={result.return_code}): {result.stderr or result.stdout}" + ) + + # ── Atexit cleanup ─────────────────────────────────────────────── + + def _register_for_cleanup(self) -> None: + global _atexit_registered + with _cleanup_lock: + _active_sandboxes[id(self)] = self + if not _atexit_registered: + atexit.register(_emergency_cleanup) + _atexit_registered = True + + def _unregister_from_cleanup(self) -> None: + with _cleanup_lock: + _active_sandboxes.pop(id(self), None) diff --git a/nemo_gym/sandbox/providers/ecs_fargate/provider.py b/nemo_gym/sandbox/providers/ecs_fargate/provider.py new file mode 100644 index 0000000000..6c4013702f --- /dev/null +++ b/nemo_gym/sandbox/providers/ecs_fargate/provider.py @@ -0,0 +1,237 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""ECS Fargate sandbox provider. + +Adapts the lifted :mod:`engine` (one stateful ``EcsFargateSandbox`` per +sandbox) to Gym's stateless ``SandboxProvider`` contract: per-sandbox engine +state lives in ``SandboxHandle.raw``; the provider methods delegate to it. + +The model endpoint is assumed reachable from inside the sandbox directly, or +routed via the SSH reverse tunnel when ``outside_endpoints`` are supplied +through ``spec.provider_options``. See the package README for the planned +Teleport follow-up. +""" + +from __future__ import annotations + +from dataclasses import replace +from pathlib import Path +from typing import Any + +from nemo_gym.sandbox.providers.base import ( + SandboxCreateError, + SandboxExecResult, + SandboxHandle, + SandboxSpec, + SandboxStatus, +) +from nemo_gym.sandbox.providers.ecs_fargate import engine + + +def _outside_endpoints(spec: SandboxSpec) -> list[engine.OutsideEndpoint]: + raw = spec.provider_options.get("outside_endpoints") or [] + endpoints = [] + for item in raw: + if isinstance(item, engine.OutsideEndpoint): + endpoints.append(item) + else: + endpoints.append(engine.OutsideEndpoint(url=item["url"], env_var=item["env_var"])) + return endpoints + + +def _volumes(spec: SandboxSpec) -> list[engine.VolumeMount]: + raw = spec.provider_options.get("volumes") or [] + volumes = [] + for item in raw: + if isinstance(item, engine.VolumeMount): + volumes.append(item) + else: + volumes.append(engine.VolumeMount(**item)) + return volumes + + +def _engine_spec(spec: SandboxSpec) -> engine.SandboxSpec: + if spec.image is None: + raise SandboxCreateError("ECS Fargate sandbox requires SandboxSpec.image") + entrypoint = " ".join(spec.entrypoint) if spec.entrypoint else None + return engine.SandboxSpec( + image=spec.image, + workdir=spec.workdir or "/workspace", + env=dict(spec.env), + files=dict(spec.files), + entrypoint=entrypoint, + volumes=_volumes(spec), + environment_dir=spec.provider_options.get("environment_dir"), + ) + + +class EcsFargateProvider: + """Run sandboxes as AWS ECS Fargate tasks behind an SSH sidecar.""" + + name = "ecs_fargate" + + def __init__(self, **config: Any) -> None: + self._cfg = engine_config_from_mapping(config) + + async def create(self, spec: SandboxSpec) -> SandboxHandle: + cfg = self._cfg + if spec.ready_timeout_s is not None: + cfg = replace(cfg, startup_timeout_sec=float(spec.ready_timeout_s)) + sandbox = engine.EcsFargateSandbox(_engine_spec(spec), ecs_config=cfg) + try: + await sandbox.start(outside_endpoints=_outside_endpoints(spec)) + except Exception as e: # noqa: BLE001 — uniform create failure surface + raise SandboxCreateError(f"ECS Fargate create failed: {e}") from e + return SandboxHandle( + sandbox_id=sandbox.task_arn or "", + provider_name=self.name, + raw=sandbox, + ) + + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + sandbox: engine.EcsFargateSandbox = handle.raw + result = await sandbox.exec( + command, + timeout_sec=180 if timeout_s is None else float(timeout_s), + cwd=cwd, + env=env, + user=user, + ) + return SandboxExecResult( + stdout=result.stdout, + stderr=result.stderr, + return_code=result.return_code, + ) + + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + sandbox: engine.EcsFargateSandbox = handle.raw + await sandbox.upload(Path(source_path), target_path) + + async def download_file(self, handle: SandboxHandle, source_path: str, target_path: Path) -> None: + sandbox: engine.EcsFargateSandbox = handle.raw + await sandbox.download(source_path, Path(target_path)) + + async def status(self, handle: SandboxHandle) -> SandboxStatus: + sandbox: engine.EcsFargateSandbox = handle.raw + return SandboxStatus.RUNNING if sandbox.is_running else SandboxStatus.STOPPED + + async def close(self, handle: SandboxHandle, *, delete: bool = False) -> None: + sandbox: engine.EcsFargateSandbox = handle.raw + await sandbox.stop() + + async def aclose(self) -> None: + return None + + +def engine_config_from_mapping(config: dict[str, Any]) -> engine.EcsFargateConfig: + """Build an ``EcsFargateConfig`` from provider kwargs. + + When ``region`` is given but ``cluster`` is omitted, infrastructure is + auto-discovered from SSM (written by the reference Terraform); explicit + kwargs always win over SSM. Mirrors NEL's orchestrator resolution. + """ + config = dict(config) + region = config.get("region") + cluster = config.get("cluster") + ssm_project = config.get("ssm_project", engine.DEFAULT_SSM_PROJECT) + + ssm: dict[str, Any] = {} + if region is not None and cluster is None: + ssm = engine.resolve_ecs_config_from_ssm(region, ssm_project) + + def pick(key: str, default: Any = None) -> Any: + val = config.get(key) + if val is not None: + return val + return ssm.get(key, default) + + ssh_sidecar = _sidecar_config(config.get("ssh_sidecar"), ssm.get("ssh_sidecar", {})) + + return engine.EcsFargateConfig( + region=region, + cluster=pick("cluster", ""), + subnets=config.get("subnets") or ssm.get("subnets", []), + security_groups=config.get("security_groups") or ssm.get("security_groups", []), + assign_public_ip=pick("assign_public_ip", True), + task_definition=config.get("task_definition"), + task_definition_family_prefix=config.get("task_definition_family_prefix", "ecs-sandbox"), + image_template=config.get("image_template"), + container_name=config.get("container_name", "main"), + container_port=config.get("container_port"), + cpu=str(config.get("cpu", "4096")), + memory=str(config.get("memory", "8192")), + ephemeral_storage_gib=config.get("ephemeral_storage_gib"), + platform_version=config.get("platform_version"), + execution_role_arn=pick("execution_role_arn"), + task_role_arn=pick("task_role_arn"), + extra_env=config.get("extra_env"), + log_group=pick("log_group"), + log_stream_prefix=config.get("log_stream_prefix"), + max_task_lifetime_sec=config.get("max_task_lifetime_sec") or 14400, + startup_timeout_sec=float(config.get("startup_timeout_sec", 300.0)), + ssh_sidecar=ssh_sidecar, + s3_bucket=pick("s3_bucket"), + s3_prefix=config.get("s3_prefix"), + ecr_repository=pick("ecr_repository"), + environment_dir=config.get("environment_dir"), + codebuild_project=config.get("codebuild_project"), + codebuild_service_role=pick("codebuild_service_role"), + codebuild_compute_type=config.get("codebuild_compute_type") or "BUILD_GENERAL1_MEDIUM", + codebuild_build_timeout=config.get("codebuild_build_timeout") or 60, + auto_mirror=config.get("auto_mirror", True), + dockerhub_secret_arn=pick("dockerhub_secret_arn"), + efs_filesystem_id=pick("efs_filesystem_id"), + efs_access_point_id=pick("efs_access_point_id"), + ssm_project=ssm_project, + ) + + +def _sidecar_config(yaml_sidecar: Any, ssm_ssh: dict[str, Any]) -> engine.SshSidecarConfig | None: + if yaml_sidecar is not None: + sc = dict(yaml_sidecar) if isinstance(yaml_sidecar, dict) else yaml_sidecar + if isinstance(sc, engine.SshSidecarConfig): + return sc + pub = sc.get("public_key_secret_arn") or ssm_ssh.get("public_key_secret_arn", "") + priv = sc.get("private_key_secret_arn") or ssm_ssh.get("private_key_secret_arn", "") + if not pub or not priv: + raise ValueError( + "ssh_sidecar.public_key_secret_arn and ssh_sidecar.private_key_secret_arn " + "are required (set explicitly or auto-discovered from SSM)." + ) + return engine.SshSidecarConfig( + sshd_port=sc.get("sshd_port", engine.DEFAULT_SSHD_PORT), + ssh_ready_timeout_sec=sc.get("ssh_ready_timeout_sec", 300.0), + public_key_secret_arn=pub, + private_key_secret_arn=priv, + image=sc.get("image"), + exec_server_port=sc.get("exec_server_port", engine.DEFAULT_EXEC_SERVER_PORT), + ) + if ssm_ssh.get("public_key_secret_arn") and ssm_ssh.get("private_key_secret_arn"): + return engine.SshSidecarConfig( + sshd_port=ssm_ssh.get("sshd_port", engine.DEFAULT_SSHD_PORT), + public_key_secret_arn=ssm_ssh["public_key_secret_arn"], + private_key_secret_arn=ssm_ssh["private_key_secret_arn"], + exec_server_port=ssm_ssh.get("exec_server_port", engine.DEFAULT_EXEC_SERVER_PORT), + ) + return None diff --git a/nemo_gym/sandbox/providers/opensandbox/__init__.py b/nemo_gym/sandbox/providers/opensandbox/__init__.py new file mode 100644 index 0000000000..c676205dd4 --- /dev/null +++ b/nemo_gym/sandbox/providers/opensandbox/__init__.py @@ -0,0 +1,38 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""OpenSandbox provider package.""" + +from nemo_gym.sandbox.providers.opensandbox.provider import ( + OpenSandboxConnectionConfig, + OpenSandboxCreateConfig, + OpenSandboxCreateError, + OpenSandboxCreateTimeoutError, + OpenSandboxCreateVerificationError, + OpenSandboxOperationConfig, + OpenSandboxProbeConfig, + OpenSandboxProvider, +) + + +__all__ = [ + "OpenSandboxConnectionConfig", + "OpenSandboxCreateConfig", + "OpenSandboxCreateError", + "OpenSandboxCreateTimeoutError", + "OpenSandboxCreateVerificationError", + "OpenSandboxOperationConfig", + "OpenSandboxProbeConfig", + "OpenSandboxProvider", +] diff --git a/nemo_gym/sandbox/providers/opensandbox/provider.py b/nemo_gym/sandbox/providers/opensandbox/provider.py new file mode 100644 index 0000000000..1ec7543b79 --- /dev/null +++ b/nemo_gym/sandbox/providers/opensandbox/provider.py @@ -0,0 +1,954 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""OpenSandbox provider implementation.""" + +import asyncio +import logging +import re +import shlex +from collections.abc import Mapping +from dataclasses import dataclass, replace +from datetime import timedelta +from pathlib import Path +from typing import Any, Awaitable, Callable + +from nemo_gym.sandbox.providers.base import ( + SandboxCreateError, + SandboxCreateVerificationError, + SandboxExecResult, + SandboxHandle, + SandboxSpec, + SandboxStatus, +) + + +LOGGER = logging.getLogger(__name__) + + +class OpenSandboxCreateError(SandboxCreateError): + """Raised when OpenSandbox cannot create a sandbox.""" + + +class OpenSandboxCreateTimeoutError(OpenSandboxCreateError): + """Raised when OpenSandbox sandbox creation exceeds the client timeout.""" + + +class OpenSandboxCreateVerificationError(SandboxCreateVerificationError): + """Raised when a newly-created sandbox cannot execute a probe command.""" + + +RETRYABLE_HTTP_STATUS_CODES = {408, 409, 425, 429, 500, 502, 503, 504} +RETRYABLE_ERROR_MARKERS = ( + "all connection attempts failed", + "connection refused", + "connection reset", + "gateway timeout", + "http 408", + "http 409", + "http 425", + "http 429", + "http 500", + "http 502", + "http 503", + "http 504", + "incomplete chunked read", + "peer closed connection", + "pod ip is not yet available", + "pod may still be starting", + "errimagepull", + "get endpoint for sandbox", + "imagepullbackoff", + "pod failed", + "podfailed", + "remote protocol error", + "service unavailable", + "server disconnected", + "status code: 408", + "status code: 409", + "status code: 425", + "status code: 429", + "status code: 500", + "status code: 502", + "status code: 503", + "status code: 504", + "temporarily unavailable", + "timed out", + "timeout", +) +METADATA_VALUE_RE = re.compile(r"[^A-Za-z0-9_.-]+") +DEFAULT_IMAGE_PULL_POLICY = "IfNotPresent" +IMAGE_PULL_POLICY_EXTENSION_KEY = "imagePullPolicy" +IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY = "opensandbox.extensions.image-pull-policy" +PROVIDER_OPTION_PLATFORM = "platform" +PROVIDER_OPTION_SKIP_HEALTH_CHECK = "skip_health_check" +PROVIDER_OPTION_SNAPSHOT_ID = "snapshot_id" +PROVIDER_OPTION_VOLUMES = "volumes" +VALID_IMAGE_PULL_POLICIES = {"Always", "IfNotPresent", "Never"} +STATUS_CODE_RE = re.compile(r"(?:status code|http)\D+(\d{3})", re.IGNORECASE) + + +def validate_image_pull_policy(image_pull_policy: str) -> str: + """Validate a Kubernetes-compatible container image pull policy.""" + if image_pull_policy not in VALID_IMAGE_PULL_POLICIES: + allowed = ", ".join(sorted(VALID_IMAGE_PULL_POLICIES)) + raise ValueError(f"image_pull_policy must be one of: {allowed}") + return image_pull_policy + + +def _require_opensandbox_sdk() -> tuple[Any, Any, Any, Any, Any]: + try: + from opensandbox import Sandbox + from opensandbox.config import ConnectionConfig + from opensandbox.models.execd import RunCommandOpts + from opensandbox.models.sandboxes import PlatformSpec, Volume + except ModuleNotFoundError as e: + raise ModuleNotFoundError( + "OpenSandbox SDK is required for the opensandbox sandbox provider. " + "Install nemo-gym[sandbox] in the runtime image before using " + "env.sandbox.provider.name=opensandbox." + ) from e + + return Sandbox, ConnectionConfig, RunCommandOpts, PlatformSpec, Volume + + +def _require_tenacity() -> tuple[Any, Any, Any, Any]: + try: + from tenacity import AsyncRetrying, retry_if_exception, stop_after_attempt, wait_random_exponential + except ModuleNotFoundError as e: + raise ModuleNotFoundError( + "tenacity is required for OpenSandbox retry handling. Install nemo-gym[sandbox] before using " + "env.sandbox.provider.name=opensandbox." + ) from e + + return AsyncRetrying, retry_if_exception, stop_after_attempt, wait_random_exponential + + +def _has_retryable_error_marker(exception: BaseException) -> bool: + message = str(exception).lower() + return any(marker in message for marker in RETRYABLE_ERROR_MARKERS) + + +def _exception_status_code(exception: BaseException) -> int | None: + status_code = getattr(exception, "status_code", None) + if isinstance(status_code, int): + return status_code + + match = STATUS_CODE_RE.search(str(exception)) + if match is None: + return None + return int(match.group(1)) + + +def _sdk_error_attributes( + exception: BaseException, + *, + operation: str, + sandbox_id: str, + attempt_number: int | None = None, + max_attempts: int | None = None, + sleep_s: float | None = None, +) -> dict[str, Any]: + attrs: dict[str, Any] = { + "provider": OpenSandboxProvider.name, + "operation": operation, + "sandbox_id": sandbox_id, + "error_type": type(exception).__name__, + "error_message": str(exception)[:500], + } + status_code = _exception_status_code(exception) + if status_code is not None: + attrs["status_code"] = status_code + if attempt_number is not None: + attrs["attempt_number"] = attempt_number + if max_attempts is not None: + attrs["max_attempts"] = max_attempts + if sleep_s is not None: + attrs["next_sleep_s"] = sleep_s + return attrs + + +def _is_retryable_create_error(exception: BaseException) -> bool: + """Return whether a sandbox create failure is likely transient.""" + if isinstance(exception, SandboxCreateVerificationError): + return True + if isinstance(exception, SandboxCreateError): + return True + if isinstance(exception, (ConnectionError, OSError, TimeoutError)): + return True + + try: + from opensandbox.exceptions import ( + InvalidArgumentException, + SandboxApiException, + SandboxException, + SandboxInternalException, + SandboxReadyTimeoutException, + SandboxUnhealthyException, + ) + except ModuleNotFoundError: + return _has_retryable_error_marker(exception) + + if isinstance(exception, InvalidArgumentException): + return False + if isinstance( + exception, + ( + SandboxInternalException, + SandboxReadyTimeoutException, + SandboxUnhealthyException, + ), + ): + return True + if isinstance(exception, SandboxApiException): + status_code = getattr(exception, "status_code", None) + if status_code in RETRYABLE_HTTP_STATUS_CODES: + return True + if status_code is not None and status_code < 500: + return False + if not isinstance(exception, SandboxException): + return _has_retryable_error_marker(exception) + + return _has_retryable_error_marker(exception) + + +def _is_retryable_sdk_operation_error(exception: BaseException, seen: set[int] | None = None) -> bool: + """Return whether an SDK operation can be retried.""" + if isinstance(exception, TimeoutError): + return False + seen = set() if seen is None else seen + exception_id = id(exception) + if exception_id in seen: + return False + seen.add(exception_id) + if isinstance(exception, (ConnectionError, OSError)): + return True + if _is_retryable_create_error(exception): + return True + cause = exception.__cause__ + if isinstance(cause, BaseException): + return _is_retryable_sdk_operation_error(cause, seen) + return False + + +def _is_missing_sandbox_delete_error(exception: BaseException) -> bool: + message = str(exception).lower() + return "sandbox" in message and "not found" in message + + +def _log_create_retry(retry_state: Any) -> None: + exception = retry_state.outcome.exception() if retry_state.outcome else None + sleep_s = retry_state.next_action.sleep if retry_state.next_action else None + LOGGER.warning( + "Retrying OpenSandbox sandbox create after attempt %s; next_sleep_s=%s; error=%r", + retry_state.attempt_number, + sleep_s, + exception, + ) + + +def _log_operation_retry(retry_state: Any) -> None: + exception = retry_state.outcome.exception() if retry_state.outcome else None + sleep_s = retry_state.next_action.sleep if retry_state.next_action else None + LOGGER.warning( + "Retrying OpenSandbox SDK operation after attempt %s; next_sleep_s=%s; error=%r", + retry_state.attempt_number, + sleep_s, + exception, + ) + + +def _string_map(values: dict[str, Any]) -> dict[str, str]: + return {str(key): str(value) for key, value in values.items()} + + +def _metadata_value(value: Any) -> str: + normalized = METADATA_VALUE_RE.sub("_", str(value)).strip("._-") + normalized = normalized[:63].strip("._-") + return normalized or "metadata" + + +def _metadata_map(values: dict[str, Any]) -> dict[str, str]: + return {str(key): _metadata_value(value) for key, value in values.items()} + + +def _normalize_spec(spec: SandboxSpec) -> SandboxSpec: + return replace( + spec, + env=_string_map(spec.env), + metadata=_metadata_map(spec.metadata), + resources=_string_map(spec.resources), + ) + + +def _to_platform_spec(platform: dict[str, Any]) -> Any: + _, _, _, PlatformSpec, _ = _require_opensandbox_sdk() + return PlatformSpec(**platform) + + +def _to_volumes(volumes: list[Mapping[str, Any]]) -> list[Any]: + _, _, _, _, Volume = _require_opensandbox_sdk() + return [Volume(**dict(volume)) for volume in volumes] + + +def _spec_volumes(spec: SandboxSpec) -> list[Mapping[str, Any]] | None: + return spec.provider_options.get(PROVIDER_OPTION_VOLUMES) + + +def _spec_extensions(spec: SandboxSpec) -> dict[str, str]: + value = spec.provider_options.get("extensions", {}) + if not isinstance(value, Mapping): + raise TypeError("OpenSandbox provider option 'extensions' must be a mapping") + return _string_map(dict(value)) + + +def _provider_option_bool(provider_options: dict[str, Any], key: str) -> bool | None: + value = provider_options.get(key) + if value is None: + return None + if not isinstance(value, bool): + raise TypeError(f"OpenSandbox provider option {key!r} must be a bool") + return value + + +def _to_sandbox_status(state: Any) -> SandboxStatus: + normalized = str(state or "").lower() + if normalized in {"active", "ready", "running"}: + return SandboxStatus.RUNNING + if normalized in {"creating", "initializing", "pending", "starting"}: + return SandboxStatus.STARTING + if normalized in {"completed", "deleted", "exited", "stopped", "terminated"}: + return SandboxStatus.STOPPED + if normalized in {"crashed", "error", "failed", "unhealthy"}: + return SandboxStatus.ERROR + return SandboxStatus.UNKNOWN + + +@dataclass(frozen=True) +class OpenSandboxConnectionConfig: + """OpenSandbox server connection settings.""" + + domain: str | None = None + api_key: str | None = None + protocol: str | None = None + request_timeout_s: int | None = None + use_server_proxy: bool = False + + +@dataclass(frozen=True) +class OpenSandboxCreateConfig: + """OpenSandbox create/reconnect retry settings.""" + + request_timeout_s: int | None = None + timeout_s: float | None = None + retries: int = 2 + retry_delay_s: float = 5.0 + retry_max_delay_s: float = 60.0 + image_pull_policy: str | None = DEFAULT_IMAGE_PULL_POLICY + skip_health_check: bool = False + connect_attempt_timeout_s: float = 30.0 + connect_poll_s: float = 2.0 + + def __post_init__(self) -> None: + if self.image_pull_policy is not None: + validate_image_pull_policy(self.image_pull_policy) + if self.timeout_s is not None and self.timeout_s <= 0: + raise ValueError("create.timeout_s must be > 0") + if self.retries < 0: + raise ValueError("create.retries must be >= 0") + if self.retry_delay_s < 0: + raise ValueError("create.retry_delay_s must be >= 0") + if self.retry_max_delay_s < 0: + raise ValueError("create.retry_max_delay_s must be >= 0") + if self.connect_attempt_timeout_s <= 0: + raise ValueError("create.connect_attempt_timeout_s must be > 0") + if self.connect_poll_s <= 0: + raise ValueError("create.connect_poll_s must be > 0") + + +@dataclass(frozen=True) +class OpenSandboxProbeConfig: + """Post-create probe settings.""" + + command: str | None = "printf nemo-gym-sandbox-ready" + expected_stdout: str | None = "nemo-gym-sandbox-ready" + timeout_s: int = 30 + deadline_s: float | None = None + stable_count: int = 1 + stable_delay_s: float = 0.0 + + def __post_init__(self) -> None: + if self.command is not None and self.timeout_s <= 0: + raise ValueError("probe.timeout_s must be > 0") + if self.deadline_s is not None and self.deadline_s <= 0: + raise ValueError("probe.deadline_s must be > 0") + if self.stable_count < 1: + raise ValueError("probe.stable_count must be >= 1") + if self.stable_delay_s < 0: + raise ValueError("probe.stable_delay_s must be >= 0") + + +@dataclass(frozen=True) +class OpenSandboxOperationConfig: + """Retry and timeout settings for SDK operations after create.""" + + retries: int = 3 + retry_delay_s: float = 1.0 + retry_max_delay_s: float = 15.0 + command_retries: int | None = None + close_timeout_s: float | None = 30.0 + + def __post_init__(self) -> None: + if self.retries < 0: + raise ValueError("operations.retries must be >= 0") + if self.retry_delay_s < 0: + raise ValueError("operations.retry_delay_s must be >= 0") + if self.retry_max_delay_s < 0: + raise ValueError("operations.retry_max_delay_s must be >= 0") + if self.command_retries is not None and self.command_retries < 0: + raise ValueError("operations.command_retries must be >= 0") + if self.close_timeout_s is not None and self.close_timeout_s <= 0: + raise ValueError("operations.close_timeout_s must be > 0") + + +def _coerce_config(value: Any, config_cls: type[Any]) -> Any: + if value is None: + return config_cls() + if isinstance(value, config_cls): + return value + if isinstance(value, Mapping): + return config_cls(**value) + raise TypeError(f"{config_cls.__name__} must be a mapping or {config_cls.__name__} instance") + + +class OpenSandboxProvider: + """Provider backed by the OpenSandbox SDK/server API.""" + + name = "opensandbox" + + def __init__( + self, + *, + connection: OpenSandboxConnectionConfig | Mapping[str, Any] | None = None, + create: OpenSandboxCreateConfig | Mapping[str, Any] | None = None, + probe: OpenSandboxProbeConfig | Mapping[str, Any] | None = None, + operations: OpenSandboxOperationConfig | Mapping[str, Any] | None = None, + ) -> None: + self._connection = _coerce_config(connection, OpenSandboxConnectionConfig) + self._create = _coerce_config(create, OpenSandboxCreateConfig) + self._probe = _coerce_config(probe, OpenSandboxProbeConfig) + self._operations = _coerce_config(operations, OpenSandboxOperationConfig) + + def _with_default_image_pull_policy(self, spec: SandboxSpec) -> SandboxSpec: + """Ensure SDK create requests carry the desired image pull policy.""" + if self._create.image_pull_policy is None: + return spec + + provider_options = dict(spec.provider_options) + extensions = _spec_extensions(spec) + image_pull_policy = extensions.get(IMAGE_PULL_POLICY_EXTENSION_KEY) or extensions.get( + IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY + ) + if image_pull_policy is None: + image_pull_policy = self._create.image_pull_policy + image_pull_policy = validate_image_pull_policy(image_pull_policy) + extensions.setdefault(IMAGE_PULL_POLICY_EXTENSION_KEY, image_pull_policy) + extensions.setdefault(IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY, image_pull_policy) + provider_options["extensions"] = extensions + return replace(spec, provider_options=provider_options) + + def _connection_config( + self, + request_timeout_s: int | float | None = None, + ) -> Any: + _, ConnectionConfig, _, _, _ = _require_opensandbox_sdk() + kwargs: dict[str, Any] = {} + if self._connection.domain is not None: + kwargs["domain"] = self._connection.domain + if self._connection.api_key is not None: + kwargs["api_key"] = self._connection.api_key + if self._connection.protocol is not None: + kwargs["protocol"] = self._connection.protocol + if request_timeout_s is None: + request_timeout_s = self._connection.request_timeout_s + if request_timeout_s is not None: + kwargs["request_timeout"] = timedelta(seconds=request_timeout_s) + if self._connection.use_server_proxy: + kwargs["use_server_proxy"] = True + return ConnectionConfig(**kwargs) + + async def aclose(self) -> None: + """Close provider-owned resources.""" + return None + + async def _await_sdk_call( + self, + awaitable: Any, + *, + operation: str, + sandbox_id: str, + timeout_s: float | None, + ) -> Any: + if timeout_s is None: + return await awaitable + + try: + return await asyncio.wait_for(awaitable, timeout=timeout_s) + except asyncio.TimeoutError as e: + raise TimeoutError( + f"Timed out during OpenSandbox {operation} after {timeout_s:g}s; sandbox_id={sandbox_id!r}" + ) from e + + async def _await_sdk_operation( + self, + operation_factory: Callable[[], Awaitable[Any]], + *, + operation: str, + sandbox_id: str, + timeout_s: float | None, + retries: int | None = None, + ) -> Any: + AsyncRetrying, retry_if_exception, stop_after_attempt, wait_random_exponential = _require_tenacity() + retry_count = self._operations.retries if retries is None else retries + max_attempts = retry_count + 1 + + def _before_sleep(retry_state: Any) -> None: + _log_operation_retry(retry_state) + + retry_policy = AsyncRetrying( + retry=retry_if_exception(_is_retryable_sdk_operation_error), + stop=stop_after_attempt(max_attempts), + wait=wait_random_exponential( + multiplier=self._operations.retry_delay_s, + max=self._operations.retry_max_delay_s, + ), + before_sleep=_before_sleep, + reraise=True, + ) + async for attempt in retry_policy: + with attempt: + return await self._await_sdk_call( + operation_factory(), + operation=operation, + sandbox_id=sandbox_id, + timeout_s=timeout_s, + ) + + raise RuntimeError("OpenSandbox SDK operation retry loop did not run") + + async def _verify_created_handle(self, handle: SandboxHandle) -> None: + if self._probe.command is None: + return + + loop = asyncio.get_running_loop() + deadline_s = self._probe.deadline_s or float(self._probe.timeout_s) + deadline = loop.time() + deadline_s + successful_probes = 0 + attempt_number = 0 + last_exception: BaseException | None = None + + while successful_probes < self._probe.stable_count: + remaining_s = deadline - loop.time() + if remaining_s <= 0: + error = OpenSandboxCreateVerificationError( + "OpenSandbox sandbox failed create probe command before " + "the startup deadline; " + f"sandbox_id={handle.sandbox_id!r}, " + f"command={self._probe.command!r}, " + f"successful_probes={successful_probes}/{self._probe.stable_count}, " + f"attempts={attempt_number}, deadline_s={deadline_s:g}" + ) + raise error from last_exception + + attempt_number += 1 + if self._probe.deadline_s is None: + command_timeout_s = float(self._probe.timeout_s) + else: + command_timeout_s = min(float(self._probe.timeout_s), remaining_s) + try: + result = await asyncio.wait_for( + self._exec( + handle, + self._probe.command, + timeout_s=command_timeout_s, + user="root", + ), + timeout=command_timeout_s, + ) + except asyncio.CancelledError: + raise + except Exception as e: + last_exception = e + successful_probes = 0 + sleep_s = min(self._create.connect_poll_s, max(deadline - loop.time(), 0.0)) + if sleep_s > 0: + await asyncio.sleep(sleep_s) + continue + + stdout = result.stdout or "" + expected = self._probe.expected_stdout + if result.return_code != 0 or (expected is not None and expected not in stdout): + last_exception = OpenSandboxCreateVerificationError( + "OpenSandbox sandbox create probe command returned an " + f"unexpected result; sandbox_id={handle.sandbox_id!r}, " + f"return_code={result.return_code}, expected_stdout={expected!r}, " + f"stdout={stdout[:200]!r}, stderr={(result.stderr or '')[:200]!r}, " + f"probe={successful_probes + 1}/{self._probe.stable_count}" + ) + successful_probes = 0 + sleep_s = min(self._create.connect_poll_s, max(deadline - loop.time(), 0.0)) + if sleep_s > 0: + await asyncio.sleep(sleep_s) + continue + + successful_probes += 1 + if successful_probes < self._probe.stable_count and self._probe.stable_delay_s: + await asyncio.sleep(self._probe.stable_delay_s) + + async def _cleanup_failed_create_handle(self, handle: SandboxHandle) -> None: + try: + await self.close(handle, delete=True) + except Exception as e: + LOGGER.warning( + "Failed to clean up OpenSandbox sandbox after create probe failure; sandbox_id=%s; error=%r", + handle.sandbox_id, + e, + ) + + async def _connect_after_create(self, handle: SandboxHandle, spec: SandboxSpec) -> SandboxHandle: + """Reconnect after SDK create so follow-up calls use a fresh SDK handle.""" + timeout_s = spec.ready_timeout_s + if timeout_s is None: + timeout_s = self._create.timeout_s + if timeout_s is None: + timeout_s = self._create.connect_attempt_timeout_s + + Sandbox, _, _, _, _ = _require_opensandbox_sdk() + loop = asyncio.get_running_loop() + deadline = loop.time() + float(timeout_s) + last_exception: BaseException | None = None + + while True: + remaining_s = deadline - loop.time() + if remaining_s <= 0: + error = OpenSandboxCreateTimeoutError( + "Timed out connecting to OpenSandbox sandbox after SDK create; " + f"sandbox_id={handle.sandbox_id!r}, timeout_s={timeout_s:g}" + ) + raise error from last_exception + + attempt_timeout_s = min(self._create.connect_attempt_timeout_s, remaining_s) + try: + sandbox = await asyncio.wait_for( + Sandbox.connect( + handle.sandbox_id, + connection_config=self._connection_config(request_timeout_s=attempt_timeout_s), + connect_timeout=timedelta(seconds=attempt_timeout_s), + skip_health_check=True, + ), + timeout=attempt_timeout_s, + ) + return SandboxHandle(sandbox_id=str(sandbox.id), provider_name=self.name, raw=sandbox) + except asyncio.CancelledError: + raise + except BaseException as e: + last_exception = e + if not _is_retryable_create_error(e): + raise + sleep_s = min(self._create.connect_poll_s, max(deadline - loop.time(), 0.0)) + if sleep_s > 0: + await asyncio.sleep(sleep_s) + + async def _create_once(self, spec: SandboxSpec) -> SandboxHandle: + """Create a sandbox through ``opensandbox.Sandbox.create``.""" + Sandbox, _, _, _, _ = _require_opensandbox_sdk() + + kwargs: dict[str, Any] = { + "env": spec.env, + "metadata": spec.metadata, + "resource": spec.resources, + "extensions": _spec_extensions(spec), + "connection_config": self._connection_config(request_timeout_s=self._create.request_timeout_s), + } + if spec.image is not None: + kwargs["image"] = spec.image + snapshot_id = spec.provider_options.get(PROVIDER_OPTION_SNAPSHOT_ID) + if snapshot_id is not None: + kwargs["snapshot_id"] = snapshot_id + if spec.timeout_s is not None: + kwargs["timeout"] = timedelta(seconds=spec.timeout_s) + if spec.ready_timeout_s is not None: + kwargs["ready_timeout"] = timedelta(seconds=spec.ready_timeout_s) + if spec.entrypoint is not None: + kwargs["entrypoint"] = spec.entrypoint + platform = spec.provider_options.get(PROVIDER_OPTION_PLATFORM) + volumes = _spec_volumes(spec) + if platform is not None: + kwargs["platform"] = _to_platform_spec(platform) + if volumes is not None: + kwargs["volumes"] = _to_volumes(volumes) + if self._create.skip_health_check: + kwargs["skip_health_check"] = True + else: + skip_health_check = _provider_option_bool(spec.provider_options, PROVIDER_OPTION_SKIP_HEALTH_CHECK) + if skip_health_check is not None: + kwargs["skip_health_check"] = skip_health_check + + timeout_s = self._create.timeout_s + if timeout_s is None and self._connection.request_timeout_s is not None: + timeout_s = float(self._connection.request_timeout_s) + + sandbox_id: str | None = None + sandbox: Any | None = None + try: + if timeout_s is None: + sandbox = await Sandbox.create(**kwargs) + else: + sandbox = await asyncio.wait_for( + Sandbox.create(**kwargs), + timeout=timeout_s, + ) + if sandbox is None: + raise RuntimeError("OpenSandbox SDK create returned no sandbox handle") + sandbox_id = str(sandbox.id) + except TimeoutError as e: + error = OpenSandboxCreateTimeoutError( + "Timed out creating OpenSandbox sandbox after " + f"{timeout_s:g}s; image={spec.image!r}, " + f"ready_timeout_s={spec.ready_timeout_s!r}" + ) + raise error from e + if sandbox_id is None: + raise RuntimeError("OpenSandbox SDK create returned no sandbox handle") + created_handle = SandboxHandle( + sandbox_id=sandbox_id, + provider_name=self.name, + raw=sandbox, + ) + handle = created_handle + try: + if self._create.skip_health_check: + handle = await self._connect_after_create(created_handle, spec) + await self._verify_created_handle(handle) + except Exception: + await self._cleanup_failed_create_handle(created_handle) + raise + return handle + + async def _create_with_retries( + self, + spec: SandboxSpec, + ) -> SandboxHandle: + AsyncRetrying, retry_if_exception, stop_after_attempt, wait_random_exponential = _require_tenacity() + retry_policy = AsyncRetrying( + retry=retry_if_exception(_is_retryable_create_error), + stop=stop_after_attempt(self._create.retries + 1), + wait=wait_random_exponential( + multiplier=self._create.retry_delay_s, + max=self._create.retry_max_delay_s, + ), + before_sleep=_log_create_retry, + reraise=True, + ) + async for attempt in retry_policy: + with attempt: + return await self._create_once(spec) + + raise OpenSandboxCreateError("OpenSandbox create retry loop did not run") + + async def create(self, spec: SandboxSpec) -> SandboxHandle: + """Create one sandbox through the configured OpenSandbox path.""" + spec = self._with_default_image_pull_policy(_normalize_spec(spec)) + return await self._create_with_retries(spec) + + async def status(self, handle: SandboxHandle) -> SandboxStatus: + """Return the current OpenSandbox lifecycle status.""" + get_info = getattr(handle.raw, "get_info", None) + if get_info is None: + return SandboxStatus.UNKNOWN + info = await self._await_sdk_operation( + get_info, + operation="get_info", + sandbox_id=handle.sandbox_id, + timeout_s=float(self._connection.request_timeout_s) + if self._connection.request_timeout_s is not None + else None, + ) + raw_status = getattr(info, "status", None) + return _to_sandbox_status(getattr(raw_status, "state", None) if raw_status is not None else None) + + def _command_retry_count(self) -> int: + return ( + self._operations.retries if self._operations.command_retries is None else self._operations.command_retries + ) + + async def _exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + retries: int | None = None, + ) -> SandboxExecResult: + """Run a command inside an OpenSandbox sandbox.""" + _, _, RunCommandOpts, _, _ = _require_opensandbox_sdk() + + opts_kwargs: dict[str, Any] = {} + if cwd is not None: + opts_kwargs["working_directory"] = cwd + if env is not None: + opts_kwargs["envs"] = env + if timeout_s is not None: + opts_kwargs["timeout"] = timedelta(seconds=timeout_s) + + effective_command = command + if isinstance(user, int): + opts_kwargs["uid"] = user + elif isinstance(user, str) and user != "root": + effective_command = f"su -s /bin/sh -c {shlex.quote(command)} {shlex.quote(user)}" + + sdk_timeout_s = ( + float(timeout_s) + 60.0 + if timeout_s is not None + else ( + float(self._connection.request_timeout_s) if self._connection.request_timeout_s is not None else None + ) + ) + effective_retries = self._command_retry_count() if retries is None else retries + execution = await self._await_sdk_operation( + lambda: handle.raw.commands.run(effective_command, opts=RunCommandOpts(**opts_kwargs)), + operation="command run", + sandbox_id=handle.sandbox_id, + timeout_s=sdk_timeout_s, + retries=effective_retries, + ) + stdout = "\n".join(msg.text for msg in execution.logs.stdout) or None + stderr_parts = [msg.text for msg in execution.logs.stderr] + if execution.error is not None: + stderr_parts.append(f"{execution.error.name}: {execution.error.value}") + stderr = "\n".join(stderr_parts) or None + error_type = None + if execution.exit_code is not None: + return_code = execution.exit_code + elif execution.error is not None: + return_code = 125 + error_type = "sandbox" + else: + return_code = 0 + + return SandboxExecResult(stdout=stdout, stderr=stderr, return_code=return_code, error_type=error_type) + + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + """Run a command inside an OpenSandbox sandbox.""" + return await self._exec( + handle, + command, + cwd=cwd, + env=env, + timeout_s=timeout_s, + user=user, + retries=self._command_retry_count(), + ) + + async def _write_file(self, handle: SandboxHandle, target_path: str, data: str | bytes) -> None: + """Write one file into an OpenSandbox sandbox.""" + await self._await_sdk_operation( + lambda: handle.raw.files.write_file(target_path, data), + operation=f"write_file({target_path})", + sandbox_id=handle.sandbox_id, + timeout_s=float(self._connection.request_timeout_s) + if self._connection.request_timeout_s is not None + else None, + ) + + async def _read_file(self, handle: SandboxHandle, source_path: str) -> bytes: + """Read one file from an OpenSandbox sandbox.""" + return await self._await_sdk_operation( + lambda: handle.raw.files.read_bytes(source_path), + operation=f"read_file({source_path})", + sandbox_id=handle.sandbox_id, + timeout_s=float(self._connection.request_timeout_s) + if self._connection.request_timeout_s is not None + else None, + ) + + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + """Upload one local file into an OpenSandbox sandbox.""" + await self._write_file(handle, target_path, source_path.read_bytes()) + + async def download_file(self, handle: SandboxHandle, source_path: str, target_path: Path) -> None: + """Download one file from an OpenSandbox sandbox.""" + target_path.parent.mkdir(parents=True, exist_ok=True) + target_path.write_bytes(await self._read_file(handle, source_path)) + + async def close(self, handle: SandboxHandle, *, delete: bool) -> None: + """Close local SDK resources and optionally terminate the sandbox.""" + kill_error: Exception | None = None + if delete: + try: + await self._await_sdk_operation( + lambda: handle.raw.kill(), + operation="kill", + sandbox_id=handle.sandbox_id, + timeout_s=self._operations.close_timeout_s, + ) + except Exception as e: + if not _is_missing_sandbox_delete_error(e): + kill_error = e + else: + LOGGER.info( + "OpenSandbox sandbox %r was already deleted during close", + handle.sandbox_id, + ) + + close_error: Exception | None = None + try: + await self._await_sdk_call( + handle.raw.close(), + operation="close", + sandbox_id=handle.sandbox_id, + timeout_s=self._operations.close_timeout_s, + ) + except Exception as e: + close_error = e + LOGGER.warning( + "Timed out or failed while closing local OpenSandbox SDK handle for sandbox %r: %r", + handle.sandbox_id, + e, + ) + + if kill_error is not None: + if close_error is not None: + raise RuntimeError( + "Failed to delete and close OpenSandbox sandbox " + f"{handle.sandbox_id!r}: delete_error={kill_error!r}, " + f"close_error={close_error!r}" + ) from kill_error + raise kill_error + if close_error is not None: + if delete: + return + raise close_error diff --git a/nemo_gym/sandbox/providers/registry.py b/nemo_gym/sandbox/providers/registry.py new file mode 100644 index 0000000000..bd98c1b885 --- /dev/null +++ b/nemo_gym/sandbox/providers/registry.py @@ -0,0 +1,85 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Provider registration utilities.""" + +from collections.abc import Callable, Mapping +from typing import Any, TypeAlias + +from nemo_gym.sandbox.providers.base import SandboxProvider + + +ProviderClass: TypeAlias = type[SandboxProvider] +ProviderLoader: TypeAlias = Callable[[], ProviderClass] + +_PROVIDER_REGISTRY: dict[str, ProviderClass] = {} +_BUILTIN_PROVIDER_LOADERS: dict[str, ProviderLoader] = {} + + +def register_provider(name: str, provider_class: ProviderClass, *, override: bool = False) -> None: + """Register a sandbox provider class.""" + if not name: + raise ValueError("Provider name must be non-empty") + if not override and (name in _PROVIDER_REGISTRY or name in _BUILTIN_PROVIDER_LOADERS): + raise ValueError(f"Sandbox provider {name!r} is already registered") + _PROVIDER_REGISTRY[name] = provider_class + + +def get_provider_class(name: str) -> ProviderClass: + """Return a registered provider class.""" + try: + return _PROVIDER_REGISTRY[name] + except KeyError as e: + loader = _BUILTIN_PROVIDER_LOADERS.get(name) + if loader is not None: + return loader() + available = ", ".join(list_providers()) or "" + raise ValueError(f"Unknown sandbox provider {name!r}. Available providers: {available}") from e + + +def create_provider(config: Mapping[str, Any]) -> SandboxProvider: + """Instantiate a provider from a single-key provider config.""" + if len(config) != 1: + raise ValueError("Sandbox provider config must contain exactly one provider name") + provider_name, provider_kwargs = next(iter(config.items())) + if not isinstance(provider_name, str) or not provider_name: + raise ValueError("Sandbox provider name must be a non-empty string") + if provider_kwargs is None: + provider_kwargs = {} + if not isinstance(provider_kwargs, Mapping): + raise TypeError(f"Sandbox provider {provider_name!r} config must be a mapping") + + provider_class = get_provider_class(provider_name) + return provider_class(**dict(provider_kwargs)) + + +def list_providers() -> list[str]: + """List registered provider names.""" + return sorted({*_PROVIDER_REGISTRY, *_BUILTIN_PROVIDER_LOADERS}) + + +def _load_opensandbox_provider() -> ProviderClass: + from nemo_gym.sandbox.providers.opensandbox import OpenSandboxProvider + + return OpenSandboxProvider + + +def _load_ecs_fargate_provider() -> ProviderClass: + from nemo_gym.sandbox.providers.ecs_fargate import EcsFargateProvider + + return EcsFargateProvider + + +_BUILTIN_PROVIDER_LOADERS["opensandbox"] = _load_opensandbox_provider +_BUILTIN_PROVIDER_LOADERS["ecs_fargate"] = _load_ecs_fargate_provider diff --git a/nemo_gym/sandbox/utils.py b/nemo_gym/sandbox/utils.py new file mode 100644 index 0000000000..b25f0f6962 --- /dev/null +++ b/nemo_gym/sandbox/utils.py @@ -0,0 +1,27 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Sandbox utility helpers.""" + + +def rewrite_image(image: str | None, rewrites: list[dict[str, str]]) -> str | None: + """Apply ordered image-prefix rewrites used by sandbox configs.""" + if image is None: + return None + for rewrite in rewrites: + from_prefix = rewrite["from"] + to_prefix = rewrite["to"] + if image.startswith(from_prefix): + return to_prefix + image[len(from_prefix) :] + return image diff --git a/pyproject.toml b/pyproject.toml index 45298ffb4f..5b6572a811 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -220,6 +220,26 @@ docs = [ ] [project.optional-dependencies] +sandbox = [ + # Tenacity: Retry helpers used by sandbox providers. + # Updated: Sat May 09, 2026 with tenacity==9.1.4 + # License: Apache 2.0 https://github.com/jd/tenacity/blob/master/LICENSE + "tenacity>=9.1.4", + + # OpenSandbox SDK: used by the OpenSandbox sandbox provider for create/exec/delete and SDK pool creation. + # Updated: Sat May 16, 2026 with opensandbox>=0.1.9 + # License: Apache 2.0 + "opensandbox>=0.1.9", +] + +sandbox-ecs = [ + # boto3: AWS SDK used by the ECS Fargate sandbox provider for ECS/EC2/ECR/ + # CodeBuild/S3/SSM/Secrets Manager calls. + # Updated: Tue Jun 03, 2026 with boto3>=1.34 + # License: Apache 2.0 https://github.com/boto/boto3/blob/develop/LICENSE + "boto3>=1.34", +] + # We include dev dependencies as an extra since technically each server module is a consumer (which means we cannot use dependency groups, which are intended to be within a project). dev = [ # Pre-commit: Used for pre-commit hooks. @@ -395,7 +415,18 @@ ng_reinstall = "nemo_gym.cli:reinstall" [tool.setuptools.packages.find] where = ["."] -include = ["benchmarks", "resources_servers", "responses_api_agents", "responses_api_models", "nemo_gym"] +include = [ + "benchmarks", + "benchmarks.*", + "resources_servers", + "resources_servers.*", + "responses_api_agents", + "responses_api_agents.*", + "responses_api_models", + "responses_api_models.*", + "nemo_gym", + "nemo_gym.*", +] ################################################ # Testing diff --git a/responses_api_agents/mini_swe_agent_2/.gitignore b/responses_api_agents/mini_swe_agent_2/.gitignore new file mode 100644 index 0000000000..68bcbc9609 --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/.gitignore @@ -0,0 +1 @@ +results/ \ No newline at end of file diff --git a/responses_api_agents/mini_swe_agent_2/README.md b/responses_api_agents/mini_swe_agent_2/README.md new file mode 100644 index 0000000000..c021cedb13 --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/README.md @@ -0,0 +1,397 @@ +# Mini-SWE-Agent 2 Sandbox Agent + +A NeMo Gym Responses API agent that integrates +[mini-swe-agent](https://github.com/SWE-agent/mini-swe-agent) v2 for evaluating +language models on SWE-bench style software engineering tasks through the public +`nemo_gym.sandbox` API. + +This agent intentionally keeps only the sandbox-backed path. It does not carry +over the older Docker/Singularity mini-SWE integration. + +## Contents + +- [Mini-SWE-Agent 2 Sandbox Agent](#mini-swe-agent-2-sandbox-agent) + - [Contents](#contents) + - [Overview](#overview) + - [Dataset Information](#dataset-information) + - [Configuration](#configuration) + - [Agent Configuration](#agent-configuration) + - [Model Parameters](#model-parameters) + - [Usage](#usage) + - [Server](#server) + - [Collect Rollouts](#collect-rollouts) + - [Running SWE-bench on ECS Fargate](#running-swe-bench-on-ecs-fargate) + - [Sandbox Environment Adapter](#sandbox-environment-adapter) + - [Environment Lifecycle](#environment-lifecycle) + - [Contributing](#contributing) + - [Licensing Information](#licensing-information) + - [Dependencies](#dependencies) + +## Overview + +`mini_swe_agent_2` runs mini-swe-agent's synchronous SWE-bench harness while +creating and executing each task environment through Gym's provider-neutral +sandbox facade. The validated path in this directory is: + +- mini-swe-agent `2.1.0` +- SWE-bench task rows, including SWE-bench Verified +- `env: sandbox` +- `responses_api_agents.mini_swe_agent_2.sandbox_environment.MiniSWESandboxEnvironment` +- OpenSandbox through `nemo_gym.sandbox.providers.opensandbox` + +For each `/run` request, `MiniSWEAgent.run()` loads mini-swe-agent's built-in +`swebench.yaml`, injects sandbox settings, runs mini-swe-agent in a Ray remote +task, evaluates the generated patch with the SWE-bench harness, and returns a +Gym verify response with reward `1.0` only when the instance is resolved and the +evaluation report includes test status. + +`MiniSWEAgent.setup_webserver()` also registers `/v1/responses`, but +`MiniSWEAgent.responses()` is intentionally not implemented in this agent. The +supported eval path is `/run`, typically via `ng_collect_rollouts`. + +## Dataset Information + +- Eval data - [princeton-nlp/SWE-bench_Verified](https://huggingface.co/datasets/princeton-nlp/SWE-bench_Verified) + is the primary validation target. It contains 500 human-validated SWE-bench + test instances. +- The rollout input JSONL should preserve the SWE-bench instance fields needed + by `swebench`, such as `instance_id`, `repo`, `base_commit`, + `problem_statement`, `patch`, `test_patch`, `FAIL_TO_PASS`, `PASS_TO_PASS`, + and related version fields. +- Each row must also include `responses_create_params`. Extra top-level + SWE-bench fields are accepted by the agent request model and passed into + mini-swe-agent as the instance dictionary. + +Example row shape: + +```json +{ + "instance_id": "django__django-13410", + "repo": "django/django", + "base_commit": "...", + "problem_statement": "...", + "patch": "...", + "test_patch": "...", + "FAIL_TO_PASS": ["..."], + "PASS_TO_PASS": ["..."], + "responses_create_params": { + "input": [], + "temperature": 0.6, + "top_p": 1.0, + "max_output_tokens": 16384 + } +} +``` + +When `image_name` is present on a row, the agent uses it directly. Otherwise it +derives the SWE-bench image from `instance_id` and `subset`: + +- `subset: verified` uses `docker.io/swebench/sweb.eval.x86_64.:latest` + with `__` replaced by `_1776_`. +- Other subsets use `docker.io/xingyaoww/sweb.eval.x86_64.:latest` with + `__` replaced by `_s_`. + +The default OpenSandbox config uses explicit Docker Hub image refs so cluster +mirroring can happen in the container runtime instead of Gym-side image +rewrites. + +## Configuration + +### Agent Configuration + +Path - `responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_opensandbox.yaml` + +```yaml +mini_swe_agent_2: + responses_api_agents: + mini_swe_agent_2: + entrypoint: app.py + domain: coding + description: Software engineering tasks driven by mini-swe-agent harness on OpenSandbox. + value: Improve agentic software engineering capabilities. + model_server: + type: responses_api_models + name: policy_model + concurrency: 64 + env: sandbox + sandbox_provider: + opensandbox: + connection: + domain: opensandbox-server.opensandbox-system.svc.cluster.local + api_key: ${oc.env:OPENSANDBOX_API_KEY} + protocol: http + request_timeout_s: 300 + use_server_proxy: true + create: + request_timeout_s: 1200 + timeout_s: 1200 + skip_health_check: true + retries: 10 + retry_delay_s: 5.0 + retry_max_delay_s: 90.0 + probe: + timeout_s: 60 + deadline_s: 180 + stable_count: 2 + stable_delay_s: 1.0 + operations: + retries: 5 + retry_delay_s: 1.0 + retry_max_delay_s: 45.0 + command_retries: 3 + close_timeout_s: 30 + sandbox_spec: + timeout_s: 18000 + ready_timeout_s: 1200 + resources: + cpu: "2" + memory: 8Gi + ephemeral-storage: 20Gi + provider_options: + platform: + os: linux + arch: amd64 + metadata: + benchmark: swebench-verified + harness: mini-swe-agent + sandbox-api: opensandbox-sdk + sandbox_environment_kwargs: + cwd: /testbed + conda_env: testbed + activate_conda: true + user: root + delete: true + run_golden: false + step_timeout: 600 + eval_timeout: 1800 + skip_if_exists: false + step_limit: 250 +``` + +Optional `sandbox_resource_profiles` can be configured as a list of resource +maps. When present, the agent hashes `instance_id` and deterministically merges +one profile into `sandbox_spec.resources`. This is useful for spreading +SWE-bench tasks across a small set of resource sizes without changing the input +data. + +### Model Parameters + +`MiniSWEAgent.run()` maps supported Responses API fields into mini-swe-agent +chat-completions kwargs: + +- `temperature`, `top_p`, `top_logprobs`, and `parallel_tool_calls` pass through. +- `max_output_tokens` becomes `max_tokens`. +- `responses_create_params.metadata.extra_body` must be a JSON object and is + passed as `extra_body`. +- `responses_create_params.metadata.chat_template_kwargs` must be a JSON object + and is nested under `extra_body.chat_template_kwargs`. +- `tool_choice` comes from the agent config when set, otherwise from the request. + The special value `bash` expands to the OpenAI function choice for the `bash` + tool. + +Keep the requested generation budget compatible with the live vLLM deployment. +For example, a deployment served with `--max-model-len 32768` will reject +`max_output_tokens=49152`. In earlier smoke testing, that upstream vLLM rejection +surfaced in mini-swe-agent as repeated: + +```text +No tool calls found in the response. Every response MUST include at least one tool call. +``` + +That symptom was not a sandbox failure and was not a reason to force the `bash` +tool. The successful smoke kept `tool_choice=auto` and lowered +`max_output_tokens` to `16384`. + +## Usage + +### Server + +Set the policy model endpoint in `env.yaml` or with equivalent Hydra overrides: + +```yaml +policy_base_url: http://..svc.cluster.local:8000/v1 +policy_api_key: dummy-key +policy_model_name: +``` + +Start the mini-swe-agent 2 server with the OpenSandbox provider and a policy +model server. The values below show a representative SWE-bench eval setup: + +```bash +CONFIG_PATHS="responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_opensandbox.yaml,responses_api_models/vllm_model/configs/vllm_model.yaml" + +ng_run "+config_paths=[$CONFIG_PATHS]" \ + +mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.concurrency=64 \ + +mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.step_timeout=600 \ + +mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.eval_timeout=1800 \ + +mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.step_limit=50 \ + +mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.run_golden=false \ + '+mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.sandbox_spec.resources={cpu: 500m, memory: 4Gi, ephemeral-storage: 8Gi}' \ + '+mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.sandbox_spec.metadata={benchmark: swebench-verified, harness: mini_swe_agent_2, endpoint_label: hosted-vllm, run_family: mini-swe-agent-2-pass8}' +``` + +Use a model server config that matches the policy endpoint you are serving. The +example above uses `vllm_model`, which is the common path for hosted vLLM +`/v1/chat/completions` endpoints. + +### Collect Rollouts + +Collect eval rollouts from a SWE-bench-style JSONL file: + +```bash +ng_collect_rollouts \ + +agent_name=mini_swe_agent_2 \ + +input_jsonl_fpath=data/mini_swe_verified_smoke8.jsonl \ + +output_jsonl_fpath=results/mini_swe_agent_2_pass8.jsonl \ + +limit=8 \ + +num_repeats=8 \ + +num_samples_in_parallel=64 \ + '+responses_create_params={max_output_tokens: 32768, temperature: 0.6, top_p: 0.95, metadata: {chat_template_kwargs: "{\"enable_thinking\": true}"}}' +``` + +`ng_collect_rollouts` also writes +`results/mini_swe_agent_2_pass8_aggregate_metrics.json` +with per-task eval status, pass@k, resolved task counts, and eval error rates. +After collecting repeated rollouts, run `ng_reward_profile` on the collected +output when you want the standalone profiler JSONL as well: + +```bash +ng_reward_profile \ + +input_jsonl_fpath=data/mini_swe_verified_smoke8.jsonl \ + +materialized_inputs_jsonl_fpath=results/mini_swe_agent_2_pass8_materialized_inputs.jsonl \ + +rollouts_jsonl_fpath=results/mini_swe_agent_2_pass8.jsonl \ + +pass_threshold=1.0 +``` + +The profiler writes `*_reward_profiling.jsonl` and `*_agent_metrics.json` +next to the rollouts file. + +The agent writes per-instance mini-swe-agent configs and result artifacts under +`results///`. + +Use the agent's `step_timeout` and `eval_timeout` overrides above to bound tool +and verifier execution. If you launch from a custom Kubernetes wrapper, add any +outer per-sample guard there. + +## Running SWE-bench on ECS Fargate + +Swap OpenSandbox for ECS Fargate by using +`configs/mini_swe_agent_ecs_fargate.yaml` — the agent loop and SWE-bench +verifier are unchanged. For provider setup (AWS infra/SSM, credentials, the +`:52222` network requirement, and automatic image mirroring via `auto_mirror`) +see the [ECS Fargate provider README](../../nemo_gym/sandbox/providers/ecs_fargate/README.md). +SWE-bench task images are pulled into the ECR mirror on first use, so no manual +image staging is needed. + +Each input row needs the SWE-bench instance fields shown under +[Dataset Information](#dataset-information) plus `subset` (`verified`), `split` +(`test`), and `responses_create_params`. + +**Golden smoke (no model)** — `run_golden=true` applies the gold patch and runs +the verifier in-sandbox, so the model is never called: + +```bash +CONFIG_PATHS="responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_ecs_fargate.yaml,responses_api_models/vllm_model/configs/vllm_model.yaml" + +AWS_REGION=us-east-1 ng_run "+config_paths=[$CONFIG_PATHS]" \ + ++mini_swe_agent_2.responses_api_agents.mini_swe_agent_2.run_golden=true + +ng_collect_rollouts +agent_name=mini_swe_agent_2 \ + +input_jsonl_fpath=data/swe_verified_smoke.jsonl \ + +output_jsonl_fpath=results/ecs_golden.jsonl \ + +limit=1 +num_repeats=1 +num_samples_in_parallel=1 +``` + +A resolved instance returns reward `1.0` with `tests_status` populated. + +**Real rollout** — drop `run_golden` and point the model server at a live +OpenAI-compatible endpoint (`policy_model_name` is the model id sent upstream): + +```bash +AWS_REGION=us-east-1 ng_run "+config_paths=[$CONFIG_PATHS]" \ + ++policy_base_url=https:///v1 ++policy_api_key= ++policy_model_name= + +ng_collect_rollouts +agent_name=mini_swe_agent_2 \ + +input_jsonl_fpath=data/swe_verified_smoke.jsonl \ + +output_jsonl_fpath=results/ecs_rollout.jsonl \ + +limit=8 +num_repeats=1 +num_samples_in_parallel=8 \ + '+responses_create_params={max_output_tokens: 16384, temperature: 0.6, top_p: 0.95}' +``` + +Reasoning models often only accept the default `temperature` (`1`) and reject a +custom `top_p`; in that case use +`'+responses_create_params={temperature: 1, max_output_tokens: 16384}'` and keep +`max_output_tokens` large enough for reasoning tokens. + +## Sandbox Environment Adapter + +`MiniSWESandboxEnvironment` adapts mini-swe-agent's synchronous environment +contract to `nemo_gym.sandbox.Sandbox`. + +When `env` is `sandbox`, Gym injects this environment config before calling +mini-swe-agent: + +```yaml +environment: + environment_class: responses_api_agents.mini_swe_agent_2.sandbox_environment.MiniSWESandboxEnvironment + image: + provider: + opensandbox: + connection: ... + spec: + resources: ... + provider_options: + platform: ... + metadata: ... +``` + +### Environment Lifecycle + +`MiniSWESandboxEnvironment.__init__()`: + +- Validates that a sandbox provider was configured. +- Builds a `SandboxSpec` from the task image, environment variables, metadata, + resources, and provider-specific options. +- Adds standard metadata such as `nemo_gym_agent=mini_swe_agent_2` and + `instance_id`. +- Creates a `Sandbox` facade and calls `Sandbox.start(...)`. + +`execute()`: + +- Receives mini-swe-agent's command action. +- Applies the configured working directory and timeout. +- Optionally wraps the command in `conda activate ` for SWE-bench images + that expect a prebuilt conda environment. +- Calls `Sandbox.exec(...)` as the configured user, root by default. +- Returns mini-swe-agent's expected sync response shape: + +```python +{ + "output": "...", + "returncode": 0, + "exception_info": "", +} +``` + +`_check_finished()` preserves mini-swe-agent's submit sentinel behavior. If the +command output begins with `COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT` and the +command succeeded, it raises `minisweagent.exceptions.Submitted` with the final +submission payload. + +`cleanup()` calls `Sandbox.stop(...)` to release provider-owned resources and +stop the sync facade's private loop. + +## Contributing + +Please refer to the main NeMo Gym documentation for contributing guidelines. + +## Licensing Information + +- **Code**: Apache 2.0 +- **SWE-bench Verified**: MIT + +### Dependencies + +- **nemo_gym**: Apache 2.0 +- **mini-swe-agent**: MIT +- **SWE-bench / swebench**: MIT diff --git a/responses_api_agents/mini_swe_agent_2/__init__.py b/responses_api_agents/mini_swe_agent_2/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/responses_api_agents/mini_swe_agent_2/app.py b/responses_api_agents/mini_swe_agent_2/app.py new file mode 100644 index 0000000000..2de794e1ec --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/app.py @@ -0,0 +1,790 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import asyncio +import hashlib +import json +import sys +import time +import traceback +from asyncio import Semaphore +from pathlib import Path +from typing import Any, Callable, Literal, Optional, cast +from uuid import uuid4 + +import ray +import yaml +from fastapi import Body, FastAPI +from minisweagent.config import builtin_config_dir, get_config_path +from pydantic import ConfigDict + +from nemo_gym.base_resources_server import ( + BaseRunRequest, + BaseVerifyRequest, + BaseVerifyResponse, +) +from nemo_gym.base_responses_api_agent import ( + BaseResponsesAPIAgentConfig, + SimpleResponsesAPIAgent, +) +from nemo_gym.config_types import ModelServerRef +from nemo_gym.global_config import TASK_INDEX_KEY_NAME +from nemo_gym.openai_utils import ( + NeMoGymResponse, + NeMoGymResponseCreateParamsNonStreaming, +) +from nemo_gym.reward_profile import compute_pass_majority_metrics, highest_k_metrics +from nemo_gym.server_utils import ( + ServerClient, + get_first_server_config_dict, +) + + +class MiniSWEAgentConfig(BaseResponsesAPIAgentConfig): + model_server: ModelServerRef + env: Literal["sandbox"] + concurrency: int + sandbox_provider: Optional[dict[str, Any]] = None + sandbox_spec: Optional[dict[str, Any]] = None + sandbox_environment_kwargs: Optional[dict[str, Any]] = None + run_golden: bool = False + step_timeout: int = 600 + eval_timeout: int = 1800 + skip_if_exists: bool = False + step_limit: int = 250 + tool_choice: Optional[str | dict[str, Any]] = None + sandbox_resource_profiles: Optional[list[dict[str, str]]] = None + + +class MiniSWEAgentRunRequest(BaseRunRequest): + model_config = ConfigDict(extra="allow") + + +class MiniSWEAgentVerifyRequest(BaseVerifyRequest): + model_config = ConfigDict(extra="allow") + + +class MiniSWEAgentVerifyResponse(BaseVerifyResponse): + model_config = ConfigDict(extra="allow") + + +@ray.remote( + scheduling_strategy="SPREAD", + runtime_env={ + "py_executable": sys.executable, + }, +) +def runner_ray_remote(runner: Callable, params: dict[str, Any]) -> Any: + return runner(**params) + + +def _json_dict_from_metadata(value: Any, *, field_name: str) -> dict[str, Any]: + if value is None: + return {} + if isinstance(value, dict): + return value + if isinstance(value, str): + parsed = json.loads(value) + if isinstance(parsed, dict): + return parsed + raise ValueError(f"responses_create_params.metadata.{field_name} must be a JSON object") + + +def _responses_create_params_to_model_kwargs( + params: dict[str, Any], + *, + default_tool_choice: Any = None, +) -> dict[str, Any]: + """Convert Gym Responses API rollout params into mini-swe-agent chat-completions kwargs.""" + model_kwargs: dict[str, Any] = {} + for key in ("temperature", "top_p", "top_logprobs", "parallel_tool_calls"): + value = params.get(key) + if value is not None: + model_kwargs[key] = value + + max_output_tokens = params.get("max_output_tokens") + if max_output_tokens is not None: + model_kwargs["max_tokens"] = max_output_tokens + + metadata = params.get("metadata") or {} + extra_body = _json_dict_from_metadata(metadata.get("extra_body"), field_name="extra_body") + chat_template_kwargs = _json_dict_from_metadata( + metadata.get("chat_template_kwargs"), + field_name="chat_template_kwargs", + ) + if chat_template_kwargs: + extra_body["chat_template_kwargs"] = chat_template_kwargs + if extra_body: + model_kwargs["extra_body"] = extra_body + + tool_choice = default_tool_choice if default_tool_choice is not None else params.get("tool_choice") + if tool_choice == "bash": + model_kwargs["tool_choice"] = _bash_tool_choice() + elif tool_choice is not None: + model_kwargs["tool_choice"] = tool_choice + + return model_kwargs + + +def _bash_tool_choice() -> dict[str, Any]: + return {"type": "function", "function": {"name": "bash"}} + + +def _sandbox_spec_for_instance( + spec: dict[str, Any] | None, + *, + resource_profiles: list[dict[str, str]] | None, + instance_id: str, +) -> dict[str, Any]: + instance_spec = dict(spec or {}) + if not resource_profiles: + return instance_spec + + resources = dict(instance_spec.get("resources") or {}) + digest = hashlib.sha256(instance_id.encode("utf-8")).digest() + profile = resource_profiles[int.from_bytes(digest[:4], "big") % len(resource_profiles)] + resources.update(profile) + instance_spec["resources"] = resources + return instance_spec + + +def _swebench_config_path() -> Path: + for candidate in ( + builtin_config_dir / "extra" / "swebench.yaml", + builtin_config_dir / "benchmarks" / "swebench.yaml", + ): + if candidate.exists(): + return candidate + return builtin_config_dir / "extra" / "swebench.yaml" + + +def _swebench_image_name(instance: dict[str, Any], subset: str) -> str: + image_name = instance.get("image_name") + if image_name: + return str(image_name) + + instance_id = instance["instance_id"] + if subset == "verified": + docker_compatible_id = instance_id.replace("__", "_1776_") + return f"docker.io/swebench/sweb.eval.x86_64.{docker_compatible_id}:latest".lower() + + docker_compatible_id = instance_id.replace("__", "_s_") + return f"docker.io/xingyaoww/sweb.eval.x86_64.{docker_compatible_id}:latest".lower() + + +def _message_content_to_text(content: Any) -> str: + if isinstance(content, str): + return content + if isinstance(content, list): + parts = [] + for item in content: + if isinstance(item, dict): + parts.append(str(item.get("text") or item.get("content") or "")) + else: + parts.append(str(item)) + return "\n".join(part for part in parts if part) + return "" if content is None else str(content) + + +def _strip_extra(item: Any) -> dict[str, Any]: + if hasattr(item, "model_dump"): + item = item.model_dump() + if not isinstance(item, dict): + return {"type": "message", "role": "user", "content": str(item)} + return {key: value for key, value in item.items() if key != "extra"} + + +def _split_trajectory_for_responses( + messages: list[dict[str, Any]], +) -> tuple[list[dict[str, Any]], list[dict[str, Any]], list[dict[str, Any]]]: + input_messages: list[dict[str, Any]] = [] + output_items: list[dict[str, Any]] = [] + raw_responses: list[dict[str, Any]] = [] + in_initial_prompt = True + + for message in messages: + role = message.get("role") + if in_initial_prompt and role in {"system", "user"}: + input_messages.append( + {"type": "message", "role": role, "content": _message_content_to_text(message.get("content"))} + ) + continue + + in_initial_prompt = False + if message.get("object") == "response": + response = _strip_extra(message) + raw_responses.append(response) + output_items.extend(_strip_extra(item) for item in response.get("output", [])) + elif role == "assistant": + content = _message_content_to_text(message.get("content")) + if content: + output_items.append( + { + "id": message.get("id") or f"msg_{uuid4()}", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": content, "annotations": []}], + } + ) + for tool_call in message.get("tool_calls") or []: + function = tool_call.get("function") or {} + output_items.append( + { + "id": tool_call.get("id") or f"fc_{uuid4()}", + "type": "function_call", + "name": function.get("name") or tool_call.get("name") or "", + "call_id": tool_call.get("id") or tool_call.get("call_id") or "", + "arguments": function.get("arguments") or tool_call.get("arguments") or "{}", + } + ) + elif role == "tool": + output_items.append( + { + "type": "function_call_output", + "call_id": message.get("tool_call_id") or message.get("call_id") or "", + "output": _message_content_to_text(message.get("content")), + } + ) + elif message.get("type") == "function_call_output": + output_items.append(_strip_extra(message)) + + return input_messages, output_items, raw_responses + + +def _default_response_object() -> dict[str, Any]: + return { + "id": f"resp_{str(uuid4())}", + "created_at": int(time.time()), + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": {}, + "object": "response", + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + "background": False, + "max_output_tokens": None, + "max_tool_calls": None, + "previous_response_id": None, + "prompt": None, + "reasoning": { + "effort": None, + "generate_summary": None, + "summary": None, + }, + "service_tier": "default", + "status": "completed", + "text": {"format": {"type": "text"}, "verbosity": "medium"}, + "top_logprobs": 0, + "truncation": "disabled", + "usage": { + "input_tokens": 0, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens": 0, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": 0, + }, + "user": None, + "prompt_cache_key": None, + "safety_identifier": None, + "store": True, + } + + +def _is_resolved(instance_id: str, eval_report: dict[str, Any]) -> bool: + try: + if not eval_report: + return False + report = eval_report["eval_report"][instance_id] + resolved = bool(report["resolved"]) + if not report.get("tests_status"): + return False + + tests_status = report["tests_status"] + f2f = tests_status.get("FAIL_TO_PASS", {}) + p2p = tests_status.get("PASS_TO_PASS", {}) + total_reported = ( + len(f2f.get("success", [])) + + len(f2f.get("failure", [])) + + len(p2p.get("success", [])) + + len(p2p.get("failure", [])) + ) + return resolved and total_reported > 0 + except Exception as exc: + print(f"Error in _is_resolved: {exc}", flush=True) + return False + + +def _metadata_dict(verify_response: dict[str, Any]) -> dict[str, Any]: + metadata = verify_response.get("metadata") or {} + return metadata if isinstance(metadata, dict) else {} + + +def _eval_report_map(verify_response: dict[str, Any]) -> dict[str, Any]: + report = _metadata_dict(verify_response).get("eval_report") or {} + return report if isinstance(report, dict) else {} + + +def _eval_instance_report(verify_response: dict[str, Any]) -> dict[str, Any]: + report_map = _eval_report_map(verify_response) + instance_id = verify_response.get("instance_id") or _metadata_dict(verify_response).get("instance_id") + if instance_id is not None: + report = report_map.get(str(instance_id)) + if isinstance(report, dict): + return report + + for report in report_map.values(): + if isinstance(report, dict) and "resolved" in report: + return report + return {} + + +def _test_status_counts(verify_response: dict[str, Any]) -> dict[str, int]: + report = _eval_instance_report(verify_response) + tests_status = report.get("tests_status") if isinstance(report, dict) else None + if not isinstance(tests_status, dict): + return {} + + counts: dict[str, int] = {} + for suite_name, suite_report in tests_status.items(): + if not isinstance(suite_report, dict): + continue + prefix = str(suite_name).lower() + counts[f"{prefix}_success"] = len(suite_report.get("success") or []) + counts[f"{prefix}_failure"] = len(suite_report.get("failure") or []) + return counts + + +def _run_eval_v2( + *, + instance: dict[str, Any], + env: Any, + model_patch: str, + instance_dir: Path, + run_id: str, + is_golden: bool, +) -> dict[str, Any]: + from swebench.harness.constants import SWEbenchInstance + from swebench.harness.docker_build import setup_logger + from swebench.harness.grading import get_eval_report + from swebench.harness.test_spec.test_spec import make_test_spec + + swebench_instance = cast(SWEbenchInstance, instance) + test_spec = make_test_spec(swebench_instance) + pred = {"instance_id": test_spec.instance_id, "model_patch": model_patch} + + instance_dir.mkdir(parents=True, exist_ok=True) + log_file = instance_dir / f"run_instance_{run_id}.log" + report_path = instance_dir / f"report_{run_id}.json" + patch_file = instance_dir / f"patch_{run_id}.diff" + patch_file.write_text(model_patch) + + logger = setup_logger(test_spec.instance_id, log_file) + logger.info(f"DEBUG test_spec {test_spec}") + logger.info(f"DEBUG eval_script {test_spec.eval_script}") + + if is_golden: + env.execute(f"cat > patch.diff <<'EOF'\n{model_patch}\n\nEOF") + env.execute("git status --porcelain") + env.execute("git apply --check patch.diff") + env.execute("git apply patch.diff") + + eval_script = test_spec.eval_script.replace("#!/bin/bash", "") + result = env.execute(eval_script, is_eval=True) + test_output = result["output"] + returncode = result["returncode"] + print(f"[EVAL]{test_spec.instance_id} returncode: {returncode}", flush=True) + + test_output_path = instance_dir / f"test_output_{run_id}.txt" + test_output_path.write_text(test_output) + print(f"[EVAL]{test_spec.instance_id} Test output written to {test_output_path}", flush=True) + + report = get_eval_report( + test_spec=test_spec, + prediction=pred, + test_log_path=str(test_output_path), + include_tests_status=True, + ) + print(f"[EVAL]{test_spec.instance_id} Result: resolved: {report[test_spec.instance_id]['resolved']}", flush=True) + + report_path.write_text(json.dumps(report, indent=4)) + return { + "instance_id": test_spec.instance_id, + "model_patch": model_patch, + "eval_report": report, + } + + +def _run_mini_swe_v2(**params: Any) -> dict[str, Any]: + from minisweagent.agents.default import DefaultAgent + from minisweagent.environments import get_environment + from minisweagent.models import get_model + + instance = params.get("instance_dict") + if isinstance(instance, str): + instance = json.loads(instance) + if not isinstance(instance, dict): + raise ValueError("mini-swe-agent v2 path requires instance_dict") + + instance = dict(instance) + instance_id = str(params.get("instance_id") or instance["instance_id"]).lower() + instance["instance_id"] = instance_id + + output_dir = Path(params["output"]) + instance_dir = output_dir / instance_id + output_dir.mkdir(parents=True, exist_ok=True) + instance_dir.mkdir(parents=True, exist_ok=True) + + config = yaml.safe_load(get_config_path(params["config"]).read_text()) + model_config = config.setdefault("model", {}) + model_config["model_class"] = "litellm" + model_config["model_name"] = params["model"] + model_config.setdefault("cost_tracking", "ignore_errors") + model_kwargs = model_config.setdefault("model_kwargs", {}) + model_kwargs["api_key"] = params["api_key"] + model_kwargs["base_url"] = params["base_url"] + model_kwargs.pop("api_base", None) + max_output_tokens = model_kwargs.pop("max_output_tokens", None) + if max_output_tokens is not None and "max_tokens" not in model_kwargs: + model_kwargs["max_tokens"] = max_output_tokens + + environment_config = config.setdefault("environment", {}) + environment_config["image"] = _swebench_image_name(instance, params["subset"]) + environment_config["step_timeout"] = params["step_timeout"] + environment_config["eval_timeout"] = params["eval_timeout"] + environment_config["instance_id"] = instance_id + environment_config["environment_class"] = ( + "responses_api_agents.mini_swe_agent_2.sandbox_environment.MiniSWESandboxEnvironment" + ) + + agent_config = config.get("agent", {}) + agent_config["step_limit"] = params["step_limit"] + agent_config.pop("collapse_limit", None) + + run_id = f"{int(time.time())}_{uuid4()}" + trajectory_path = instance_dir / f"{instance_id}_{run_id}.traj.json" + agent_config["output_path"] = trajectory_path + env = None + agent = None + try: + print(f"[EVAL]{instance_id} Creating environment...", flush=True) + env = get_environment(environment_config) + print(f"[EVAL]{instance_id} Environment created", flush=True) + + model = get_model(config=model_config) + agent = DefaultAgent(model, env, **agent_config) + + if params["run_golden"]: + exit_status = "Gold Patch Applied" + model_patch = instance.get("patch", "") + data = agent.save(None, {"messages": []}) + else: + print(f"[EVAL]{instance_id} Running mini-swe-agent v2...", flush=True) + info = agent.run(instance["problem_statement"]) + exit_status = info.get("exit_status", "") + model_patch = info.get("submission", "") + data = agent.save( + trajectory_path, + {"instance_id": instance_id}, + ) + + print(f"[EVAL]{instance_id} Running eval", flush=True) + eval_report = _run_eval_v2( + instance=instance, + env=env, + model_patch=model_patch, + instance_dir=instance_dir, + run_id=run_id, + is_golden=params["run_golden"], + ) + print(f"[EVAL]{instance_id} Eval completed", flush=True) + + input_messages, response_output, responses = _split_trajectory_for_responses(data.get("messages", [])) + + return { + instance_id: { + "input_messages": input_messages, + "response_output": response_output, + "responses": responses, + "eval_report": eval_report, + "exit_status": exit_status, + } + } + finally: + if env and hasattr(env, "cleanup"): + env.cleanup() + + +def run_mini_swe_with_sandbox(**params: Any) -> Any: + return _run_mini_swe_v2(**params) + + +class MiniSWEAgent(SimpleResponsesAPIAgent): + config: MiniSWEAgentConfig + sem: Semaphore = None + model_config = ConfigDict(arbitrary_types_allowed=True) + + def model_post_init(self, __context: Any) -> None: + self.sem = Semaphore(self.config.concurrency) + + def setup_webserver(self) -> FastAPI: + app = FastAPI() + self.setup_session_middleware(app) + app.post("/v1/responses")(self.responses) + app.post("/run")(self.run) + app.post("/aggregate_metrics")(self.aggregate_metrics) + return app + + def compute_metrics(self, tasks: list[list[dict[str, Any]]]) -> dict[str, Any]: + metrics, _, _, max_k = compute_pass_majority_metrics(tasks) + metrics.pop("per_sample_aggregate", None) + + all_rollouts = [rollout for task in tasks for rollout in task] + rollout_count = len(all_rollouts) + resolved_task_count = sum(1 for task in tasks if any(float(r.get("reward", 0.0) or 0.0) >= 1.0 for r in task)) + eval_error_count = sum(1 for rollout in all_rollouts if _metadata_dict(rollout).get("error")) + eval_report_count = sum(1 for rollout in all_rollouts if _eval_report_map(rollout)) + tests_status_count = sum(1 for rollout in all_rollouts if _eval_instance_report(rollout).get("tests_status")) + patch_applied_count = sum( + 1 for rollout in all_rollouts if _eval_instance_report(rollout).get("patch_successfully_applied") + ) + + metrics.update( + { + "task_count": len(tasks), + "rollout_count": rollout_count, + "max_rollouts_per_task": max_k, + "resolved_task_count": resolved_task_count, + "resolved_task_rate": 100.0 * resolved_task_count / len(tasks) if tasks else 0.0, + "eval_error_rollout_count": eval_error_count, + "eval_error_rate": 100.0 * eval_error_count / rollout_count if rollout_count else 0.0, + "eval_report_rollout_count": eval_report_count, + "eval_report_rate": 100.0 * eval_report_count / rollout_count if rollout_count else 0.0, + "tests_status_rollout_count": tests_status_count, + "tests_status_rate": 100.0 * tests_status_count / rollout_count if rollout_count else 0.0, + "patch_applied_rollout_count": patch_applied_count, + "patch_applied_rate": 100.0 * patch_applied_count / rollout_count if rollout_count else 0.0, + "per_task_metrics": self._compute_per_task_eval_metrics(tasks), + } + ) + + test_status_totals: dict[str, int] = {} + for rollout in all_rollouts: + for key, value in _test_status_counts(rollout).items(): + test_status_totals[key] = test_status_totals.get(key, 0) + value + metrics.update({f"tests_status/{key}": value for key, value in sorted(test_status_totals.items())}) + + return metrics + + def _compute_per_task_eval_metrics(self, tasks: list[list[dict[str, Any]]]) -> list[dict[str, Any]]: + per_task_metrics: list[dict[str, Any]] = [] + for fallback_idx, rollouts in enumerate(tasks): + if not rollouts: + continue + + first = rollouts[0] + task_index = first.get(TASK_INDEX_KEY_NAME, fallback_idx) + instance_id = first.get("instance_id") or _metadata_dict(first).get("instance_id") + resolved_count = sum(1 for rollout in rollouts if float(rollout.get("reward", 0.0) or 0.0) >= 1.0) + error_count = sum(1 for rollout in rollouts if _metadata_dict(rollout).get("error")) + eval_report_count = sum(1 for rollout in rollouts if _eval_report_map(rollout)) + tests_status_count = sum(1 for rollout in rollouts if _eval_instance_report(rollout).get("tests_status")) + patch_applied_count = sum( + 1 for rollout in rollouts if _eval_instance_report(rollout).get("patch_successfully_applied") + ) + + task_metrics: dict[str, Any] = { + TASK_INDEX_KEY_NAME: task_index, + "instance_id": instance_id, + "rollout_count": len(rollouts), + "resolved": resolved_count > 0, + "resolved_rollout_count": resolved_count, + "eval_error_rollout_count": error_count, + "eval_report_rollout_count": eval_report_count, + "tests_status_rollout_count": tests_status_count, + "patch_applied_rollout_count": patch_applied_count, + } + + test_status_totals: dict[str, int] = {} + for rollout in rollouts: + for key, value in _test_status_counts(rollout).items(): + test_status_totals[key] = test_status_totals.get(key, 0) + value + task_metrics.update({f"tests_status/{key}": value for key, value in sorted(test_status_totals.items())}) + per_task_metrics.append(task_metrics) + + return per_task_metrics + + def get_key_metrics(self, agent_metrics: dict[str, Any]) -> dict[str, Any]: + key_metrics: dict[str, Any] = {} + key_metrics.update(highest_k_metrics(agent_metrics, "pass@{k}", score_names=["accuracy"])) + key_metrics.update(highest_k_metrics(agent_metrics, "pass@1[avg-of-{k}]", score_names=["accuracy"])) + for key in ( + "mean/reward", + "resolved_task_count", + "task_count", + "resolved_task_rate", + "eval_error_rate", + "tests_status_rate", + ): + if key in agent_metrics: + key_metrics[key] = agent_metrics[key] + return key_metrics + + async def responses(self, body: NeMoGymResponseCreateParamsNonStreaming = Body()) -> NeMoGymResponse: + raise NotImplementedError + + async def run(self, body: MiniSWEAgentRunRequest) -> MiniSWEAgentVerifyResponse: + async with self.sem: + model_server_name = self.config.model_server.name + global_config_dict = ServerClient.load_from_global_config().global_config_dict + + model_server_config = get_first_server_config_dict( + global_config_dict, + model_server_name, + ) + + policy_model_name = global_config_dict["policy_model_name"] + + ##### MINI-SWE-AGENT CONFIG ##### + subset = body.subset + split = body.split + workers = 1 + run_golden = self.config.run_golden + base_url = f"http://{model_server_config['host']}:{model_server_config['port']}/v1" + dummy_key = "dummy_key" + model_name = f"hosted_vllm/{policy_model_name}" + step_timeout = self.config.step_timeout + eval_timeout = self.config.eval_timeout + step_limit = self.config.step_limit + + instance_id = body.instance_id + + mini_swe_config_path = _swebench_config_path() + config = yaml.safe_load(get_config_path(mini_swe_config_path).read_text()) + responses_create_params_dict = body.responses_create_params.model_dump(exclude_none=True) + + default_model_kwargs = config["model"]["model_kwargs"] + temperature = ( + body.responses_create_params.temperature + if body.responses_create_params.temperature is not None + else default_model_kwargs["temperature"] + ) + top_p = ( + body.responses_create_params.top_p + if body.responses_create_params.top_p is not None + else default_model_kwargs["top_p"] + ) + model_kwargs = _responses_create_params_to_model_kwargs( + responses_create_params_dict, + default_tool_choice=self.config.tool_choice, + ) + if model_kwargs: + config.setdefault("model", {}).setdefault("model_kwargs", {}).update(model_kwargs) + + output_file_dir = f"{Path.cwd()}/results/{subset}/{policy_model_name}" + config_path = mini_swe_config_path + should_write_config = bool(model_kwargs) + if self.config.sandbox_provider is None: + raise ValueError("mini_swe_agent_2 requires sandbox_provider") + config.setdefault("environment", {}).update(self.config.sandbox_environment_kwargs or {}) + config["environment"]["provider"] = self.config.sandbox_provider + config["environment"]["spec"] = _sandbox_spec_for_instance( + self.config.sandbox_spec, + resource_profiles=self.config.sandbox_resource_profiles, + instance_id=instance_id, + ) + should_write_config = True + + if should_write_config: + config_output_dir = Path(output_file_dir) / "_configs" + config_output_dir.mkdir(parents=True, exist_ok=True) + config_path = config_output_dir / f"{instance_id}.sandbox.yaml" + config_path.write_text(yaml.safe_dump(config, sort_keys=False)) + + if self.config.skip_if_exists: + if Path(f"{output_file_dir}/{instance_id}/{instance_id}.json").exists(): + with open(f"{output_file_dir}/{instance_id}/{instance_id}.json", "r") as f: + print(f"Skipping {instance_id} because it already exists") + verify_response = MiniSWEAgentVerifyResponse.model_validate_json(f.read()) + return verify_response + + #### RUN MINI-SWE-AGENT ##### + try: + params = dict( + subset=subset, + split=split, + workers=workers, + output=output_file_dir, + model=model_name, + api_key=dummy_key, + base_url=base_url, + env="sandbox", + run_golden=run_golden, + instance_id=instance_id, + config=config_path, + # TODO: add this later + instance_dict=body.model_dump(), + responses_create_params=json.dumps(responses_create_params_dict), + step_timeout=step_timeout, + eval_timeout=eval_timeout, + step_limit=step_limit, + ) + future = runner_ray_remote.remote(run_mini_swe_with_sandbox, params) + result = await asyncio.to_thread(ray.get, future) + result = result[instance_id] + input_messages = result["input_messages"] + response_output = result["response_output"] + responses = result["responses"] + reward = 1.0 if _is_resolved(instance_id, result["eval_report"]) else 0.0 + + except Exception as e: + error_info = {"error": str(e), "traceback": traceback.format_exc()} + print(f"Error running mini-swe-agent: {e}\n{error_info['traceback']}", flush=True) + result = {"eval_report": error_info} + input_messages = [] + response_output = [] + responses = [] + reward = 0.0 + + body.responses_create_params.input = input_messages + response = _default_response_object() + if responses: + response.update(dict(responses[-1])) + response.pop("extra", None) + response["model"] = policy_model_name + response["temperature"] = temperature + response["top_p"] = top_p + response["output"] = response_output + + verify_response = MiniSWEAgentVerifyResponse( + responses_create_params=body.responses_create_params, + reward=reward, + response=response, + instance_id=instance_id, + metadata=result.get("eval_report", {}) if result else {}, + ) + + output_path = Path(f"{output_file_dir}/{instance_id}") + output_path.mkdir(parents=True, exist_ok=True) + + with open(f"{output_file_dir}/{instance_id}/{instance_id}.json", "w") as f: + json.dump(verify_response.model_dump(), f) + + return verify_response + + +if __name__ == "__main__": + MiniSWEAgent.run_webserver() diff --git a/responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_ecs_fargate.yaml b/responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_ecs_fargate.yaml new file mode 100644 index 0000000000..89f8e3876c --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_ecs_fargate.yaml @@ -0,0 +1,39 @@ +mini_swe_agent_2: + responses_api_agents: + mini_swe_agent_2: + entrypoint: app.py + domain: coding + description: Software engineering tasks driven by mini-swe-agent harness on ECS Fargate. + value: Improve agentic software engineering capabilities. + model_server: + type: responses_api_models + name: policy_model + concurrency: 8 + env: sandbox + # Region-only config: the provider auto-discovers cluster/subnets/SGs/roles + # and the SSH-sidecar key ARNs from SSM (//ecs-sandbox/config, + # ssm_project defaults to "harbor"), exactly like NEL. + sandbox_provider: + ecs_fargate: + region: ${oc.env:AWS_REGION} + cpu: "2048" + memory: "8192" + ephemeral_storage_gib: 50 + sandbox_spec: + timeout_s: 18000 + ready_timeout_s: 1200 + metadata: + benchmark: swebench-verified + harness: mini-swe-agent + sandbox-api: ecs-fargate + sandbox_environment_kwargs: + cwd: /testbed + conda_env: testbed + activate_conda: true + user: root + delete: true + run_golden: false + step_timeout: 600 + eval_timeout: 1800 + skip_if_exists: false + step_limit: 250 diff --git a/responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_opensandbox.yaml b/responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_opensandbox.yaml new file mode 100644 index 0000000000..3c5e6f503b --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/configs/mini_swe_agent_opensandbox.yaml @@ -0,0 +1,64 @@ +mini_swe_agent_2: + responses_api_agents: + mini_swe_agent_2: + entrypoint: app.py + domain: coding + description: Software engineering tasks driven by mini-swe-agent harness on OpenSandbox. + value: Improve agentic software engineering capabilities. + model_server: + type: responses_api_models + name: policy_model + concurrency: 64 + env: sandbox + sandbox_provider: + opensandbox: + connection: + domain: opensandbox-server.opensandbox-system.svc.cluster.local + api_key: ${oc.env:OPENSANDBOX_API_KEY} + protocol: http + request_timeout_s: 300 + use_server_proxy: true + create: + request_timeout_s: 1200 + timeout_s: 1200 + skip_health_check: true + retries: 10 + retry_delay_s: 5.0 + retry_max_delay_s: 90.0 + probe: + timeout_s: 60 + deadline_s: 180 + stable_count: 2 + stable_delay_s: 1.0 + operations: + retries: 5 + retry_delay_s: 1.0 + retry_max_delay_s: 45.0 + command_retries: 3 + close_timeout_s: 30 + sandbox_spec: + timeout_s: 18000 + ready_timeout_s: 1200 + resources: + cpu: "2" + memory: 8Gi + ephemeral-storage: 20Gi + provider_options: + platform: + os: linux + arch: amd64 + metadata: + benchmark: swebench-verified + harness: mini-swe-agent + sandbox-api: opensandbox-sdk + sandbox_environment_kwargs: + cwd: /testbed + conda_env: testbed + activate_conda: true + user: root + delete: true + run_golden: false + step_timeout: 600 + eval_timeout: 1800 + skip_if_exists: false + step_limit: 250 diff --git a/responses_api_agents/mini_swe_agent_2/requirements.txt b/responses_api_agents/mini_swe_agent_2/requirements.txt new file mode 100644 index 0000000000..22fb4a3220 --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/requirements.txt @@ -0,0 +1,3 @@ +-e nemo-gym[dev,sandbox,sandbox-ecs] @ ../../ +mini-swe-agent==2.1.0 +swebench==4.1.0 diff --git a/responses_api_agents/mini_swe_agent_2/sandbox_environment.py b/responses_api_agents/mini_swe_agent_2/sandbox_environment.py new file mode 100644 index 0000000000..06f737dd42 --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/sandbox_environment.py @@ -0,0 +1,212 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""mini-swe-agent environment adapter backed by the Gym sandbox API.""" + +import os +import shlex +from dataclasses import dataclass, field +from typing import Any + + +try: + from minisweagent.exceptions import Submitted +except ModuleNotFoundError: + + class Submitted(Exception): + """Compatibility shim for local mini-swe-agent versions before v2.""" + + def __init__(self, *messages: dict[str, Any]) -> None: + self.messages = messages + super().__init__() + + +from nemo_gym.sandbox import Sandbox, SandboxSpec +from nemo_gym.sandbox.utils import rewrite_image + + +@dataclass +class MiniSWESandboxEnvironmentConfig: + """Configuration for mini-swe-agent runs inside a sandbox.""" + + image: str + cwd: str = "/workspace" + env: dict[str, str] = field(default_factory=dict) + forward_env: list[str] = field(default_factory=list) + timeout: int = 60 + step_timeout: int = 600 + eval_timeout: int = 1800 + interpreter: list[str] = field(default_factory=lambda: ["bash", "-c"]) + executable: str = "sandbox" + run_args: list[str] = field(default_factory=list) + start_args: list[str] = field(default_factory=list) + container_timeout: str = "2h" + instance_id: str | None = None + provider: dict[str, Any] = field(default_factory=dict) + spec: dict[str, Any] = field(default_factory=dict) + conda_env: str | None = None + activate_conda: bool = False + user: str | int | None = "root" + delete: bool = True + + +class MiniSWESandboxEnvironment: + """mini-swe-agent sync environment implemented with ``nemo_gym.sandbox.Sandbox``.""" + + def __init__( + self, + *, + config_class: type = MiniSWESandboxEnvironmentConfig, + **kwargs: Any, + ) -> None: + self.config = config_class(**kwargs) + if not self.config.provider: + raise ValueError("MiniSWESandboxEnvironment requires provider") + + self._sandbox: Sandbox | None = None + self._closed = False + + spec_config = dict(self.config.spec) + image = spec_config.pop("image", None) or self.config.image + image = rewrite_image(image, spec_config.pop("image_rewrites", [])) + provider_options = dict(spec_config.pop("provider_options", {})) + for option_key in ("platform", "volumes", "skip_health_check", "extensions"): + if option_key in spec_config: + provider_options[option_key] = spec_config.pop(option_key) + if "snapshot_id" in spec_config: + provider_options["snapshot_id"] = spec_config.pop("snapshot_id") + + env = dict(spec_config.pop("env", {})) + for key in self.config.forward_env: + value = os.getenv(key) + if value is not None: + env[key] = value + env.update(self.config.env) + + self._sandbox = Sandbox(self.config.provider).start( + SandboxSpec( + image=image, + timeout_s=spec_config.pop("timeout_s", None), + ready_timeout_s=spec_config.pop("ready_timeout_s", None), + workdir=spec_config.pop("workdir", self.config.cwd), + env=env, + files=spec_config.pop("files", {}), + metadata={ + **spec_config.pop("metadata", {}), + "nemo_gym_agent": "mini_swe_agent_2", + "instance_id": (self.config.instance_id or "unknown")[:63], + }, + resources=spec_config.pop("resources", {}), + entrypoint=spec_config.pop("entrypoint", None), + provider_options=provider_options, + ), + delete_on_stop=self.config.delete, + ) + + def get_template_vars(self, **kwargs: Any) -> dict[str, Any]: + return {**self.config.__dict__, **kwargs} + + def serialize(self) -> dict[str, Any]: + return { + "info": { + "config": { + "environment": self.config.__dict__, + "environment_type": f"{self.__class__.__module__}.{self.__class__.__name__}", + } + } + } + + def _command(self, command: str, cwd: str) -> str: + if not self.config.activate_conda or not self.config.conda_env: + return command + quoted_cwd = shlex.quote(cwd) + quoted_env = shlex.quote(self.config.conda_env) + # Resolve conda from common install roots before activating. The exec + # shell is non-login/non-interactive (e.g. `bash -c` on the ECS exec + # server), so `conda` is not on PATH and `conda info --base` cannot be + # relied on. Source the first available conda.sh (SWE-bench images ship + # /opt/miniconda3), falling back to `conda info --base` only when conda + # already happens to be on PATH. Using a grouped loop (not an `&&` + # chain) keeps a missing path from aborting the whole command. + return ( + f"cd {quoted_cwd} && " + '{ for __base in /opt/miniconda3 /opt/conda "$HOME/miniconda3" ' + '"$(command -v conda >/dev/null 2>&1 && conda info --base 2>/dev/null)"; do ' + '[ -n "$__base" ] && [ -f "$__base/etc/profile.d/conda.sh" ] && ' + '. "$__base/etc/profile.d/conda.sh" && break; done; } && ' + f"conda activate {quoted_env} && " + f"{command}" + ) + + def execute( + self, + action: dict[str, Any] | str, + cwd: str = "", + is_eval: bool = False, + timeout: int | None = None, + ) -> dict[str, Any]: + command = action.get("command", "") if isinstance(action, dict) else action + timeout_s = timeout or (self.config.eval_timeout if is_eval else self.config.step_timeout) + exec_cwd = cwd or self.config.cwd + if self._sandbox is None: + raise RuntimeError("Sandbox is not available") + + result = self._sandbox.exec( + self._command(command, exec_cwd), + timeout_s=timeout_s, + cwd="/", + user=self.config.user, + ) + output = "\n".join(part for part in (result.stdout, result.stderr) if part) + response = { + "output": output, + "returncode": result.return_code, + "exception_info": "", + } + self._check_finished(response) + return response + + def _check_finished(self, output: dict[str, Any]) -> None: + """Match mini-swe-agent's submit sentinel handling for sandbox-backed runs.""" + lines = output.get("output", "").lstrip().splitlines(keepends=True) + if lines and lines[0].strip() == "COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT" and output["returncode"] == 0: + submission = "".join(lines[1:]) + raise Submitted( + { + "role": "exit", + "content": submission, + "extra": {"exit_status": "Submitted", "submission": submission}, + } + ) + + def cleanup(self) -> None: + if self._closed: + return + self._closed = True + if self._sandbox is not None: + self._sandbox.stop() + self._sandbox = None + + def __enter__(self) -> "MiniSWESandboxEnvironment": + return self + + def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: + self.cleanup() + + def __del__(self) -> None: # pragma: no cover + if hasattr(self, "_closed") and not self._closed: + try: + self.cleanup() + except Exception: + pass diff --git a/responses_api_agents/mini_swe_agent_2/tests/test_app.py b/responses_api_agents/mini_swe_agent_2/tests/test_app.py new file mode 100644 index 0000000000..77feb64695 --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/tests/test_app.py @@ -0,0 +1,916 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import json +import sys +from pathlib import Path +from types import ModuleType, SimpleNamespace +from typing import Any, Dict, Optional +from unittest.mock import MagicMock, patch + +import pytest +import yaml +from fastapi.testclient import TestClient + +from nemo_gym.config_types import AggregateMetricsRequest, ModelServerRef +from nemo_gym.global_config import ROLLOUT_INDEX_KEY_NAME, TASK_INDEX_KEY_NAME +from nemo_gym.openai_utils import ( + NeMoGymChatCompletionCreateParamsNonStreaming, + NeMoGymResponseCreateParamsNonStreaming, +) +from nemo_gym.server_utils import ServerClient + + +try: + __import__("minisweagent.config") +except ModuleNotFoundError as exc: + if exc.name not in {"minisweagent", "minisweagent.config"}: + raise + minisweagent_module = ModuleType("minisweagent") + minisweagent_module.__path__ = [] + minisweagent_config_module = ModuleType("minisweagent.config") + minisweagent_config_module.builtin_config_dir = Path("/tmp/minisweagent/config") + minisweagent_config_module.get_config_path = Path + sys.modules["minisweagent"] = minisweagent_module + sys.modules["minisweagent.config"] = minisweagent_config_module + +from responses_api_agents.mini_swe_agent_2 import app as mini_swe_app_module +from responses_api_agents.mini_swe_agent_2.app import ( + MiniSWEAgent, + MiniSWEAgentConfig, + MiniSWEAgentRunRequest, + MiniSWEAgentVerifyResponse, + _is_resolved, + _json_dict_from_metadata, + _message_content_to_text, + _responses_create_params_to_model_kwargs, + _run_mini_swe_v2, + _sandbox_spec_for_instance, + _split_trajectory_for_responses, + _swebench_config_path, + _swebench_image_name, + run_mini_swe_with_sandbox, +) + + +DEFAULT_RUN_MINI_SWE_RESULT = { + "test_instance_123": { + "input_messages": [ + {"type": "message", "role": "system", "content": "You are a helpful assistant."}, + {"type": "message", "role": "user", "content": "Fix this bug."}, + ], + "response_output": [ + { + "id": "msg-1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "I'll help you fix the bug.", "annotations": []}], + } + ], + "responses": [ + { + "id": "resp-1", + "object": "response", + "output": [], + } + ], + "eval_report": { + "eval_report": { + "test_instance_123": { + "resolved": True, + "tests_status": { + "FAIL_TO_PASS": {"success": ["test1"], "failure": []}, + "PASS_TO_PASS": {"success": ["test2"], "failure": []}, + }, + } + } + }, + } +} + +DEFAULT_CONFIG_YAML = """ +model: + model_kwargs: + temperature: 0.5 + top_p: 0.8 +""" + +DEFAULT_CHAT_COMPLETION = { + "id": "chatcmpl-123", + "object": "chat.completion", + "created": 1677652288, + "model": "test_model", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello! How can I help you today?", + }, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 12, "total_tokens": 21}, +} + + +def create_test_config( + host: str = "0.0.0.0", + port: int = 8080, + model_name: str = "test_model", +) -> MiniSWEAgentConfig: + return MiniSWEAgentConfig( + name="mini_swe_agent_2", + host=host, + port=port, + entrypoint="", + model_server=ModelServerRef( + type="responses_api_models", + name=model_name, + ), + env="sandbox", + concurrency=1, + sandbox_provider={"opensandbox": {}}, + sandbox_spec={}, + ) + + +def setup_server_client_mocks(mock_load_from_global_config, mock_get_first_server_config_dict): + mock_server_client_instance = MagicMock() + mock_server_client_instance.global_config_dict = {"policy_model_name": "test_model"} + mock_load_from_global_config.return_value = mock_server_client_instance + + mock_get_first_server_config_dict.return_value = { + "host": "0.0.0.0", + "port": 8080, + } + + +def setup_config_path_mock(mock_get_config_path, config_yaml: str = DEFAULT_CONFIG_YAML): + mock_config_path = MagicMock() + mock_config_path.read_text.return_value = config_yaml + mock_get_config_path.return_value = mock_config_path + + +def setup_run_mini_swe_mock( + mock_to_thread, + mock_runner_ray_remote, + run_mini_swe_result: Dict[str, Any] = None, +): + """Setup mock for Ray-based run_mini_swe execution""" + if run_mini_swe_result is None: + run_mini_swe_result = DEFAULT_RUN_MINI_SWE_RESULT + + # Mock the Ray remote function to return a future-like object + mock_future = MagicMock() + mock_runner_ray_remote.remote.return_value = mock_future + + # Mock asyncio.to_thread (which calls ray.get) to return the result + mock_to_thread.return_value = run_mini_swe_result + + +def create_run_request( + instance_id: str = "test_instance_123", + temperature: float = 0.5, + top_p: float = 0.8, + max_output_tokens: int | None = None, + metadata: dict[str, Any] | None = None, + subset: str = "gym", + split: str = "train", + input_data: list = None, +) -> MiniSWEAgentRunRequest: + """Create a test run request with default values.""" + if input_data is None: + input_data = [] + + return MiniSWEAgentRunRequest( + instance_id=instance_id, + subset=subset, + split=split, + responses_create_params=NeMoGymResponseCreateParamsNonStreaming( + temperature=temperature, + top_p=top_p, + max_output_tokens=max_output_tokens, + metadata=metadata, + input=input_data, + ), + ) + + +def create_chat_completion_request( + model: str = "test_model", + messages: list = None, + temperature: float = 0.7, + max_tokens: Optional[int] = None, +) -> NeMoGymChatCompletionCreateParamsNonStreaming: + if messages is None: + messages = [{"role": "user", "content": "Hello!"}] + + kwargs = {"model": model, "messages": messages, "temperature": temperature} + if max_tokens is not None: + kwargs["max_tokens"] = max_tokens + + return NeMoGymChatCompletionCreateParamsNonStreaming(**kwargs) + + +def assert_run_response( + response: MiniSWEAgentVerifyResponse, + expected_reward: float = 1.0, + expected_temperature: float = 0.5, + expected_top_p: float = 0.8, + expected_input_length: int = 2, +): + assert isinstance(response, MiniSWEAgentVerifyResponse) + assert response.reward == expected_reward + assert response.responses_create_params.temperature == expected_temperature + assert response.responses_create_params.top_p == expected_top_p + assert len(response.responses_create_params.input) == expected_input_length + + if expected_input_length >= 2: + assert response.responses_create_params.input[0]["role"] == "system" + assert response.responses_create_params.input[1]["role"] == "user" + + +def assert_run_mini_swe_called( + mock_to_thread, + subset: str = "gym", + split: str = "train", + instance_id: str = "test_instance_123", +): + mock_to_thread.assert_called_once() + call_args = mock_to_thread.call_args + args = call_args[0] + assert len(args) >= 1 + + +class TestApp: + def test_sanity(self) -> None: + config = create_test_config(model_name="") + MiniSWEAgent(config=config, server_client=MagicMock(spec=ServerClient)) + + def test_response_param_helpers_cover_metadata_and_tool_choice_modes(self) -> None: + assert _json_dict_from_metadata(None, field_name="extra_body") == {} + assert _json_dict_from_metadata({"top_k": 20}, field_name="extra_body") == {"top_k": 20} + + kwargs = _responses_create_params_to_model_kwargs( + { + "temperature": 0.6, + "top_p": 0.95, + "max_output_tokens": 123, + "metadata": { + "extra_body": json.dumps({"top_k": 20}), + "chat_template_kwargs": json.dumps({"enable_thinking": True}), + }, + "tool_choice": {"type": "function", "function": {"name": "python"}}, + } + ) + + assert kwargs == { + "temperature": 0.6, + "top_p": 0.95, + "max_tokens": 123, + "extra_body": {"top_k": 20, "chat_template_kwargs": {"enable_thinking": True}}, + "tool_choice": {"type": "function", "function": {"name": "python"}}, + } + assert _responses_create_params_to_model_kwargs({"tool_choice": "bash"})["tool_choice"] == { + "type": "function", + "function": {"name": "bash"}, + } + assert ( + _responses_create_params_to_model_kwargs({"tool_choice": "auto"}, default_tool_choice="none")[ + "tool_choice" + ] + == "none" + ) + + with pytest.raises(ValueError, match="extra_body"): + _json_dict_from_metadata("[]", field_name="extra_body") + + def test_sandbox_resource_profiles_override_static_resources(self) -> None: + spec = _sandbox_spec_for_instance( + {"resources": {"cpu": "1", "memory": "8Gi", "ephemeral-storage": "20Gi"}}, + resource_profiles=[ + {"cpu": "250m", "memory": "3Gi", "ephemeral-storage": "1Gi"}, + {"cpu": "500m", "memory": "4Gi", "ephemeral-storage": "1Gi"}, + ], + instance_id="django__django-12345", + ) + + assert spec["resources"] in ( + {"cpu": "250m", "memory": "3Gi", "ephemeral-storage": "1Gi"}, + {"cpu": "500m", "memory": "4Gi", "ephemeral-storage": "1Gi"}, + ) + assert _sandbox_spec_for_instance(None, resource_profiles=None, instance_id="task") == {} + + def test_split_trajectory_and_resolution_helpers_cover_edge_cases(self) -> None: + input_messages, output_items, raw_responses = _split_trajectory_for_responses( + [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "user"}, + { + "role": "assistant", + "content": "answer", + "tool_calls": [{"id": "call-1", "function": {"name": "bash", "arguments": '{"command":"pwd"}'}}], + }, + {"role": "tool", "tool_call_id": "call-1", "content": "tool output"}, + {"type": "function_call_output", "call_id": "call-2", "output": "raw", "extra": {"ignored": True}}, + {"object": "response", "output": [{"type": "message", "content": "raw"}], "extra": {"ignored": True}}, + ] + ) + + assert input_messages == [ + {"type": "message", "role": "system", "content": "sys"}, + {"type": "message", "role": "user", "content": "user"}, + ] + assert any(item["type"] == "function_call" and item["call_id"] == "call-1" for item in output_items) + assert any(item["type"] == "function_call_output" and item["call_id"] == "call-1" for item in output_items) + assert any(item["type"] == "function_call_output" and item["call_id"] == "call-2" for item in output_items) + assert raw_responses == [{"object": "response", "output": [{"type": "message", "content": "raw"}]}] + + assert not _is_resolved("task", {}) + assert not _is_resolved("task", {"eval_report": {"task": {"resolved": True}}}) + assert not _is_resolved( + "task", + { + "eval_report": { + "task": { + "resolved": True, + "tests_status": {"FAIL_TO_PASS": {"success": [], "failure": []}}, + } + } + }, + ) + + def test_misc_mini_swe_helpers(self, monkeypatch, tmp_path) -> None: + assert _swebench_image_name({"instance_id": "django__django-1"}, "verified") == ( + "docker.io/swebench/sweb.eval.x86_64.django_1776_django-1:latest" + ) + assert _swebench_image_name({"instance_id": "django__django-1"}, "lite") == ( + "docker.io/xingyaoww/sweb.eval.x86_64.django_s_django-1:latest" + ) + assert _swebench_image_name({"instance_id": "x", "image_name": "custom:image"}, "verified") == "custom:image" + assert _message_content_to_text("hello") == "hello" + assert _message_content_to_text(None) == "" + assert _message_content_to_text([{"text": "one"}, {"content": "two"}, 3]) == "one\ntwo\n3" + + builtin_dir = tmp_path / "configs" + benchmark_dir = builtin_dir / "benchmarks" + benchmark_dir.mkdir(parents=True) + (benchmark_dir / "swebench.yaml").write_text("{}", encoding="utf-8") + monkeypatch.setattr(mini_swe_app_module, "builtin_config_dir", builtin_dir) + assert _swebench_config_path() == benchmark_dir / "swebench.yaml" + monkeypatch.setattr(mini_swe_app_module, "builtin_config_dir", tmp_path / "missing") + assert _swebench_config_path() == tmp_path / "missing" / "extra" / "swebench.yaml" + + def test_run_mini_swe_records_completion_and_errors(self, monkeypatch) -> None: + monkeypatch.setattr( + mini_swe_app_module, + "_run_mini_swe_v2", + lambda **_params: { + "task-1": { + "eval_report": { + "task-1": {"resolved": True}, + } + } + }, + ) + assert run_mini_swe_with_sandbox( + env="sandbox", + instance_id="task-1", + ) == {"task-1": {"eval_report": {"task-1": {"resolved": True}}}} + + def fail_runner(**_params): + raise RuntimeError("boom") + + monkeypatch.setattr(mini_swe_app_module, "_run_mini_swe_v2", fail_runner) + with pytest.raises(RuntimeError, match="boom"): + run_mini_swe_with_sandbox(env="sandbox", instance_id="task-1") + + monkeypatch.setattr( + mini_swe_app_module, + "_run_mini_swe_v2", + lambda **_params: {"task-1": {"eval_report": {"task-1": {"resolved": False}}}}, + ) + assert run_mini_swe_with_sandbox(env="sandbox", instance_id="task-1") == { + "task-1": {"eval_report": {"task-1": {"resolved": False}}} + } + + monkeypatch.setattr(mini_swe_app_module, "_run_mini_swe_v2", lambda **_params: {"task-1": "bad"}) + assert run_mini_swe_with_sandbox(env="sandbox", instance_id="task-1") == {"task-1": "bad"} + + def test_run_mini_swe_v2_success_and_golden_paths(self, monkeypatch, tmp_path) -> None: + holder: dict[str, Any] = {} + + class FakeLogger: + def info(self, _message: str) -> None: + return None + + def setup_logger(_instance_id: str, _log_file: Path) -> FakeLogger: + return FakeLogger() + + def make_test_spec(instance: dict[str, Any]) -> SimpleNamespace: + return SimpleNamespace( + instance_id=instance["instance_id"], + eval_script="#!/bin/bash\npytest -q", + ) + + def get_eval_report( + *, + test_spec: SimpleNamespace, + prediction: dict[str, Any], + test_log_path: str, + **_kwargs: Any, + ): + assert Path(test_log_path).exists() + return {test_spec.instance_id: {"resolved": True, "prediction": prediction}} + + class FakeEnv: + def __init__(self, config: dict[str, Any]) -> None: + self.config = config + self.commands: list[tuple[str, bool]] = [] + self.cleaned = False + + def execute(self, command: str, is_eval: bool = False) -> dict[str, Any]: + self.commands.append((command, is_eval)) + return {"output": "tests passed", "returncode": 0} + + def cleanup(self) -> None: + self.cleaned = True + + class FakeAgent: + def __init__(self, model: Any, env: FakeEnv, **agent_config: Any) -> None: + self.model = model + self.env = env + self.agent_config = agent_config + holder["agent_config"] = agent_config + + def run(self, problem_statement: str) -> dict[str, Any]: + assert problem_statement == "Fix the bug" + return {"exit_status": "submitted", "submission": "diff --git a/file b/file"} + + def save(self, path: Path | None, metadata: dict[str, Any]) -> dict[str, Any]: + holder["save_path"] = path + holder["save_metadata"] = metadata + if path is not None: + path.write_text("{}", encoding="utf-8") + return { + "messages": [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": [{"text": "problem"}]}, + { + "id": "resp-1", + "object": "response", + "output": [ + { + "id": "msg-1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "answer", "annotations": []}], + }, + { + "type": "function_call", + "name": "bash", + "call_id": "call-1", + "arguments": json.dumps({"command": "echo hi"}), + }, + ], + "extra": {"actions": [{"command": "echo hi", "tool_call_id": "call-1"}]}, + }, + { + "type": "function_call_output", + "call_id": "call-1", + "output": "tool output", + "extra": {"raw_output": "tool output"}, + }, + ] + } + + def get_environment(config: dict[str, Any]) -> FakeEnv: + env = FakeEnv(config) + holder["env"] = env + return env + + def get_model(config: dict[str, Any]) -> SimpleNamespace: + holder["model_config"] = config + return SimpleNamespace(config=config) + + module_specs = { + "swebench": ModuleType("swebench"), + "swebench.harness": ModuleType("swebench.harness"), + "swebench.harness.constants": ModuleType("swebench.harness.constants"), + "swebench.harness.docker_build": ModuleType("swebench.harness.docker_build"), + "swebench.harness.grading": ModuleType("swebench.harness.grading"), + "swebench.harness.test_spec": ModuleType("swebench.harness.test_spec"), + "swebench.harness.test_spec.test_spec": ModuleType("swebench.harness.test_spec.test_spec"), + "minisweagent.agents": ModuleType("minisweagent.agents"), + "minisweagent.agents.default": ModuleType("minisweagent.agents.default"), + "minisweagent.environments": ModuleType("minisweagent.environments"), + "minisweagent.models": ModuleType("minisweagent.models"), + } + module_specs["swebench.harness.constants"].SWEbenchInstance = dict + module_specs["swebench.harness.docker_build"].setup_logger = setup_logger + module_specs["swebench.harness.grading"].get_eval_report = get_eval_report + module_specs["swebench.harness.test_spec.test_spec"].make_test_spec = make_test_spec + module_specs["minisweagent.agents.default"].DefaultAgent = FakeAgent + module_specs["minisweagent.environments"].get_environment = get_environment + module_specs["minisweagent.models"].get_model = get_model + for name, module in module_specs.items(): + monkeypatch.setitem(sys.modules, name, module) + + config_path = tmp_path / "swebench.yaml" + config_path.write_text( + yaml.safe_dump( + { + "model": {"model_kwargs": {"max_output_tokens": 99}}, + "environment": {}, + "agent": {"step_limit": 1, "collapse_limit": 3}, + } + ), + encoding="utf-8", + ) + monkeypatch.setattr(mini_swe_app_module, "get_config_path", lambda _config: config_path) + monkeypatch.setattr(mini_swe_app_module, "uuid4", lambda: "uuid") + monkeypatch.setattr(mini_swe_app_module.time, "time", lambda: 1234) + + params = { + "instance_dict": { + "instance_id": "django__django-123", + "problem_statement": "Fix the bug", + "patch": "gold", + }, + "instance_id": "django__django-123", + "output": str(tmp_path / "out"), + "config": "swebench", + "model": "hosted/model", + "api_key": "key", # pragma: allowlist secret + "base_url": "http://model/v1", + "subset": "verified", + "step_timeout": 30, + "eval_timeout": 60, + "env": "sandbox", + "step_limit": 7, + "run_golden": False, + } + + result = _run_mini_swe_v2(**params) + + env = holder["env"] + assert env.cleaned is True + assert env.config["environment_class"].endswith("MiniSWESandboxEnvironment") + assert env.config["image"] == "docker.io/swebench/sweb.eval.x86_64.django_1776_django-123:latest" + assert holder["model_config"]["model_class"] == "litellm" + assert holder["model_config"]["model_name"] == "hosted/model" + assert holder["model_config"]["model_kwargs"]["max_tokens"] == 99 + assert holder["model_config"]["model_kwargs"]["base_url"] == "http://model/v1" + assert "api_base" not in holder["model_config"]["model_kwargs"] + assert holder["agent_config"]["step_limit"] == 7 + assert holder["save_metadata"] == {"instance_id": "django__django-123"} + assert result["django__django-123"]["input_messages"] == [ + {"type": "message", "role": "system", "content": "sys"}, + {"type": "message", "role": "user", "content": "problem"}, + ] + assert result["django__django-123"]["response_output"] == [ + { + "id": "msg-1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "answer", "annotations": []}], + }, + { + "type": "function_call", + "name": "bash", + "call_id": "call-1", + "arguments": json.dumps({"command": "echo hi"}), + }, + {"type": "function_call_output", "call_id": "call-1", "output": "tool output"}, + ] + assert result["django__django-123"]["responses"] == [ + { + "id": "resp-1", + "object": "response", + "output": [ + { + "id": "msg-1", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "answer", "annotations": []}], + }, + { + "type": "function_call", + "name": "bash", + "call_id": "call-1", + "arguments": json.dumps({"command": "echo hi"}), + }, + ], + } + ] + + golden_params = params | {"run_golden": True} + result = _run_mini_swe_v2(**golden_params) + + env = holder["env"] + assert env.cleaned is True + assert env.config["environment_class"].endswith("MiniSWESandboxEnvironment") + assert [command for command, _ in env.commands[:4]] == [ + "cat > patch.diff <<'EOF'\ngold\n\nEOF", + "git status --porcelain", + "git apply --check patch.diff", + "git apply patch.diff", + ] + assert result["django__django-123"]["exit_status"] == "Gold Patch Applied" + + string_params = params | { + "instance_dict": json.dumps( + {"instance_id": "django__django-123", "problem_statement": "Fix the bug", "patch": "gold"} + ), + } + assert "django__django-123" in _run_mini_swe_v2(**string_params) + + with pytest.raises(ValueError, match="instance_dict"): + _run_mini_swe_v2(**(params | {"instance_dict": None})) + + @patch("responses_api_agents.mini_swe_agent_2.app.ServerClient.load_from_global_config") + @patch("responses_api_agents.mini_swe_agent_2.app.get_first_server_config_dict") + @patch("responses_api_agents.mini_swe_agent_2.app.get_config_path") + @patch("responses_api_agents.mini_swe_agent_2.app.runner_ray_remote") + @patch("asyncio.to_thread") + async def test_run_successful_execution( + self, + mock_to_thread, + mock_runner_ray_remote, + mock_get_config_path, + mock_get_first_server_config_dict, + mock_load_from_global_config, + ) -> None: + """Test successful execution of the run method with mocked run_mini_swe.""" + + config = create_test_config() + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + setup_server_client_mocks(mock_load_from_global_config, mock_get_first_server_config_dict) + setup_config_path_mock(mock_get_config_path) + setup_run_mini_swe_mock(mock_to_thread, mock_runner_ray_remote) + + run_request = create_run_request() + + response = await server.run(run_request) + + assert_run_response(response) + + assert_run_mini_swe_called(mock_to_thread) + + @patch("responses_api_agents.mini_swe_agent_2.app.ServerClient.load_from_global_config") + @patch("responses_api_agents.mini_swe_agent_2.app.get_first_server_config_dict") + @patch("responses_api_agents.mini_swe_agent_2.app.get_config_path") + @patch("responses_api_agents.mini_swe_agent_2.app.runner_ray_remote") + @patch("asyncio.to_thread") + async def test_run_writes_generation_params_to_config( + self, + mock_to_thread, + mock_runner_ray_remote, + mock_get_config_path, + mock_get_first_server_config_dict, + mock_load_from_global_config, + tmp_path, + monkeypatch, + ) -> None: + monkeypatch.chdir(tmp_path) + config = create_test_config() + config.tool_choice = "bash" + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + setup_server_client_mocks(mock_load_from_global_config, mock_get_first_server_config_dict) + setup_config_path_mock(mock_get_config_path) + setup_run_mini_swe_mock(mock_to_thread, mock_runner_ray_remote) + + run_request = create_run_request( + temperature=0.6, + top_p=0.95, + max_output_tokens=49152, + metadata={ + "extra_body": '{"top_k":20,"min_p":0.0,"presence_penalty":0.0,"repetition_penalty":1.0}', + "chat_template_kwargs": '{"enable_thinking":true}', + }, + ) + + await server.run(run_request) + + call_args = mock_runner_ray_remote.remote.call_args + params = call_args.args[1] + generated_config = yaml.safe_load(Path(params["config"]).read_text()) + model_kwargs = generated_config["model"]["model_kwargs"] + assert model_kwargs["temperature"] == 0.6 + assert model_kwargs["top_p"] == 0.95 + assert model_kwargs["max_tokens"] == 49152 + assert "max_output_tokens" not in model_kwargs + assert model_kwargs["tool_choice"] == {"type": "function", "function": {"name": "bash"}} + assert model_kwargs["extra_body"] == { + "top_k": 20, + "min_p": 0.0, + "presence_penalty": 0.0, + "repetition_penalty": 1.0, + "chat_template_kwargs": {"enable_thinking": True}, + } + + @patch("responses_api_agents.mini_swe_agent_2.app.ServerClient.load_from_global_config") + @patch("responses_api_agents.mini_swe_agent_2.app.get_first_server_config_dict") + @patch("responses_api_agents.mini_swe_agent_2.app.get_config_path") + @patch("responses_api_agents.mini_swe_agent_2.app.runner_ray_remote") + @patch("asyncio.to_thread") + async def test_run_failed_execution( + self, + mock_to_thread, + mock_runner_ray_remote, + mock_get_config_path, + mock_get_first_server_config_dict, + mock_load_from_global_config, + ) -> None: + """Test run method when run_mini_swe fails.""" + + config = create_test_config() + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + setup_server_client_mocks(mock_load_from_global_config, mock_get_first_server_config_dict) + setup_config_path_mock(mock_get_config_path) + + # Mock Ray remote function + mock_future = MagicMock() + mock_runner_ray_remote.remote.return_value = mock_future + + # Mock asyncio.to_thread (ray.get) to raise an exception + mock_to_thread.side_effect = Exception("run_mini_swe failed") + + run_request = create_run_request(instance_id="test_instance_456", temperature=0.3, top_p=0.95) + + response = await server.run(run_request) + + assert_run_response( + response, + expected_reward=0.0, + expected_temperature=0.3, + expected_top_p=0.95, + expected_input_length=0, + ) + + assert_run_mini_swe_called(mock_to_thread, instance_id="test_instance_456") + + @patch("responses_api_agents.mini_swe_agent_2.app.ServerClient.load_from_global_config") + @patch("responses_api_agents.mini_swe_agent_2.app.get_first_server_config_dict") + @patch("responses_api_agents.mini_swe_agent_2.app.get_config_path") + @patch("responses_api_agents.mini_swe_agent_2.app.runner_ray_remote") + @patch("asyncio.to_thread") + async def test_run_mini_swe_not_found( + self, + mock_to_thread, + mock_runner_ray_remote, + mock_get_config_path, + mock_get_first_server_config_dict, + mock_load_from_global_config, + ) -> None: + config = create_test_config() + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + setup_server_client_mocks(mock_load_from_global_config, mock_get_first_server_config_dict) + setup_config_path_mock(mock_get_config_path) + + # Mock Ray remote function + mock_future = MagicMock() + mock_runner_ray_remote.remote.return_value = mock_future + + # Mock asyncio.to_thread (ray.get) to raise FileNotFoundError + mock_to_thread.side_effect = FileNotFoundError("run_mini_swe not found") + + run_request = create_run_request(instance_id="test_instance_789", temperature=0.2, top_p=1.0) + + response = await server.run(run_request) + + assert_run_response( + response, + expected_reward=0.0, + expected_temperature=0.2, + expected_top_p=1.0, + expected_input_length=0, + ) + + assert_run_mini_swe_called(mock_to_thread, instance_id="test_instance_789") + + async def test_responses_not_implemented(self) -> None: + config = create_test_config() + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + request_body = NeMoGymResponseCreateParamsNonStreaming(temperature=0.7, top_p=0.9, input=[]) + + with pytest.raises(NotImplementedError): + await server.responses(request_body) + + async def test_aggregate_metrics_includes_eval_results(self) -> None: + config = create_test_config() + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + responses = [ + { + TASK_INDEX_KEY_NAME: 0, + ROLLOUT_INDEX_KEY_NAME: 0, + "instance_id": "task-a", + "reward": 1.0, + "metadata": { + "instance_id": "task-a", + "eval_report": { + "task-a": { + "resolved": True, + "patch_successfully_applied": True, + "tests_status": { + "FAIL_TO_PASS": {"success": ["test-a"], "failure": []}, + "PASS_TO_PASS": {"success": ["test-b"], "failure": []}, + }, + } + }, + }, + }, + { + TASK_INDEX_KEY_NAME: 0, + ROLLOUT_INDEX_KEY_NAME: 1, + "instance_id": "task-a", + "reward": 0.0, + "metadata": { + "instance_id": "task-a", + "eval_report": { + "task-a": { + "resolved": False, + "patch_successfully_applied": True, + "tests_status": { + "FAIL_TO_PASS": {"success": [], "failure": ["test-a"]}, + "PASS_TO_PASS": {"success": ["test-b"], "failure": []}, + }, + } + }, + }, + }, + { + TASK_INDEX_KEY_NAME: 1, + ROLLOUT_INDEX_KEY_NAME: 0, + "instance_id": "task-b", + "reward": 0.0, + "metadata": {"error": "boom"}, + }, + { + TASK_INDEX_KEY_NAME: 1, + ROLLOUT_INDEX_KEY_NAME: 1, + "instance_id": "task-b", + "reward": 0.0, + "metadata": {"error": "boom"}, + }, + ] + + result = await server.aggregate_metrics(AggregateMetricsRequest(verify_responses=responses)) + + assert result.agent_metrics["pass@2/accuracy"] == pytest.approx(50.0) + assert result.agent_metrics["resolved_task_count"] == 1 + assert result.agent_metrics["eval_error_rollout_count"] == 2 + assert result.agent_metrics["tests_status/fail_to_pass_success"] == 1 + assert result.key_metrics["pass@2/accuracy"] == pytest.approx(50.0) + + groups = {group[TASK_INDEX_KEY_NAME]: group for group in result.group_level_metrics} + assert groups[0]["instance_id"] == "task-a" + assert groups[0]["resolved"] is True + assert groups[0]["tests_status_rollout_count"] == 2 + assert groups[1]["instance_id"] == "task-b" + assert groups[1]["eval_error_rollout_count"] == 2 + + def test_endpoints_registration(self) -> None: + config = create_test_config() + mock_server_client = MagicMock(spec=ServerClient) + server = MiniSWEAgent(config=config, server_client=mock_server_client) + + app = server.setup_webserver() + client = TestClient(app, raise_server_exceptions=False) + + response = client.post("/v1/responses", json={"temperature": 0.7, "top_p": 0.9, "input": []}) + assert response.status_code == 500 + + run_response = client.post("/run", json={}) + assert run_response.status_code != 404 + + aggregate_response = client.post("/aggregate_metrics", json={"verify_responses": []}) + assert aggregate_response.status_code == 200 diff --git a/responses_api_agents/mini_swe_agent_2/tests/test_sandbox_environment.py b/responses_api_agents/mini_swe_agent_2/tests/test_sandbox_environment.py new file mode 100644 index 0000000000..dc6ac45eb8 --- /dev/null +++ b/responses_api_agents/mini_swe_agent_2/tests/test_sandbox_environment.py @@ -0,0 +1,84 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from responses_api_agents.mini_swe_agent_2.sandbox_environment import MiniSWESandboxEnvironment, Submitted + + +def test_check_finished_raises_submitted_for_submit_sentinel() -> None: + env = MiniSWESandboxEnvironment.__new__(MiniSWESandboxEnvironment) + + try: + env._check_finished( + { + "output": "COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT\npatch contents\n", + "returncode": 0, + "exception_info": "", + } + ) + except Submitted as error: + assert error.messages == ( + { + "role": "exit", + "content": "patch contents\n", + "extra": {"exit_status": "Submitted", "submission": "patch contents\n"}, + }, + ) + else: + raise AssertionError("Expected Submitted") + + +def test_check_finished_ignores_nonzero_submit_sentinel() -> None: + env = MiniSWESandboxEnvironment.__new__(MiniSWESandboxEnvironment) + + env._check_finished( + { + "output": "COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT\npatch contents\n", + "returncode": 1, + "exception_info": "", + } + ) + + +def _env_with(activate_conda: bool, conda_env): + from responses_api_agents.mini_swe_agent_2.sandbox_environment import ( + MiniSWESandboxEnvironmentConfig, + ) + + env = MiniSWESandboxEnvironment.__new__(MiniSWESandboxEnvironment) + env.config = MiniSWESandboxEnvironmentConfig.__new__(MiniSWESandboxEnvironmentConfig) + env.config.activate_conda = activate_conda + env.config.conda_env = conda_env + return env + + +def test_command_passthrough_when_conda_disabled() -> None: + env = _env_with(activate_conda=False, conda_env="testbed") + assert env._command("git apply patch.diff", "/testbed") == "git apply patch.diff" + + +def test_command_resolves_conda_without_relying_on_path() -> None: + # The exec shell is non-login (conda not on PATH), so the wrapper must not + # depend solely on `conda info --base`; it sources conda.sh from known roots + # and must not use an `&&` chain that aborts when a root is missing. + env = _env_with(activate_conda=True, conda_env="testbed") + wrapped = env._command("git apply patch.diff", "/testbed") + + assert wrapped.startswith("cd /testbed && ") + assert "/opt/miniconda3/etc/profile.d/conda.sh" not in wrapped # built dynamically, not hardcoded mid-string + assert "/opt/miniconda3" in wrapped # but /opt/miniconda3 is one of the search roots + assert "conda activate testbed && git apply patch.diff" in wrapped + # The sourcing loop runs in a group so a missing root does not abort the command. + assert wrapped.count("&&") >= 2 + assert "for __base in" in wrapped diff --git a/tests/unit_tests/test_ecs_fargate_provider.py b/tests/unit_tests/test_ecs_fargate_provider.py new file mode 100644 index 0000000000..fe73607d5e --- /dev/null +++ b/tests/unit_tests/test_ecs_fargate_provider.py @@ -0,0 +1,372 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""ECS Fargate provider tests — all AWS/SSH/network calls mocked.""" + +from __future__ import annotations + +import contextlib +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from nemo_gym.sandbox import AsyncSandbox +from nemo_gym.sandbox.providers import ( + SandboxSpec, + SandboxStatus, + create_provider, + get_provider_class, + list_providers, +) +from nemo_gym.sandbox.providers.ecs_fargate import EcsFargateProvider, engine +from nemo_gym.sandbox.providers.ecs_fargate.provider import engine_config_from_mapping + + +_ENG = "nemo_gym.sandbox.providers.ecs_fargate.engine" + + +def _provider_config(**overrides): + cfg = dict( + region="us-west-2", + cluster="test-cluster", + subnets=["subnet-aaa"], + security_groups=["sg-bbb"], + assign_public_ip=True, + execution_role_arn="arn:aws:iam::1234:role/ecsTaskExec", + task_role_arn="arn:aws:iam::1234:role/ecsTask", + ssh_sidecar={ + "sshd_port": 2222, + "public_key_secret_arn": "arn:aws:secretsmanager:us-east-1:1234:secret:pub", + "private_key_secret_arn": "arn:aws:secretsmanager:us-east-1:1234:secret:priv", + "exec_server_port": 5000, + }, + ) + cfg.update(overrides) + return {"ecs_fargate": cfg} + + +@contextlib.contextmanager +def _mock_engine_start(exec_result=None): + """Patch every AWS/SSH seam so ``EcsFargateSandbox.start`` runs offline. + + Yields the fake exec client so delegation can be asserted. + """ + tunnel = MagicMock() + tunnel.local_port = 19000 + exec_client = MagicMock() + exec_client.exec = AsyncMock(return_value=exec_result or engine.ExecResult("out", "err", 0)) + exec_client.upload = AsyncMock() + exec_client.download = AsyncMock(return_value=b"payload") + exec_client.close = AsyncMock() + + with ( + patch.object(engine.EcsFargateSandbox, "_init_aws_clients"), + patch.object(engine.EcsFargateSandbox, "_resolve_image", return_value="python:3.12"), + patch.object(engine.EcsFargateSandbox, "_register_task_definition", return_value="task-def-arn"), + patch.object(engine.EcsFargateSandbox, "_run_task", return_value="task-arn"), + patch.object(engine.EcsFargateSandbox, "_register_for_cleanup"), + patch.object(engine.EcsFargateSandbox, "_wait_for_running"), + patch.object(engine.EcsFargateSandbox, "_get_task_public_ip", return_value="10.0.0.10"), + patch.object(engine.EcsFargateSandbox, "_wait_for_ssh_ready"), + patch(f"{_ENG}.download_secret_to_file", return_value="/tmp/key"), + patch(f"{_ENG}.download_secret_to_string", return_value="ssh-rsa fake"), + patch(f"{_ENG}.build_ssh_sidecar_container", return_value={"name": "ssh-tunnel"}), + patch(f"{_ENG}._free_port", return_value=19001), + patch(f"{_ENG}.SshTunnel", return_value=tunnel) as tunnel_cls, + patch(f"{_ENG}.ExecClient", return_value=exec_client), + ): + yield exec_client, tunnel_cls + + +# ── Registration ────────────────────────────────────────────────────── + + +def test_provider_registered(): + assert "ecs_fargate" in list_providers() + assert get_provider_class("ecs_fargate") is EcsFargateProvider + assert EcsFargateProvider.name == "ecs_fargate" + + +# ── Config resolution ───────────────────────────────────────────────── + + +def test_config_explicit_values_no_ssm(): + p = create_provider(_provider_config()) + cfg = p._cfg + assert cfg.cluster == "test-cluster" + assert cfg.subnets == ["subnet-aaa"] + assert cfg.assign_public_ip is True + assert cfg.ssh_sidecar.exec_server_port == 5000 + assert cfg.ssh_sidecar.sshd_port == 2222 + + +def test_config_ssm_autodiscovery_merges_and_yaml_wins(): + ssm_blob = { + "cluster": "ssm-cluster", + "subnets": ["subnet-ssm"], + "security_groups": ["sg-ssm"], + "execution_role_arn": "arn:ssm:exec", + "ssh_sidecar": { + "public_key_secret_arn": "arn:ssm:pub", + "private_key_secret_arn": "arn:ssm:priv", + "sshd_port": 52222, + }, + } + with patch(f"{_ENG}.resolve_ecs_config_from_ssm", return_value=ssm_blob) as resolve: + # region set, cluster omitted -> SSM is consulted. + p = create_provider({"ecs_fargate": {"region": "us-west-2", "subnets": ["subnet-override"]}}) + resolve.assert_called_once_with("us-west-2", "harbor") + cfg = p._cfg + assert cfg.cluster == "ssm-cluster" # filled from SSM + assert cfg.subnets == ["subnet-override"] # explicit YAML wins + assert cfg.execution_role_arn == "arn:ssm:exec" + assert cfg.ssh_sidecar.public_key_secret_arn == "arn:ssm:pub" + assert cfg.ssh_sidecar.sshd_port == 52222 + + +def test_config_no_ssm_when_cluster_present(): + with patch(f"{_ENG}.resolve_ecs_config_from_ssm") as resolve: + create_provider(_provider_config()) + resolve.assert_not_called() + + +def test_sidecar_missing_key_arns_raises(): + with pytest.raises(ValueError, match="public_key_secret_arn"): + engine_config_from_mapping({"cluster": "c", "ssh_sidecar": {"sshd_port": 2222}}) + + +# ── create / lifecycle delegation ───────────────────────────────────── + + +async def test_create_returns_handle_with_running_sandbox(): + provider = create_provider(_provider_config()) + spec = SandboxSpec(image="python:3.12") + with _mock_engine_start(): + handle = await provider.create(spec) + assert handle.provider_name == "ecs_fargate" + assert handle.sandbox_id == "task-arn" + assert isinstance(handle.raw, engine.EcsFargateSandbox) + assert await provider.status(handle) == SandboxStatus.RUNNING + await provider.close(handle) + assert await provider.status(handle) == SandboxStatus.STOPPED + + +async def test_exec_maps_engine_result(): + provider = create_provider(_provider_config()) + spec = SandboxSpec(image="python:3.12", workdir="/work") + with _mock_engine_start(exec_result=engine.ExecResult("hello\n", "", 0)) as (exec_client, _): + handle = await provider.create(spec) + result = await provider.exec(handle, "echo hello", timeout_s=42) + assert (result.stdout, result.stderr, result.return_code) == ("hello\n", "", 0) + # engine.exec received the gym timeout as timeout_sec + _, kwargs = exec_client.exec.call_args + assert kwargs["timeout"] == 42 + + +async def test_upload_and_download_delegate(tmp_path): + provider = create_provider(_provider_config()) + spec = SandboxSpec(image="python:3.12") + src = tmp_path / "in.txt" + src.write_text("data") + dest = tmp_path / "out.txt" + with _mock_engine_start() as (exec_client, _): + handle = await provider.create(spec) + await provider.upload_file(handle, src, "/remote/in.txt") + await provider.download_file(handle, "/remote/out.txt", dest) + exec_client.upload.assert_awaited_once() + assert dest.read_bytes() == b"payload" + + +async def test_outside_endpoints_build_reverse_tunnel(): + provider = create_provider(_provider_config()) + spec = SandboxSpec( + image="python:3.12", + provider_options={"outside_endpoints": [{"url": "http://127.0.0.1:4000/v1", "env_var": "MODEL_BASE_URL"}]}, + ) + with _mock_engine_start() as (_, tunnel_cls): + handle = await provider.create(spec) + # exec-server mode opens a forward tunnel to the exec server + a reverse + # tunnel for each outside endpoint. + _, kwargs = tunnel_cls.call_args + assert kwargs["forward_port"] == 5000 + # reverse spec format: "::" + assert any(s.endswith(":127.0.0.1:4000") for s in kwargs["reverses"]) + # the resolved endpoint is injected into the container env + routing = handle.raw._outside_endpoint_routing + assert routing.resolved_endpoint_url("MODEL_BASE_URL").startswith("http://127.0.0.1:") + + +async def test_create_missing_image_raises(): + provider = create_provider(_provider_config()) + from nemo_gym.sandbox.providers import SandboxCreateError + + with pytest.raises(SandboxCreateError, match="requires SandboxSpec.image"): + await provider.create(SandboxSpec(image=None)) + + +# ── Public AsyncSandbox surface ─────────────────────────────────────── + + +async def test_async_sandbox_end_to_end(): + spec = SandboxSpec(image="python:3.12", files={"/app/run.sh": "echo hi"}) + with _mock_engine_start() as (exec_client, _): + sb = AsyncSandbox(_provider_config(), spec) + await sb.start() + # initial files uploaded via the provider during start() + exec_client.upload.assert_awaited() + res = await sb.exec("ls") + assert res.return_code == 0 + assert await sb.status() == SandboxStatus.RUNNING + await sb.stop() + assert await sb.status() == SandboxStatus.STOPPED + + +# ── Pure engine helpers (no AWS) ────────────────────────────────────── + + +def test_task_def_hash_ignores_log_config(): + base = { + "family": "f1", + "containerDefinitions": [ + {"name": "main", "image": "x", "logConfiguration": {"options": {"awslogs-group": "g1"}}} + ], + } + other = { + "family": "f2", + "containerDefinitions": [ + {"name": "main", "image": "x", "logConfiguration": {"options": {"awslogs-group": "g2"}}} + ], + } + assert engine._compute_task_def_hash(base) == engine._compute_task_def_hash(other) + + +def test_generate_buildspec_pushes_to_ecr(): + cfg = engine.EcsFargateConfig( + region="us-west-2", + ecr_repository="123.dkr.ecr.us-west-2.amazonaws.com/sandbox", + ) + spec = engine.ImageBuilder._generate_buildspec(cfg, "sandbox", "tag1", f"{cfg.ecr_repository}:tag1") + assert "docker build -t sandbox:tag1" in spec + assert f"docker push {cfg.ecr_repository}:tag1" in spec + + +def _resolve_image_for(image, **cfg_overrides): + cfg = engine.EcsFargateConfig(region="us-east-1", **cfg_overrides) + sandbox = engine.EcsFargateSandbox(engine.SandboxSpec(image=image), ecs_config=cfg) + return sandbox._resolve_image() + + +def test_resolve_image_routes_bare_name_to_ecr_mirror(): + # Bare/public names are mirrored to the ECR tag, never pulled directly. + ecr = "463701203462.dkr.ecr.us-east-1.amazonaws.com/harbor-us-east-1" + resolved = _resolve_image_for("docker.io/swebench/sweb.eval.x86_64.astropy_1776_astropy-12907:latest", ecr_repository=ecr) + assert resolved == f"{ecr}:{engine._sanitize_id('docker.io/swebench/sweb.eval.x86_64.astropy_1776_astropy-12907:latest')}" + assert "docker.io" not in resolved.split(":", 1)[1] # origin registry not used for the pull + + +def test_resolve_image_passes_through_existing_ecr_ref(): + # A reference already in the ECR mirror is used verbatim (tag preserved, + # including the double underscore the sanitizer would otherwise collapse). + ecr = "463701203462.dkr.ecr.us-east-1.amazonaws.com/harbor-us-east-1" + existing = f"{ecr}:nel-harbor-tasks-swe-bench-astropy-1-1ccf0d50cb33__1ccf0d50" + assert _resolve_image_for(existing, ecr_repository=ecr) == existing + + +def test_resolve_image_template_takes_precedence(): + ecr = "463701203462.dkr.ecr.us-east-1.amazonaws.com/harbor-us-east-1" + resolved = _resolve_image_for( + "anything", ecr_repository=ecr, image_template="{task_id}-built" + ) + assert resolved == "anything-built" + + +def test_is_ecr_image_ref_matches_only_ecr_hosts(): + assert engine._is_ecr_image_ref("463701203462.dkr.ecr.us-east-1.amazonaws.com/repo:tag") + assert not engine._is_ecr_image_ref("docker.io/swebench/sweb.eval:latest") + assert not engine._is_ecr_image_ref("ubuntu:24.04") + + +def test_generate_mirror_buildspec_pulls_tags_and_pushes(): + cfg = engine.EcsFargateConfig( + region="us-east-1", + ecr_repository="123.dkr.ecr.us-east-1.amazonaws.com/mirror", + dockerhub_secret_arn="arn:aws:secretsmanager:us-east-1:123:secret:dh", + ) + src = "docker.io/swebench/sweb.eval.x86_64.astropy_1776_astropy-12907:latest" + ecr_url = f"{cfg.ecr_repository}:{engine._sanitize_id(src)}" + spec = engine.ImageBuilder._generate_mirror_buildspec(cfg, src, ecr_url) + assert f"docker pull --platform linux/amd64 {src}" in spec + assert f"docker tag {src} {ecr_url}" in spec + assert f"docker push {ecr_url}" in spec + assert "get-login-password" in spec # ECR login + assert "secretsmanager get-secret-value" in spec # Docker Hub login + + +def test_ensure_mirrored_skips_when_already_present(): + cfg = engine.EcsFargateConfig(region="us-east-1", ecr_repository="123.dkr.ecr.us-east-1.amazonaws.com/mirror") + with ( + patch.object(engine.ImageBuilder, "image_exists_in_ecr", return_value=True), + patch.object(engine.ImageBuilder, "run_buildspec_via_codebuild") as cb, + ): + url = engine.ImageBuilder.ensure_mirrored(cfg=cfg, src_image="ubuntu:24.04") + cb.assert_not_called() + assert url == f"{cfg.ecr_repository}:{engine._sanitize_id('ubuntu:24.04')}" + + +def test_ensure_mirrored_runs_codebuild_when_missing(): + cfg = engine.EcsFargateConfig(region="us-east-1", ecr_repository="123.dkr.ecr.us-east-1.amazonaws.com/mirror") + with ( + patch.object(engine.ImageBuilder, "image_exists_in_ecr", return_value=False), + patch.object(engine.ImageBuilder, "run_buildspec_via_codebuild") as cb, + ): + url = engine.ImageBuilder.ensure_mirrored(cfg=cfg, src_image="ubuntu:24.04") + cb.assert_called_once() + assert url.endswith(engine._sanitize_id("ubuntu:24.04")) + + +async def test_create_auto_mirrors_missing_public_image(): + ecr = "123.dkr.ecr.us-east-1.amazonaws.com/mirror" + provider = create_provider(_provider_config(ecr_repository=ecr)) + spec = SandboxSpec(image="docker.io/swebench/sweb.eval:latest") + with _mock_engine_start(), patch.object(engine.ImageBuilder, "ensure_mirrored") as m: + await provider.create(spec) + m.assert_called_once() + assert m.call_args.kwargs["src_image"] == "docker.io/swebench/sweb.eval:latest" + + +async def test_create_skips_mirror_for_existing_ecr_ref(): + ecr = "463701203462.dkr.ecr.us-east-1.amazonaws.com/mirror" + provider = create_provider(_provider_config(ecr_repository=ecr)) + spec = SandboxSpec(image=f"{ecr}:already-mirrored") + with _mock_engine_start(), patch.object(engine.ImageBuilder, "ensure_mirrored") as m: + await provider.create(spec) + m.assert_not_called() + + +async def test_create_skips_mirror_when_auto_mirror_disabled(): + ecr = "123.dkr.ecr.us-east-1.amazonaws.com/mirror" + provider = create_provider(_provider_config(ecr_repository=ecr, auto_mirror=False)) + spec = SandboxSpec(image="docker.io/swebench/sweb.eval:latest") + with _mock_engine_start(), patch.object(engine.ImageBuilder, "ensure_mirrored") as m: + await provider.create(spec) + m.assert_not_called() + + +def test_get_ecr_image_tag_is_content_addressed(tmp_path): + (tmp_path / "Dockerfile").write_text("FROM scratch") + tag1 = engine.ImageBuilder.get_ecr_image_tag(tmp_path, "env") + tag2 = engine.ImageBuilder.get_ecr_image_tag(tmp_path, "env") + assert tag1 == tag2 and tag1.startswith("env__") + (tmp_path / "Dockerfile").write_text("FROM alpine") + assert engine.ImageBuilder.get_ecr_image_tag(tmp_path, "env") != tag1 diff --git a/tests/unit_tests/test_opensandbox_provider.py b/tests/unit_tests/test_opensandbox_provider.py new file mode 100644 index 0000000000..36c4573007 --- /dev/null +++ b/tests/unit_tests/test_opensandbox_provider.py @@ -0,0 +1,667 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import asyncio +import builtins +from dataclasses import dataclass +from datetime import timedelta +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +import pytest + +from nemo_gym.sandbox.providers.base import SandboxSpec, SandboxStatus + + +pytest.importorskip("tenacity", reason="tenacity optional sandbox dependency is not installed") + +from nemo_gym.sandbox.providers.opensandbox import provider as opensandbox_provider + + +@dataclass(frozen=True) +class FakePlatformSpec: + os: str + arch: str + + +class FakeConnectionConfig: + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + + +@dataclass(frozen=True) +class FakeVolume: + name: str + + +class FakeSandbox: + created_kwargs: dict[str, Any] = {} + connected_args: tuple[Any, ...] = () + connected_kwargs: dict[str, Any] = {} + + def __init__(self, sandbox_id: str = "sandbox-1") -> None: + self.id = sandbox_id + + @classmethod + async def create(cls, *_args: Any, **kwargs: Any) -> "FakeSandbox": + cls.created_kwargs = kwargs + return cls() + + @classmethod + async def connect(cls, *args: Any, **kwargs: Any) -> "FakeSandbox": + cls.connected_args = args + cls.connected_kwargs = kwargs + return cls() + + +@pytest.fixture +def fake_opensandbox_sdk(monkeypatch: pytest.MonkeyPatch) -> None: + def require_sdk() -> tuple[Any, Any, Any, Any, Any]: + return FakeSandbox, FakeConnectionConfig, object, FakePlatformSpec, object + + monkeypatch.setattr(opensandbox_provider, "_require_opensandbox_sdk", require_sdk) + + +def test_sdk_import_helpers_and_retry_classification() -> None: + assert len(opensandbox_provider._require_opensandbox_sdk()) == 5 + assert len(opensandbox_provider._require_tenacity()) == 4 + + class StatusCodeError(Exception): + status_code = 429 + + assert opensandbox_provider._exception_status_code(StatusCodeError("rate limited")) == 429 + assert opensandbox_provider._is_retryable_create_error( + opensandbox_provider.OpenSandboxCreateError("create failed") + ) + + from opensandbox.exceptions import ( # noqa: PLC0415 + InvalidArgumentException, + SandboxApiException, + SandboxException, + SandboxInternalException, + ) + + assert opensandbox_provider._is_retryable_create_error(InvalidArgumentException("bad input")) is False + assert opensandbox_provider._is_retryable_create_error(SandboxInternalException("server failed")) is True + + retryable_api_error = SandboxApiException("busy") + retryable_api_error.status_code = 503 + assert opensandbox_provider._is_retryable_create_error(retryable_api_error) is True + + nonretryable_api_error = SandboxApiException("not found") + nonretryable_api_error.status_code = 404 + assert opensandbox_provider._is_retryable_create_error(nonretryable_api_error) is False + assert opensandbox_provider._is_retryable_create_error(SandboxException("gateway timeout")) is True + + retry_state = SimpleNamespace( + outcome=SimpleNamespace(exception=lambda: RuntimeError("temporary")), + next_action=SimpleNamespace(sleep=0.5), + attempt_number=2, + ) + opensandbox_provider._log_create_retry(retry_state) + + +def test_missing_optional_dependency_import_helpers(monkeypatch: pytest.MonkeyPatch) -> None: + real_import = builtins.__import__ + + def block_imports(*blocked_names: str) -> None: + def fake_import( + name: str, + globals_: dict[str, Any] | None = None, + locals_: dict[str, Any] | None = None, + fromlist: tuple[str, ...] = (), + level: int = 0, + ) -> Any: + if any(name == blocked or name.startswith(f"{blocked}.") for blocked in blocked_names): + raise ModuleNotFoundError(name) + return real_import(name, globals_, locals_, fromlist, level) + + monkeypatch.setattr(builtins, "__import__", fake_import) + + block_imports("opensandbox") + with pytest.raises(ModuleNotFoundError, match="OpenSandbox SDK is required"): + opensandbox_provider._require_opensandbox_sdk() + + block_imports("tenacity") + with pytest.raises(ModuleNotFoundError, match="tenacity is required"): + opensandbox_provider._require_tenacity() + + block_imports("opensandbox.exceptions") + assert opensandbox_provider._is_retryable_create_error(RuntimeError("gateway timeout")) is True + + +async def test_provider_conversion_helpers( + fake_opensandbox_sdk: None, + monkeypatch: pytest.MonkeyPatch, +) -> None: + connection_config = opensandbox_provider.OpenSandboxConnectionConfig(domain="sandbox.example") + assert ( + opensandbox_provider._coerce_config(connection_config, opensandbox_provider.OpenSandboxConnectionConfig) + is connection_config + ) + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (object, object, object, FakePlatformSpec, FakeVolume), + ) + assert opensandbox_provider._to_volumes([{"name": "workspace"}]) == [FakeVolume(name="workspace")] + + +async def test_direct_create_passes_platform_to_sdk_create( + fake_opensandbox_sdk: None, +) -> None: + provider = opensandbox_provider.OpenSandboxProvider( + connection={"request_timeout_s": 10}, + probe={"command": None}, + ) + + handle = await provider.create( + SandboxSpec( + image="mirror.gcr.io/astral/uv:python3.12-bookworm-slim", + provider_options={"platform": {"os": "linux", "arch": "amd64"}}, + ), + ) + + assert handle.sandbox_id == "sandbox-1" + assert FakeSandbox.created_kwargs["platform"] == FakePlatformSpec( + os="linux", + arch="amd64", + ) + + +def test_provider_validation_and_retry_helpers() -> None: + with pytest.raises(ValueError, match="image_pull_policy"): + opensandbox_provider.validate_image_pull_policy("Sometimes") + with pytest.raises(TypeError, match="must be a bool"): + opensandbox_provider._provider_option_bool({"skip_health_check": "true"}, "skip_health_check") + + assert opensandbox_provider._to_sandbox_status("starting") == SandboxStatus.STARTING + assert opensandbox_provider._to_sandbox_status("terminated") == SandboxStatus.STOPPED + assert opensandbox_provider._to_sandbox_status("failed") == SandboxStatus.ERROR + assert opensandbox_provider._to_sandbox_status(None) == SandboxStatus.UNKNOWN + + invalid_kwargs = [ + {"create": {"timeout_s": 0}}, + {"probe": {"timeout_s": 0}}, + {"probe": {"deadline_s": 0}}, + {"probe": {"stable_count": 0}}, + {"probe": {"stable_delay_s": -1}}, + {"create": {"retries": -1}}, + {"create": {"retry_delay_s": -1}}, + {"create": {"retry_max_delay_s": -1}}, + {"operations": {"retries": -1}}, + {"operations": {"retry_delay_s": -1}}, + {"operations": {"retry_max_delay_s": -1}}, + {"operations": {"command_retries": -1}}, + {"operations": {"close_timeout_s": 0}}, + {"create": {"connect_attempt_timeout_s": 0}}, + {"create": {"connect_poll_s": 0}}, + {"create": {"image_pull_policy": "Sometimes"}}, + ] + for kwargs in invalid_kwargs: + with pytest.raises(ValueError): + opensandbox_provider.OpenSandboxProvider(**kwargs) + with pytest.raises(TypeError): + opensandbox_provider.OpenSandboxProvider(**{"batch_" + "create_retries": 1}) + with pytest.raises(TypeError): + opensandbox_provider.OpenSandboxProvider(connection=object()) + + assert opensandbox_provider._exception_status_code(RuntimeError("HTTP status code: 503")) == 503 + assert opensandbox_provider._exception_status_code(RuntimeError("plain error")) is None + attrs = opensandbox_provider._sdk_error_attributes( + RuntimeError("HTTP 502 bad gateway"), + operation="exec", + sandbox_id="sandbox-1", + attempt_number=2, + max_attempts=3, + sleep_s=0.5, + ) + assert attrs["status_code"] == 502 + assert attrs["attempt_number"] == 2 + assert attrs["next_sleep_s"] == 0.5 + + +def test_connection_config_and_image_policy(fake_opensandbox_sdk: None) -> None: + provider = opensandbox_provider.OpenSandboxProvider( + connection={ + "domain": "sandbox.example", + "api_key": "key", # pragma: allowlist secret + "protocol": "https", + "request_timeout_s": 10, + "use_server_proxy": True, + } + ) + + config = provider._connection_config() + assert config.kwargs == { + "domain": "sandbox.example", + "api_key": "key", # pragma: allowlist secret + "protocol": "https", + "request_timeout": timedelta(seconds=10), + "use_server_proxy": True, + } + short_timeout_config = provider._connection_config(request_timeout_s=3) + assert short_timeout_config.kwargs["request_timeout"] == timedelta(seconds=3) + + spec = SandboxSpec(image="image:tag", provider_options={"extensions": {"imagePullPolicy": "Never"}}) + updated = provider._with_default_image_pull_policy(spec) + extensions = updated.provider_options["extensions"] + assert extensions["imagePullPolicy"] == "Never" + assert extensions["opensandbox.extensions.image-pull-policy"] == "Never" + + no_policy_provider = opensandbox_provider.OpenSandboxProvider(create={"image_pull_policy": None}) + assert no_policy_provider._with_default_image_pull_policy(spec) is spec + + +async def test_exec_file_operations_and_reference_validation(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + class FakeRunCommandOpts: + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + + class FakeLog: + def __init__(self, text: str) -> None: + self.text = text + + class FakeCommands: + def __init__(self) -> None: + self.calls: list[tuple[str, FakeRunCommandOpts]] = [] + + async def run(self, command: str, *, opts: FakeRunCommandOpts) -> Any: + self.calls.append((command, opts)) + if "fail" in command: + return SimpleNamespace( + logs=SimpleNamespace(stdout=[], stderr=[FakeLog("stderr")]), + error=SimpleNamespace(name="CommandError", value="failed"), + exit_code=None, + ) + return SimpleNamespace( + logs=SimpleNamespace(stdout=[FakeLog("stdout")], stderr=[]), + error=None, + exit_code=None, + ) + + class FakeFiles: + def __init__(self) -> None: + self.writes: list[tuple[str, str | bytes]] = [] + + async def write_file(self, target_path: str, data: str | bytes) -> None: + self.writes.append((target_path, data)) + + async def read_bytes(self, source_path: str) -> bytes: + return f"bytes:{source_path}".encode() + + class FakeRaw: + def __init__(self) -> None: + self.commands = FakeCommands() + self.files = FakeFiles() + + async def get_info(self) -> Any: + return SimpleNamespace(status=SimpleNamespace(state="RUNNING")) + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (object, object, FakeRunCommandOpts, object, object), + ) + + provider = opensandbox_provider.OpenSandboxProvider( + connection={"request_timeout_s": 5}, + probe={"command": None}, + ) + raw = FakeRaw() + handle = opensandbox_provider.SandboxHandle(sandbox_id="sandbox-1", provider_name="opensandbox", raw=raw) + + result = await provider.exec( + handle, + "echo hello", + cwd="/repo", + env={"A": "B"}, + timeout_s=2, + user=1000, + ) + assert result == opensandbox_provider.SandboxExecResult(stdout="stdout", stderr=None, return_code=0) + command, opts = raw.commands.calls[0] + assert command == "echo hello" + assert opts.kwargs == { + "working_directory": "/repo", + "envs": {"A": "B"}, + "timeout": timedelta(seconds=2), + "uid": 1000, + } + + result = await provider.exec(handle, "fail", user="agent") + assert result.return_code == 125 + assert result.error_type == "sandbox" + assert result.stderr == "stderr\nCommandError: failed" + assert raw.commands.calls[1][0] == "su -s /bin/sh -c fail agent" + + upload_path = tmp_path / "upload.txt" + upload_path.write_text("upload", encoding="utf-8") + await provider.upload_file(handle, upload_path, "/remote/upload.txt") + download_path = tmp_path / "nested" / "download.txt" + await provider.download_file(handle, "/remote/download.txt", download_path) + assert raw.files.writes == [("/remote/upload.txt", b"upload")] + assert download_path.read_bytes() == b"bytes:/remote/download.txt" + assert await provider.status(handle) == SandboxStatus.RUNNING + bare_handle = opensandbox_provider.SandboxHandle(sandbox_id="sandbox-2", provider_name="opensandbox", raw=object()) + assert await provider.status(bare_handle) == SandboxStatus.UNKNOWN + + +async def test_provider_create_probe_and_close_error_paths(monkeypatch: pytest.MonkeyPatch) -> None: + provider = opensandbox_provider.OpenSandboxProvider( + create={"connect_poll_s": 0.01}, + probe={ + "command": "probe", + "expected_stdout": "ready", + "timeout_s": 1, + "deadline_s": 0.01, + }, + ) + handle = opensandbox_provider.SandboxHandle(sandbox_id="sandbox-1", provider_name="opensandbox", raw=object()) + + async def bad_probe(*_args: Any, **_kwargs: Any) -> opensandbox_provider.SandboxExecResult: + return opensandbox_provider.SandboxExecResult(stdout="not ready", stderr="bad", return_code=1) + + async def no_sleep(_seconds: float) -> None: + return None + + monkeypatch.setattr(opensandbox_provider.asyncio, "sleep", no_sleep) + monkeypatch.setattr(provider, "_exec", bad_probe) + with pytest.raises(opensandbox_provider.OpenSandboxCreateVerificationError): + await provider._verify_created_handle(handle) + + provider = opensandbox_provider.OpenSandboxProvider( + probe={"command": "probe", "expected_stdout": None, "stable_count": 2, "stable_delay_s": 0.01}, + ) + sleep_calls: list[float] = [] + + async def record_sleep(seconds: float) -> None: + sleep_calls.append(seconds) + + async def good_probe(*_args: Any, **_kwargs: Any) -> opensandbox_provider.SandboxExecResult: + return opensandbox_provider.SandboxExecResult(stdout="ready", stderr=None, return_code=0) + + monkeypatch.setattr(opensandbox_provider.asyncio, "sleep", record_sleep) + monkeypatch.setattr(provider, "_exec", good_probe) + await provider._verify_created_handle(handle) + assert sleep_calls == [0.01] + + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": "probe"}) + + async def cancelled_probe(*_args: Any, **_kwargs: Any) -> opensandbox_provider.SandboxExecResult: + raise asyncio.CancelledError() + + monkeypatch.setattr(provider, "_exec", cancelled_probe) + with pytest.raises(asyncio.CancelledError): + await provider._verify_created_handle(handle) + + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": None}) + + async def close_raises(_handle: Any, *, delete: bool) -> None: + del delete + raise RuntimeError("close failed") + + monkeypatch.setattr(provider, "close", close_raises) + await provider._cleanup_failed_create_handle(handle) + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": None}) + + class DeleteAlreadyGoneRaw: + async def kill(self) -> None: + raise RuntimeError("sandbox sandbox-1 not found") + + async def close(self) -> None: + return None + + await provider.close( + opensandbox_provider.SandboxHandle( + sandbox_id="sandbox-1", + provider_name="opensandbox", + raw=DeleteAlreadyGoneRaw(), + ), + delete=True, + ) + + class DeleteAndCloseFailRaw: + async def kill(self) -> None: + raise RuntimeError("delete failed") + + async def close(self) -> None: + raise RuntimeError("close failed") + + with pytest.raises(RuntimeError, match="Failed to delete and close"): + await provider.close( + opensandbox_provider.SandboxHandle( + sandbox_id="sandbox-2", + provider_name="opensandbox", + raw=DeleteAndCloseFailRaw(), + ), + delete=True, + ) + + class DeleteFailsCloseSucceedsRaw: + async def kill(self) -> None: + raise RuntimeError("delete failed") + + async def close(self) -> None: + return None + + with pytest.raises(RuntimeError, match="delete failed"): + await provider.close( + opensandbox_provider.SandboxHandle( + sandbox_id="sandbox-3", + provider_name="opensandbox", + raw=DeleteFailsCloseSucceedsRaw(), + ), + delete=True, + ) + + +async def test_create_once_and_connect_after_create_error_paths( + fake_opensandbox_sdk: None, + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = opensandbox_provider.OpenSandboxProvider( + create={"timeout_s": 1, "skip_health_check": True}, + probe={"command": None}, + ) + monkeypatch.setattr(opensandbox_provider, "_to_volumes", lambda volumes: volumes) + spec = SandboxSpec( + image="image:tag", + timeout_s=10, + ready_timeout_s=20, + entrypoint=["/bin/sh"], + provider_options={ + "snapshot_id": "snapshot-1", + "platform": {"os": "linux", "arch": "amd64"}, + "volumes": [{"name": "workspace"}], + "skip_health_check": False, + }, + ) + handle = await provider._create_once(spec) + assert handle.sandbox_id == "sandbox-1" + assert FakeSandbox.created_kwargs["snapshot_id"] == "snapshot-1" + assert FakeSandbox.created_kwargs["timeout"] == timedelta(seconds=10) + assert FakeSandbox.created_kwargs["ready_timeout"] == timedelta(seconds=20) + assert FakeSandbox.created_kwargs["entrypoint"] == ["/bin/sh"] + assert FakeSandbox.created_kwargs["platform"] == FakePlatformSpec(os="linux", arch="amd64") + assert FakeSandbox.created_kwargs["volumes"] == [{"name": "workspace"}] + assert FakeSandbox.created_kwargs["skip_health_check"] is True + + class FailingConnectSandbox(FakeSandbox): + @classmethod + async def connect(cls, *args: Any, **kwargs: Any) -> "FakeSandbox": + del args, kwargs + raise ConnectionError("pod may still be starting") + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (FailingConnectSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider( + create={"connect_attempt_timeout_s": 0.01, "connect_poll_s": 0.01}, + probe={"command": None}, + ) + + async def no_sleep(_seconds: float) -> None: + return None + + monkeypatch.setattr(opensandbox_provider.asyncio, "sleep", no_sleep) + with pytest.raises(opensandbox_provider.OpenSandboxCreateTimeoutError): + await provider._connect_after_create( + opensandbox_provider.SandboxHandle(sandbox_id="sandbox-1", provider_name="opensandbox", raw=None), + SandboxSpec(image="image:tag"), + ) + + class CancelledConnectSandbox(FakeSandbox): + @classmethod + async def connect(cls, *args: Any, **kwargs: Any) -> "FakeSandbox": + del args, kwargs + raise asyncio.CancelledError() + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (CancelledConnectSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": None}) + with pytest.raises(asyncio.CancelledError): + await provider._connect_after_create( + opensandbox_provider.SandboxHandle(sandbox_id="sandbox-1", provider_name="opensandbox", raw=None), + SandboxSpec(image="image:tag", ready_timeout_s=1), + ) + + class NonRetryableConnectSandbox(FakeSandbox): + @classmethod + async def connect(cls, *args: Any, **kwargs: Any) -> "FakeSandbox": + del args, kwargs + raise ValueError("bad connection request") + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (NonRetryableConnectSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": None}) + with pytest.raises(ValueError, match="bad connection request"): + await provider._connect_after_create( + opensandbox_provider.SandboxHandle(sandbox_id="sandbox-1", provider_name="opensandbox", raw=None), + SandboxSpec(image="image:tag", ready_timeout_s=1), + ) + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (FakeSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider( + connection={"request_timeout_s": 3}, + probe={"command": None}, + ) + handle = await provider._create_once(SandboxSpec(image="image:tag", provider_options={"skip_health_check": True})) + assert handle.sandbox_id == "sandbox-1" + assert FakeSandbox.created_kwargs["skip_health_check"] is True + + class TimeoutSandbox(FakeSandbox): + @classmethod + async def create(cls, **_kwargs: Any) -> "FakeSandbox": + await asyncio.get_running_loop().create_future() + return cls() + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (TimeoutSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider( + create={"timeout_s": 0.01}, + probe={"command": None}, + ) + with pytest.raises(opensandbox_provider.OpenSandboxCreateTimeoutError): + await provider._create_once(SandboxSpec(image="image:tag")) + + class EmptyCreateSandbox(FakeSandbox): + @classmethod + async def create(cls, **_kwargs: Any) -> None: + return None + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (EmptyCreateSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": None}) + with pytest.raises(RuntimeError, match="returned no sandbox handle"): + await provider._create_once(SandboxSpec(image="image:tag")) + + monkeypatch.setattr( + opensandbox_provider, + "_require_opensandbox_sdk", + lambda: (FakeSandbox, FakeConnectionConfig, object, FakePlatformSpec, object), + ) + provider = opensandbox_provider.OpenSandboxProvider(probe={"command": "probe"}) + cleanup_calls: list[str] = [] + + async def fail_verify(_handle: opensandbox_provider.SandboxHandle) -> None: + raise RuntimeError("probe failed") + + async def cleanup(handle: opensandbox_provider.SandboxHandle) -> None: + cleanup_calls.append(handle.sandbox_id) + + monkeypatch.setattr(provider, "_verify_created_handle", fail_verify) + monkeypatch.setattr(provider, "_cleanup_failed_create_handle", cleanup) + with pytest.raises(RuntimeError, match="probe failed"): + await provider._create_once(SandboxSpec(image="image:tag")) + assert cleanup_calls == ["sandbox-1"] + + +async def test_retry_classification_and_await_sdk_helpers(monkeypatch: pytest.MonkeyPatch) -> None: + provider = opensandbox_provider.OpenSandboxProvider( + operations={"retries": 0}, + probe={"command": None}, + ) + assert await provider.aclose() is None + assert await provider._await_sdk_call(_return_value("ok"), operation="op", sandbox_id="sandbox-1", timeout_s=None) + assert opensandbox_provider._is_retryable_sdk_operation_error(TimeoutError("command timeout")) is False + assert opensandbox_provider._is_retryable_sdk_operation_error(ConnectionError("connection failed")) is True + wrapped = RuntimeError("wrapper") + wrapped.__cause__ = ConnectionError("connection reset") + assert opensandbox_provider._is_retryable_sdk_operation_error(wrapped) is True + wrapped.__cause__ = wrapped + assert opensandbox_provider._is_retryable_sdk_operation_error(wrapped) is False + + from opensandbox.exceptions import SandboxApiException # noqa: PLC0415 + + cyclic_api_error = SandboxApiException("proxy failed") + cyclic_api_error.status_code = 500 + cyclic_api_error.__cause__ = cyclic_api_error + assert opensandbox_provider._is_retryable_sdk_operation_error(cyclic_api_error) is True + + async def cancelled() -> None: + raise asyncio.CancelledError() + + with pytest.raises(asyncio.CancelledError): + await provider._await_sdk_operation( + cancelled, + operation="cancelled", + sandbox_id="sandbox-1", + timeout_s=None, + ) + + +async def _return_value(value: Any) -> Any: + return value diff --git a/tests/unit_tests/test_sandbox.py b/tests/unit_tests/test_sandbox.py new file mode 100644 index 0000000000..bf068d0f60 --- /dev/null +++ b/tests/unit_tests/test_sandbox.py @@ -0,0 +1,1009 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import asyncio +import importlib.util +from datetime import timedelta +from pathlib import Path +from typing import Any +from uuid import uuid4 + +import pytest + +import nemo_gym.sandbox.providers.registry as provider_registry +from nemo_gym.sandbox import ( + AsyncSandbox, + Sandbox, + SandboxCreateError, + SandboxExecResult, + SandboxHandle, + SandboxSpec, + SandboxStatus, + create_provider, + get_provider_class, + list_providers, + register_provider, +) +from nemo_gym.sandbox.api import _AsyncLoopRunner +from nemo_gym.sandbox.utils import rewrite_image +from responses_api_agents.mini_swe_agent_2.sandbox_environment import MiniSWESandboxEnvironment + + +def _has_module(module_name: str) -> bool: + try: + return importlib.util.find_spec(module_name) is not None + except ModuleNotFoundError: + return False + + +requires_tenacity = pytest.mark.skipif( + not _has_module("tenacity"), + reason="tenacity optional sandbox dependency is not installed", +) + + +def _require_opensandbox_provider() -> tuple[Any, Any, Any, str, str]: + pytest.importorskip("tenacity", reason="tenacity optional sandbox dependency is not installed") + from nemo_gym.sandbox.providers.opensandbox import provider as opensandbox_provider_module + from nemo_gym.sandbox.providers.opensandbox.provider import ( + IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY, + IMAGE_PULL_POLICY_EXTENSION_KEY, + OpenSandboxCreateVerificationError, + OpenSandboxProvider, + ) + + return ( + opensandbox_provider_module, + OpenSandboxProvider, + OpenSandboxCreateVerificationError, + IMAGE_PULL_POLICY_EXTENSION_KEY, + IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY, + ) + + +class FakeSandboxProvider: + name = "fake" + last_instance: "FakeSandboxProvider | None" = None + + def __init__(self, marker: str = "default") -> None: + self.marker = marker + self.created_specs: list[SandboxSpec] = [] + self.created_handles: list[SandboxHandle] = [] + self.exec_calls: list[dict[str, Any]] = [] + self.upload_calls: list[tuple[SandboxHandle, Path, str]] = [] + self.download_calls: list[tuple[SandboxHandle, str, Path]] = [] + self.closed: list[tuple[SandboxHandle, bool]] = [] + self.aclosed = False + FakeSandboxProvider.last_instance = self + + async def create(self, spec: SandboxSpec) -> SandboxHandle: + self.created_specs.append(spec) + handle = SandboxHandle( + sandbox_id=f"fake-{len(self.created_handles) + 1}", + provider_name=self.name, + raw={"spec": spec}, + ) + self.created_handles.append(handle) + return handle + + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + self.exec_calls.append( + { + "handle": handle, + "command": command, + "cwd": cwd, + "env": env, + "timeout_s": timeout_s, + "user": user, + } + ) + return SandboxExecResult(stdout="ok", stderr=None, return_code=0) + + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + self.upload_calls.append((handle, source_path, target_path)) + + async def download_file(self, handle: SandboxHandle, source_path: str, target_path: Path) -> None: + self.download_calls.append((handle, source_path, target_path)) + target_path.parent.mkdir(parents=True, exist_ok=True) + target_path.write_bytes(b"downloaded") + + async def status(self, handle: SandboxHandle) -> SandboxStatus: + del handle + return SandboxStatus.RUNNING + + async def close(self, handle: SandboxHandle, *, delete: bool) -> None: + self.closed.append((handle, delete)) + + async def aclose(self) -> None: + self.aclosed = True + + +class PlainSandboxProvider: + name = "plain" + + def __init__(self) -> None: + self.created_handles: list[SandboxHandle] = [] + + async def create(self, spec: SandboxSpec) -> SandboxHandle: + handle = SandboxHandle( + sandbox_id=f"plain-{len(self.created_handles) + 1}", + provider_name=self.name, + raw={"spec": spec}, + ) + self.created_handles.append(handle) + return handle + + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + del handle, command, cwd, env, timeout_s, user + return SandboxExecResult(stdout="ok", stderr=None, return_code=0) + + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + del handle, source_path, target_path + + async def download_file(self, handle: SandboxHandle, source_path: str, target_path: Path) -> None: + del handle, source_path + target_path.write_bytes(b"downloaded") + + async def status(self, handle: SandboxHandle) -> SandboxStatus: + del handle + return SandboxStatus.UNKNOWN + + async def close(self, handle: SandboxHandle, *, delete: bool = False) -> None: + del handle, delete + + async def aclose(self) -> None: + return None + + +class TransferOnlySandboxProvider: + name = "transfer-only" + + def __init__(self) -> None: + self.created_handles: list[SandboxHandle] = [] + self.upload_calls: list[tuple[SandboxHandle, Path, str]] = [] + self.download_calls: list[tuple[SandboxHandle, str, Path]] = [] + + async def create(self, spec: SandboxSpec) -> SandboxHandle: + handle = SandboxHandle( + sandbox_id=f"transfer-{len(self.created_handles) + 1}", + provider_name=self.name, + raw={"spec": spec}, + ) + self.created_handles.append(handle) + return handle + + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + del handle, command, cwd, env, timeout_s, user + return SandboxExecResult(stdout="ok", stderr=None, return_code=0) + + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + self.upload_calls.append((handle, source_path, target_path)) + + async def download_file(self, handle: SandboxHandle, source_path: str, target_path: Path) -> None: + self.download_calls.append((handle, source_path, target_path)) + target_path.write_bytes(b"fallback") + + async def status(self, handle: SandboxHandle) -> SandboxStatus: + del handle + return SandboxStatus.RUNNING + + async def close(self, handle: SandboxHandle, *, delete: bool = False) -> None: + del handle, delete + + async def aclose(self) -> None: + return None + + +class FailingUploadProvider(FakeSandboxProvider): + async def upload_file(self, handle: SandboxHandle, source_path: Path, target_path: str) -> None: + self.upload_calls.append((handle, source_path, target_path)) + raise RuntimeError("upload failed") + + +def test_sandbox_facade_uses_public_provider_api(tmp_path: Path) -> None: + asyncio.run(_assert_sandbox_facade_uses_public_provider_api(tmp_path)) + + +async def _assert_sandbox_facade_uses_public_provider_api(tmp_path: Path) -> None: + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + + sandbox = AsyncSandbox({provider_name: {"marker": "configured"}}) + await sandbox.start( + SandboxSpec( + image="image:tag", + metadata={"suite": "unit"}, + workdir="/repo", + files={"/tmp/bootstrap.txt": "hello"}, + ), + delete_on_stop=True, + ) + + provider = FakeSandboxProvider.last_instance + assert provider is not None + handle = provider.created_handles[0] + assert provider.marker == "configured" + assert provider.created_specs[0].image == "image:tag" + assert provider.created_specs[0].metadata == {"suite": "unit"} + assert provider.upload_calls[0][0] == handle + assert provider.upload_calls[0][2] == "/tmp/bootstrap.txt" + + result = await sandbox.exec("pytest -q", timeout_s=60, user="agent") + assert result == SandboxExecResult(stdout="ok", stderr=None, return_code=0) + assert provider.exec_calls[0] == { + "handle": handle, + "command": "pytest -q", + "cwd": "/repo", + "env": None, + "timeout_s": 60, + "user": "agent", + } + assert await sandbox.status() == SandboxStatus.RUNNING + + source_path = tmp_path / "source.txt" + target_path = tmp_path / "nested" / "target.txt" + source_path.write_text("local", encoding="utf-8") + await sandbox.upload(source_path, "/remote/source.txt") + await sandbox.download("/remote/source.txt", target_path) + assert provider.upload_calls[1] == (handle, source_path, "/remote/source.txt") + assert provider.download_calls == [(handle, "/remote/source.txt", target_path)] + assert target_path.read_bytes() == b"downloaded" + + await sandbox.stop() + await sandbox.stop() + assert provider.closed[-1] == (handle, True) + assert await sandbox.status() == SandboxStatus.STOPPED + assert provider.aclosed is True + + context_provider = FakeSandboxProvider() + async with AsyncSandbox(context_provider) as context_sandbox: + await context_sandbox.start(SandboxSpec(image="image:tag"), delete_on_stop=True) + context_handle = context_provider.created_handles[0] + assert context_provider.closed[-1] == (context_handle, True) + + +def test_async_sandbox_initial_file_error_paths() -> None: + asyncio.run(_assert_async_sandbox_initial_file_error_paths()) + + +async def _assert_async_sandbox_initial_file_error_paths() -> None: + failing_provider = FailingUploadProvider() + failing_sandbox = AsyncSandbox(failing_provider) + with pytest.raises(RuntimeError, match="upload failed"): + await failing_sandbox.start(SandboxSpec(image="image:tag", files={"/tmp/bootstrap.txt": "hello"})) + assert failing_provider.closed == [ + ( + SandboxHandle( + sandbox_id="fake-1", + provider_name="fake", + raw={"spec": SandboxSpec(image="image:tag", files={"/tmp/bootstrap.txt": "hello"})}, + ), + True, + ) + ] + + unstarted = AsyncSandbox(FakeSandboxProvider()) + with pytest.raises(RuntimeError, match="not been started"): + await unstarted.exec("pwd") + + started = AsyncSandbox(FakeSandboxProvider()) + await started.start(SandboxSpec(image="image:tag")) + with pytest.raises(RuntimeError, match="already started"): + await started.start(SandboxSpec(image="image:tag")) + await started.stop() + with pytest.raises(RuntimeError, match="has been stopped"): + await started.start(SandboxSpec(image="image:tag")) + + +def test_rewrite_image_validation() -> None: + assert rewrite_image(None, []) is None + assert rewrite_image("image:tag", [{"from": "other/", "to": "mirror/"}]) == "image:tag" + + +def test_provider_registry_validation_and_listing(monkeypatch: pytest.MonkeyPatch) -> None: + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + + assert get_provider_class(provider_name) is FakeSandboxProvider + assert "opensandbox" in list_providers() + assert provider_name in list_providers() + with pytest.raises(ValueError, match="must be non-empty"): + register_provider("", FakeSandboxProvider) + with pytest.raises(ValueError, match="already registered"): + register_provider(provider_name, FakeSandboxProvider) + with pytest.raises(ValueError, match="already registered"): + register_provider("opensandbox", FakeSandboxProvider) + register_provider(provider_name, FakeSandboxProvider, override=True) + with pytest.raises(ValueError, match="Unknown sandbox provider"): + get_provider_class(f"missing-{uuid4().hex}") + + builtin_name = f"builtin-{uuid4().hex}" + monkeypatch.setitem(provider_registry._BUILTIN_PROVIDER_LOADERS, builtin_name, lambda: FakeSandboxProvider) + assert get_provider_class(builtin_name) is FakeSandboxProvider + register_provider(builtin_name, PlainSandboxProvider, override=True) + assert get_provider_class(builtin_name) is PlainSandboxProvider + assert builtin_name in list_providers() + + +def test_create_provider_validation_and_constructor_cleanup() -> None: + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + provider = create_provider({provider_name: None}) + assert isinstance(provider, FakeSandboxProvider) + assert provider.marker == "default" + + with pytest.raises(ValueError, match="exactly one provider name"): + create_provider({}) + with pytest.raises(ValueError, match="non-empty string"): + create_provider({"": {}}) + with pytest.raises(TypeError, match="must be a mapping"): + create_provider({provider_name: "not-a-mapping"}) + + class FailingProvider(FakeSandboxProvider): + def __init__(self) -> None: + raise RuntimeError("provider constructor failed") + + failing_provider_name = f"failing-{uuid4().hex}" + register_provider(failing_provider_name, FailingProvider) + with pytest.raises(RuntimeError, match="provider constructor failed"): + Sandbox({failing_provider_name: {}}) + + +def test_async_sandbox_transfer_fallback_and_unknown_status(tmp_path: Path) -> None: + asyncio.run(_assert_async_sandbox_transfer_fallback_and_unknown_status(tmp_path)) + + +async def _assert_async_sandbox_transfer_fallback_and_unknown_status(tmp_path: Path) -> None: + transfer_provider = TransferOnlySandboxProvider() + transfer_sandbox = AsyncSandbox(transfer_provider) + await transfer_sandbox.start(SandboxSpec(image="image:tag", files={"/remote/inline.txt": "fallback"})) + transfer_handle = transfer_provider.created_handles[0] + assert transfer_provider.upload_calls[0][0] == transfer_handle + assert transfer_provider.upload_calls[0][2] == "/remote/inline.txt" + source_path = tmp_path / "source.txt" + target_path = tmp_path / "target.txt" + source_path.write_text("local", encoding="utf-8") + await transfer_sandbox.upload(source_path, "/remote/source.txt") + await transfer_sandbox.download("/remote/inline.txt", target_path) + assert transfer_provider.upload_calls[1] == (transfer_handle, source_path, "/remote/source.txt") + assert transfer_provider.download_calls == [(transfer_handle, "/remote/inline.txt", target_path)] + assert target_path.read_bytes() == b"fallback" + + plain_provider = PlainSandboxProvider() + plain_sandbox = AsyncSandbox(plain_provider) + await plain_sandbox.start(SandboxSpec(image="image:tag")) + assert await plain_sandbox.status() == SandboxStatus.UNKNOWN + + +def test_sync_sandbox_facade_uses_public_provider_api(tmp_path: Path) -> None: + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + + with Sandbox({provider_name: {"marker": "configured"}}) as sandbox: + sandbox.start( + SandboxSpec( + image="image:tag", + metadata={"suite": "unit"}, + workdir="/repo", + files={"/tmp/bootstrap.txt": "hello"}, + ), + delete_on_stop=True, + ) + + provider = FakeSandboxProvider.last_instance + assert provider is not None + handle = provider.created_handles[0] + assert provider.marker == "configured" + assert provider.created_specs[0].image == "image:tag" + assert provider.created_specs[0].metadata == {"suite": "unit"} + assert provider.upload_calls[0][0] == handle + assert provider.upload_calls[0][2] == "/tmp/bootstrap.txt" + + result = sandbox.exec("pytest -q", timeout_s=60, user="agent") + assert result == SandboxExecResult(stdout="ok", stderr=None, return_code=0) + assert provider.exec_calls[0] == { + "handle": handle, + "command": "pytest -q", + "cwd": "/repo", + "env": None, + "timeout_s": 60, + "user": "agent", + } + assert sandbox.status() == SandboxStatus.RUNNING + + upload_path = tmp_path / "sync-upload.txt" + upload_path.write_text("sync", encoding="utf-8") + download_path = tmp_path / "sync-download.txt" + sandbox.upload(upload_path, "/tmp/sync-upload.txt") + sandbox.download("/tmp/sync-download.txt", download_path) + assert download_path.read_bytes() == b"downloaded" + sandbox.stop() + assert provider.closed[-1] == (handle, True) + assert sandbox.status() == SandboxStatus.STOPPED + assert provider.aclosed is True + try: + sandbox.exec("pwd") + except RuntimeError as e: + assert "sync loop is closed" in str(e) + else: + raise AssertionError("expected closed sync sandbox to reject further calls") + + +def test_sync_loop_runner_close_is_idempotent() -> None: + runner = _AsyncLoopRunner() + runner.close() + runner.close() + + +def test_sync_sandbox_file_operations(tmp_path: Path) -> None: + provider = FakeSandboxProvider() + with Sandbox(provider) as sandbox: + sandbox.start(SandboxSpec(image="image:tag")) + handle = provider.created_handles[0] + source_path = tmp_path / "source.txt" + target_path = tmp_path / "target.txt" + source_path.write_text("local", encoding="utf-8") + sandbox.upload(source_path, "/remote/source.txt") + sandbox.download("/remote/source.txt", target_path) + + assert provider.upload_calls == [(handle, source_path, "/remote/source.txt")] + assert provider.download_calls == [(handle, "/remote/source.txt", target_path)] + assert target_path.read_bytes() == b"downloaded" + + +def test_sync_sandbox_facade_rejects_async_context() -> None: + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + + async def _create_sync_sandbox_in_async_context() -> None: + Sandbox({provider_name: {}}) + + try: + asyncio.run(_create_sync_sandbox_in_async_context()) + except RuntimeError as e: + assert "use AsyncSandbox in async code" in str(e) + else: + raise AssertionError("expected sync Sandbox to reject async context") + + +@requires_tenacity +def test_opensandbox_sdk_create_receives_default_image_pull_policy(monkeypatch) -> None: + asyncio.run(_assert_opensandbox_sdk_create_receives_default_image_pull_policy(monkeypatch)) + + +async def _assert_opensandbox_sdk_create_receives_default_image_pull_policy(monkeypatch) -> None: + ( + opensandbox_provider_module, + OpenSandboxProvider, + _OpenSandboxCreateVerificationError, + IMAGE_PULL_POLICY_EXTENSION_KEY, + IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY, + ) = _require_opensandbox_provider() + del _OpenSandboxCreateVerificationError + + class FakeSDKSandbox: + create_calls: list[dict[str, Any]] = [] + + def __init__(self, sandbox_id: str) -> None: + self.id = sandbox_id + + @classmethod + async def create(cls, **kwargs: Any) -> "FakeSDKSandbox": + cls.create_calls.append(kwargs) + return cls("sdk-sandbox-1") + + monkeypatch.setattr( + opensandbox_provider_module, + "_require_opensandbox_sdk", + lambda: (FakeSDKSandbox, object, object, object, object), + ) + + provider = OpenSandboxProvider(probe={"command": None}) + monkeypatch.setattr(provider, "_connection_config", lambda request_timeout_s=None: object()) + + handle = await provider.create( + SandboxSpec( + image="image:tag", + metadata={ + "harbor_instance_id": "swebench::django__django-10880", + "long": f"bad:{'x' * 80}:", + }, + ) + ) + + assert handle.sandbox_id == "sdk-sandbox-1" + metadata = FakeSDKSandbox.create_calls[0]["metadata"] + assert metadata["harbor_instance_id"] == "swebench_django__django-10880" + assert metadata["long"] == ("bad_" + "x" * 59) + extensions = FakeSDKSandbox.create_calls[0]["extensions"] + assert extensions[IMAGE_PULL_POLICY_EXTENSION_KEY] == "IfNotPresent" + assert extensions[IMAGE_PULL_POLICY_ANNOTATION_EXTENSION_KEY] == "IfNotPresent" + + +@requires_tenacity +def test_opensandbox_connect_after_create_uses_connection_config(monkeypatch) -> None: + asyncio.run(_assert_opensandbox_connect_after_create_uses_connection_config(monkeypatch)) + + +async def _assert_opensandbox_connect_after_create_uses_connection_config(monkeypatch) -> None: + opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + class FakeConnectionConfig: + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + + class FakeSDKSandbox: + connect_calls: list[dict[str, Any]] = [] + + def __init__(self, sandbox_id: str) -> None: + self.id = sandbox_id + + @classmethod + async def connect(cls, sandbox_id: str, **kwargs: Any) -> "FakeSDKSandbox": + cls.connect_calls.append({"sandbox_id": sandbox_id, **kwargs}) + return cls(sandbox_id) + + monkeypatch.setattr( + opensandbox_provider_module, + "_require_opensandbox_sdk", + lambda: (FakeSDKSandbox, FakeConnectionConfig, object, object, object), + ) + + provider = OpenSandboxProvider( + connection={"domain": "sandbox.example", "protocol": "https"}, + create={"connect_attempt_timeout_s": 1}, + probe={"command": None}, + ) + handle = await provider._connect_after_create( + SandboxHandle(sandbox_id="sdk-sandbox-1", provider_name="opensandbox", raw=None), + SandboxSpec(image="image:tag", ready_timeout_s=10), + ) + + assert handle.sandbox_id == "sdk-sandbox-1" + assert isinstance(handle.raw, FakeSDKSandbox) + connect_call = FakeSDKSandbox.connect_calls[0] + assert connect_call["skip_health_check"] is True + assert connect_call["connection_config"].kwargs == { + "domain": "sandbox.example", + "protocol": "https", + "request_timeout": timedelta(seconds=1), + } + + +@requires_tenacity +def test_opensandbox_create_probe_can_require_stable_successes(monkeypatch) -> None: + asyncio.run(_assert_opensandbox_create_probe_can_require_stable_successes(monkeypatch)) + + +async def _assert_opensandbox_create_probe_can_require_stable_successes(monkeypatch) -> None: + _opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + provider = OpenSandboxProvider( + probe={ + "command": "true", + "expected_stdout": None, + "stable_count": 3, + "stable_delay_s": 0, + }, + ) + calls: list[dict[str, Any]] = [] + + async def fake_exec( + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + calls.append( + { + "handle": handle, + "command": command, + "cwd": cwd, + "env": env, + "timeout_s": timeout_s, + "user": user, + } + ) + return SandboxExecResult(stdout="", stderr="", return_code=0) + + monkeypatch.setattr(provider, "_exec", fake_exec) + handle = SandboxHandle(sandbox_id="sdk-sandbox-0", provider_name="opensandbox", raw=object()) + + await provider._verify_created_handle(handle) + + assert [call["command"] for call in calls] == ["true", "true", "true"] + assert all(call["timeout_s"] == 30 for call in calls) + assert all(call["user"] == "root" for call in calls) + + +@requires_tenacity +def test_opensandbox_create_probe_polls_same_sandbox_after_transient_errors(monkeypatch) -> None: + asyncio.run(_assert_opensandbox_create_probe_polls_same_sandbox_after_transient_errors(monkeypatch)) + + +async def _assert_opensandbox_create_probe_polls_same_sandbox_after_transient_errors(monkeypatch) -> None: + _opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + provider = OpenSandboxProvider( + create={"connect_poll_s": 0.01}, + probe={ + "command": "true", + "expected_stdout": None, + "timeout_s": 1, + "deadline_s": 2, + "stable_count": 2, + "stable_delay_s": 0, + }, + ) + attempts = 0 + handles: list[SandboxHandle] = [] + + async def fake_exec( + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + del command, cwd, env, timeout_s, user + nonlocal attempts + attempts += 1 + handles.append(handle) + if attempts <= 2: + raise ConnectionError("direct execd endpoint is not accepting connections yet") + return SandboxExecResult(stdout="", stderr="", return_code=0) + + monkeypatch.setattr(provider, "_exec", fake_exec) + handle = SandboxHandle(sandbox_id="sdk-sandbox-0", provider_name="opensandbox", raw=object()) + + await provider._verify_created_handle(handle) + + assert attempts == 4 + assert {seen_handle.sandbox_id for seen_handle in handles} == {"sdk-sandbox-0"} + + +def test_opensandbox_create_probe_failures_are_retryable() -> None: + ( + opensandbox_provider_module, + _OpenSandboxProvider, + OpenSandboxCreateVerificationError, + *_unused, + ) = _require_opensandbox_provider() + + error = OpenSandboxCreateVerificationError("pod sdk-sandbox-0 failed create probe") + + assert isinstance(error, SandboxCreateError) + assert opensandbox_provider_module._is_retryable_create_error(error) is True + + +def test_opensandbox_starting_pod_endpoint_errors_are_retryable() -> None: + opensandbox_provider_module, *_unused = _require_opensandbox_provider() + + error = RuntimeError( + "Get endpoint for sandbox sdk-sandbox-0 port 44772 failed: " + "Pod IP is not yet available. The Pod may still be starting." + ) + + assert opensandbox_provider_module._is_retryable_create_error(error) is True + + +@requires_tenacity +def test_opensandbox_exec_retries_retryable_sdk_failures(monkeypatch) -> None: + asyncio.run(_assert_opensandbox_exec_retries_retryable_sdk_failures(monkeypatch)) + + +async def _assert_opensandbox_exec_retries_retryable_sdk_failures(monkeypatch) -> None: + opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + class FakeRunCommandOpts: + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + + class FakeLog: + def __init__(self, text: str) -> None: + self.text = text + + class FakeLogs: + stdout = [FakeLog("ok")] + stderr: list[FakeLog] = [] + + class FakeExecution: + logs = FakeLogs() + error = None + exit_code = 0 + + class FakeCommands: + def __init__(self) -> None: + self.calls = 0 + + async def run(self, command: str, *, opts: FakeRunCommandOpts) -> FakeExecution: + del command, opts + self.calls += 1 + if self.calls <= 2: + raise ConnectionError("transient connection failure") + return FakeExecution() + + class FakeRaw: + def __init__(self) -> None: + self.commands = FakeCommands() + + monkeypatch.setattr( + opensandbox_provider_module, + "_require_opensandbox_sdk", + lambda: (object, object, FakeRunCommandOpts, object, object), + ) + + provider = OpenSandboxProvider( + operations={ + "retries": 2, + "retry_delay_s": 0, + "retry_max_delay_s": 0, + "command_retries": 2, + }, + probe={"command": None}, + ) + raw = FakeRaw() + handle = SandboxHandle(sandbox_id="sdk-sandbox-1", provider_name="opensandbox", raw=raw) + + result = await provider.exec(handle, "echo hello", timeout_s=30) + + assert result.stdout == "ok" + assert result.return_code == 0 + assert raw.commands.calls == 3 + + +@requires_tenacity +def test_opensandbox_command_retries_can_be_disabled(monkeypatch) -> None: + asyncio.run(_assert_opensandbox_command_retries_can_be_disabled(monkeypatch)) + + +async def _assert_opensandbox_command_retries_can_be_disabled(monkeypatch) -> None: + opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + class FakeRunCommandOpts: + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + + class FakeCommands: + def __init__(self) -> None: + self.calls = 0 + + async def run(self, command: str, *, opts: FakeRunCommandOpts) -> None: + del command, opts + self.calls += 1 + raise ConnectionError("transient connection failure") + + class FakeRaw: + def __init__(self) -> None: + self.commands = FakeCommands() + + monkeypatch.setattr( + opensandbox_provider_module, + "_require_opensandbox_sdk", + lambda: (object, object, FakeRunCommandOpts, object, object), + ) + + provider = OpenSandboxProvider( + operations={ + "retries": 2, + "retry_delay_s": 0, + "retry_max_delay_s": 0, + "command_retries": 0, + }, + probe={"command": None}, + ) + raw = FakeRaw() + handle = SandboxHandle(sandbox_id="sdk-sandbox-1", provider_name="opensandbox", raw=raw) + + try: + await provider.exec(handle, "echo hello", timeout_s=30) + except ConnectionError: + pass + else: + raise AssertionError("expected provider.exec to propagate the command failure") + + assert raw.commands.calls == 1 + + +@requires_tenacity +def test_opensandbox_close_timeout_does_not_fail_after_delete() -> None: + asyncio.run(_assert_opensandbox_close_timeout_does_not_fail_after_delete()) + + +async def _assert_opensandbox_close_timeout_does_not_fail_after_delete() -> None: + _opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + class SlowCloseRaw: + def __init__(self) -> None: + self.killed = False + + async def kill(self) -> None: + self.killed = True + + async def close(self) -> None: + await asyncio.sleep(60) + + raw = SlowCloseRaw() + provider = OpenSandboxProvider( + operations={"close_timeout_s": 0.01}, + probe={"command": None}, + ) + handle = SandboxHandle(sandbox_id="sdk-sandbox-1", provider_name="opensandbox", raw=raw) + + await provider.close(handle, delete=True) + + assert raw.killed is True + + +@requires_tenacity +def test_opensandbox_close_timeout_still_fails_without_delete() -> None: + asyncio.run(_assert_opensandbox_close_timeout_still_fails_without_delete()) + + +async def _assert_opensandbox_close_timeout_still_fails_without_delete() -> None: + _opensandbox_provider_module, OpenSandboxProvider, *_unused = _require_opensandbox_provider() + + class SlowCloseRaw: + async def close(self) -> None: + await asyncio.sleep(60) + + provider = OpenSandboxProvider( + operations={"close_timeout_s": 0.01}, + probe={"command": None}, + ) + handle = SandboxHandle(sandbox_id="sdk-sandbox-1", provider_name="opensandbox", raw=SlowCloseRaw()) + + try: + await provider.close(handle, delete=False) + except TimeoutError: + pass + else: + raise AssertionError("expected close timeout to fail when delete=False") + + +def test_mini_swe_sandbox_environment_owns_conda_setup(monkeypatch) -> None: + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + monkeypatch.setenv("FORWARDED_KEY", "forwarded-value") + + env = MiniSWESandboxEnvironment( + image="upstream/image:tag", + cwd="/testbed", + provider={provider_name: {"marker": "configured"}}, + spec={ + "image_rewrites": [{"from": "upstream/", "to": "mirror/"}], + "metadata": {"suite": "unit"}, + "resources": {"cpu": "1"}, + }, + env={"STATIC_KEY": "static-value"}, + forward_env=["FORWARDED_KEY"], + conda_env="testbed", + activate_conda=True, + user="agent", + delete=True, + ) + + try: + assert env.get_template_vars(extra="value")["extra"] == "value" + serialized = env.serialize() + assert serialized["info"]["config"]["environment_type"].endswith("MiniSWESandboxEnvironment") + env.config.activate_conda = False + assert env._command("echo plain", "/tmp/work") == "echo plain" + env.config.activate_conda = True + + provider = FakeSandboxProvider.last_instance + assert provider is not None + assert provider.marker == "configured" + assert provider.created_specs[0].image == "mirror/image:tag" + assert provider.created_specs[0].env == { + "FORWARDED_KEY": "forwarded-value", + "STATIC_KEY": "static-value", + } + + result = env.execute("pytest -q", is_eval=True) + assert result == {"output": "ok", "returncode": 0, "exception_info": ""} + exec_call = provider.exec_calls[0] + assert exec_call["cwd"] == "/" + assert exec_call["timeout_s"] == 1800 + assert exec_call["user"] == "agent" + assert "conda activate testbed" in exec_call["command"] + assert exec_call["command"].endswith("pytest -q") + finally: + env.cleanup() + env.cleanup() + + assert FakeSandboxProvider.last_instance is not None + assert FakeSandboxProvider.last_instance.closed[0][1] is True + + +def test_mini_swe_sandbox_environment_validation_and_context_manager() -> None: + with pytest.raises(ValueError, match="requires provider"): + MiniSWESandboxEnvironment(image="image:tag") + + provider_name = f"fake-{uuid4().hex}" + register_provider(provider_name, FakeSandboxProvider) + with MiniSWESandboxEnvironment( + image="image:tag", + provider={provider_name: {}}, + delete=False, + ) as env: + assert env._sandbox is not None + assert FakeSandboxProvider.last_instance is not None + assert FakeSandboxProvider.last_instance.created_handles[0].sandbox_id == "fake-1" + + assert FakeSandboxProvider.last_instance is not None + assert FakeSandboxProvider.last_instance.closed[-1][1] is False + + +def test_mini_swe_sandbox_environment_submit_sentinel() -> None: + class SubmitSandboxProvider(FakeSandboxProvider): + async def exec( + self, + handle: SandboxHandle, + command: str, + *, + cwd: str | None = None, + env: dict[str, str] | None = None, + timeout_s: int | float | None = None, + user: str | int | None = None, + ) -> SandboxExecResult: + del handle, command, cwd, env, timeout_s, user + return SandboxExecResult( + stdout="COMPLETE_TASK_AND_SUBMIT_FINAL_OUTPUT\nfinal answer", + stderr=None, + return_code=0, + ) + + provider_name = f"submit-{uuid4().hex}" + register_provider(provider_name, SubmitSandboxProvider) + env = MiniSWESandboxEnvironment(image="image:tag", provider={provider_name: {}}) + + try: + with pytest.raises(Exception) as exc_info: + env.execute("submit") + assert exc_info.value.messages[0]["extra"]["submission"] == "final answer" + finally: + env.cleanup() diff --git a/uv.lock b/uv.lock index c34d4a2c0c..1707c4a0d0 100644 --- a/uv.lock +++ b/uv.lock @@ -268,6 +268,34 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/04/eb/f4151e0c7377a6e08a38108609ba5cede57986802757848688aeedd1b9e8/beautifulsoup4-4.13.5-py3-none-any.whl", hash = "sha256:642085eaa22233aceadff9c69651bc51e8bf3f874fb6d7104ece2beb24b47c4a", size = 105113, upload-time = "2025-08-24T14:06:14.884Z" }, ] +[[package]] +name = "boto3" +version = "1.43.21" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "botocore" }, + { name = "jmespath" }, + { name = "s3transfer" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/1e/f2/0ef88b6584285002a8a8000e34340f56e82681ad2b8a76ea8bd3ecaa5cb9/boto3-1.43.21.tar.gz", hash = "sha256:6dfeb70bf4f9a3514b91c7199f475f71f939199d62f9c63cd555b033fb283f89", size = 113157, upload-time = "2026-06-03T07:09:23.263Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/01/ea/5352950cbee9d1e8392e5396ddbc6defc982414e2abc004b501139bce13c/boto3-1.43.21-py3-none-any.whl", hash = "sha256:8bb863b32dabe5baa4f2e3701778c259243ba117e4811a595411819958c4fb1b", size = 140534, upload-time = "2026-06-03T07:09:21.18Z" }, +] + +[[package]] +name = "botocore" +version = "1.43.21" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "jmespath" }, + { name = "python-dateutil" }, + { name = "urllib3" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/7b/97/d9d26cebf0a8533105e183d8438931c0b196e52484cd5bf00e8443ac1b2d/botocore-1.43.21.tar.gz", hash = "sha256:17604607efe28894e947401379e569cc8f0fe2d69337ece98bd0c82d1bcfaf92", size = 15451979, upload-time = "2026-06-03T07:09:11.489Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/bb/c7/8c60049357e96d663980d66b98a44cb3626a7e5eaca66480b97826eb5379/botocore-1.43.21-py3-none-any.whl", hash = "sha256:f021ba3e844c36031106fc531ec90259ef005ba5d04f691df53bc4ecbd08f0dd", size = 15135303, upload-time = "2026-06-03T07:09:06.115Z" }, +] + [[package]] name = "cachetools" version = "5.5.2" @@ -1021,6 +1049,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b3/4a/4175a563579e884192ba6e81725fc0448b042024419be8d83aa8a80a3f44/jiter-0.10.0-cp314-cp314t-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:3aa96f2abba33dc77f79b4cf791840230375f9534e5fac927ccceb58c5e604a5", size = 354213, upload-time = "2025-05-18T19:04:41.894Z" }, ] +[[package]] +name = "jmespath" +version = "1.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d3/59/322338183ecda247fb5d1763a6cbe46eff7222eaeebafd9fa65d4bf5cb11/jmespath-1.1.0.tar.gz", hash = "sha256:472c87d80f36026ae83c6ddd0f1d05d4e510134ed462851fd5f754c8c3cbb88d", size = 27377, upload-time = "2026-01-22T16:35:26.279Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/14/2f/967ba146e6d58cf6a652da73885f52fc68001525b4197effc174321d70b4/jmespath-1.1.0-py3-none-any.whl", hash = "sha256:a5663118de4908c91729bea0acadca56526eb2698e83de10cd116ae0f4e97c64", size = 20419, upload-time = "2026-01-22T16:35:24.919Z" }, +] + [[package]] name = "jsonschema" version = "4.25.1" @@ -1414,6 +1451,13 @@ dev = [ { name = "requests-mock" }, { name = "ruff" }, ] +sandbox = [ + { name = "opensandbox" }, + { name = "tenacity" }, +] +sandbox-ecs = [ + { name = "boto3" }, +] [package.dev-dependencies] docs = [ @@ -1432,6 +1476,7 @@ docs = [ [package.metadata] requires-dist = [ { name = "aiohttp", specifier = ">=3.13.3" }, + { name = "boto3", marker = "extra == 'sandbox-ecs'", specifier = ">=1.34" }, { name = "coverage", extras = ["toml"], marker = "extra == 'dev'" }, { name = "datasets" }, { name = "devtools" }, @@ -1446,6 +1491,7 @@ requires-dist = [ { name = "mypy", marker = "extra == 'dev'", specifier = ">=1.8.0" }, { name = "omegaconf" }, { name = "openai", specifier = "<=2.7.2" }, + { name = "opensandbox", marker = "extra == 'sandbox'", specifier = ">=0.1.9" }, { name = "orjson" }, { name = "pre-commit", marker = "extra == 'dev'", specifier = ">=3.6.0" }, { name = "psutil" }, @@ -1461,6 +1507,7 @@ requires-dist = [ { name = "requests-mock", marker = "extra == 'dev'" }, { name = "rich" }, { name = "ruff", marker = "extra == 'dev'" }, + { name = "tenacity", marker = "extra == 'sandbox'", specifier = ">=9.1.4" }, { name = "tqdm" }, { name = "urllib3", specifier = ">=2.7.0" }, { name = "uvicorn" }, @@ -1468,7 +1515,7 @@ requires-dist = [ { name = "wandb" }, { name = "yappi" }, ] -provides-extras = ["dev"] +provides-extras = ["sandbox", "sandbox-ecs", "dev"] [package.metadata.requires-dev] docs = [ @@ -1623,6 +1670,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/10/68/162c97ea78c957d68ecf78a5c5041d2e25bd5562bdf5d89a6cbf7f8429bf/opencensus_context-0.1.3-py2.py3-none-any.whl", hash = "sha256:073bb0590007af276853009fac7e4bab1d523c3f03baf4cb4511ca38967c6039", size = 5060, upload-time = "2022-08-03T22:20:20.352Z" }, ] +[[package]] +name = "opensandbox" +version = "0.1.9" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "attrs" }, + { name = "httpx" }, + { name = "pydantic" }, + { name = "python-dateutil" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/5a/2a/ab3cc141e041f71a373c97fcda8749dba9328f1b9bf80401378c0611556f/opensandbox-0.1.9.tar.gz", hash = "sha256:670fbf292c498f8467963d21e91ade9ea8b8f63f4ef18d18fff9581e0952ec03", size = 160034, upload-time = "2026-05-12T12:27:20.692Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/4e/9b/553f8d7a30eddb12785711b2a1c682386878e2bb95450acd806f9fa62930/opensandbox-0.1.9-py3-none-any.whl", hash = "sha256:17faed35b60a982fee5a643fed8e4e12f041e5432d5ea0665d2828d1f2082759", size = 360945, upload-time = "2026-05-12T12:27:19.465Z" }, +] + [[package]] name = "opentelemetry-api" version = "1.36.0" @@ -2509,6 +2571,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/84/a8/001d4a7c2b37623a3fd7463208267fb906df40ff31db496157549cfd6e72/ruff-0.12.11-py3-none-win_arm64.whl", hash = "sha256:bae4d6e6a2676f8fb0f98b74594a048bae1b944aab17e9f5d504062303c6dbea", size = 12135290, upload-time = "2025-08-28T13:59:06.933Z" }, ] +[[package]] +name = "s3transfer" +version = "0.18.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "botocore" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/e0/1f/12417f7f493fc45e1f9fd5d4a9b6c125cf8d2cf3f8ddbdfab3e76406e9d6/s3transfer-0.18.0.tar.gz", hash = "sha256:3760b8b7ec1315da54048b2d626276732bee4300d054d492d4e1d43e20d4ecbd", size = 160560, upload-time = "2026-05-28T19:39:09.124Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2b/58/a58fc997655386daa2e25784e30c288aa3e3819e401f77029ee4899fb55a/s3transfer-0.18.0-py3-none-any.whl", hash = "sha256:239c13b09e65ad0346e1be7348b8a202dcad44ac7ea7c6eb858fc881dce739b6", size = 88572, upload-time = "2026-05-28T19:39:07.999Z" }, +] + [[package]] name = "sentry-sdk" version = "2.53.0" @@ -2777,6 +2851,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/51/f0/1098f6628bbe04b086ce59692d09b116ec751286eb7d33e88c5bf0c2e210/swagger_plugin_for_sphinx-6.0.0-py3-none-any.whl", hash = "sha256:35dc646d759a44ce78aefde2fe34f54e7b8c3439d0a52541a6a8b9924a711832", size = 11253, upload-time = "2025-10-16T06:26:08.504Z" }, ] +[[package]] +name = "tenacity" +version = "9.1.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/47/c6/ee486fd809e357697ee8a44d3d69222b344920433d3b6666ccd9b374630c/tenacity-9.1.4.tar.gz", hash = "sha256:adb31d4c263f2bd041081ab33b498309a57c77f9acf2db65aadf0898179cf93a", size = 49413, upload-time = "2026-02-07T10:45:33.841Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/d7/c1/eb8f9debc45d3b7918a32ab756658a0904732f75e555402972246b0b8e71/tenacity-9.1.4-py3-none-any.whl", hash = "sha256:6095a360c919085f28c6527de529e76a06ad89b23659fa881ae0649b867a9d55", size = 28926, upload-time = "2026-02-07T10:45:32.24Z" }, +] + [[package]] name = "tqdm" version = "4.67.1"