diff --git a/.github/scripts/run_integration_tests.py b/.github/scripts/run_integration_tests.py index 9c032e5ccf..08f4d16564 100644 --- a/.github/scripts/run_integration_tests.py +++ b/.github/scripts/run_integration_tests.py @@ -602,6 +602,10 @@ def main() -> None: additional_env["OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_EXTRA"] = ( installation.extra ) + if installation.distribution is not None: + additional_env[ + "OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_DISTRIBUTION" + ] = installation.distribution run_suite( python, wheel, diff --git a/examples/sandbox/README.md b/examples/sandbox/README.md index 2248de310e..b40ae6b144 100644 --- a/examples/sandbox/README.md +++ b/examples/sandbox/README.md @@ -25,7 +25,7 @@ Most examples call a model through `Runner`, so set `OPENAI_API_KEY` in the repo ## Cloud backend examples -Cloud-provider examples live under [`extensions/`](./extensions/). They cover E2B, Modal, and Daytona sandbox backends and require provider-specific credentials in addition to `OPENAI_API_KEY`. +Cloud-provider examples live under [`extensions/`](./extensions/). They cover CreateOS, E2B, Modal, Daytona, Cloudflare, Runloop, Blaxel, and Vercel sandbox backends and require provider-specific credentials in addition to `OPENAI_API_KEY`. ## Tutorial scaffold diff --git a/examples/sandbox/extensions/README.md b/examples/sandbox/extensions/README.md index 3037543df4..acac1236d5 100644 --- a/examples/sandbox/extensions/README.md +++ b/examples/sandbox/extensions/README.md @@ -6,10 +6,28 @@ They intentionally keep the flow simple: 1. Build a tiny manifest in memory. 2. Create a `SandboxAgent` that inspects that workspace through one shell tool. -3. Run the agent against E2B, Modal, Daytona, Cloudflare, Runloop, Blaxel, or Vercel. +3. Run the agent against CreateOS, E2B, Modal, Daytona, Cloudflare, Runloop, Blaxel, or Vercel. All of these examples require `OPENAI_API_KEY`, because they call the model through the normal `Runner` path. Each cloud backend also needs its own provider credentials. +## CreateOS + +Install the CreateOS extra and configure the provider API key: + +```bash +uv sync --extra createos +export CREATEOS_API_KEY=... +export OPENAI_API_KEY=... +``` + +Run the minimal agent example: + +```bash +uv run python examples/sandbox/extensions/createos_runner.py --stream +``` + +The example defaults to the `s-4vcpu-4gb` shape and `devbox:1` root filesystem. Override them with `--shape` and `--rootfs` when your CreateOS environment uses different catalog entries. Add `--pause-on-exit` to preserve the sandbox for a later resumed run; otherwise the runner destroys it during cleanup. + ## E2B ### Setup diff --git a/examples/sandbox/extensions/createos_runner.py b/examples/sandbox/extensions/createos_runner.py new file mode 100644 index 0000000000..8278116a29 --- /dev/null +++ b/examples/sandbox/extensions/createos_runner.py @@ -0,0 +1,128 @@ +"""Minimal CreateOS-backed sandbox example for manual validation.""" + +import argparse +import asyncio +import os +import sys +from pathlib import Path + +from openai.types.responses import ResponseTextDeltaEvent + +from agents import ModelSettings, Runner +from agents.run import RunConfig +from agents.sandbox import Manifest, SandboxAgent, SandboxRunConfig + +if __package__ is None or __package__ == "": + sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +from examples.sandbox.misc.example_support import text_manifest +from examples.sandbox.misc.workspace_shell import WorkspaceShellCapability + +try: + from agents.extensions.sandbox import ( + DEFAULT_CREATEOS_WORKSPACE_ROOT, + CreateOSSandboxClient, + CreateOSSandboxClientOptions, + ) +except Exception as exc: # pragma: no cover - import path depends on optional extras + raise SystemExit( + "CreateOS sandbox examples require the optional repo extra.\n" + "Install it with: uv sync --extra createos" + ) from exc + + +DEFAULT_QUESTION = "Summarize this cloud sandbox workspace in 2 sentences." + + +def _manifest() -> Manifest: + manifest = text_manifest( + { + "README.md": ( + "# CreateOS Demo Workspace\n\n" + "This workspace validates the CreateOS sandbox backend for the Agents SDK.\n" + ), + "status.md": ( + "# Status\n\n" + "- Sandbox creation is configured.\n" + "- Command execution and file transfer are ready for validation.\n" + ), + } + ) + return manifest.model_copy(update={"root": DEFAULT_CREATEOS_WORKSPACE_ROOT}) + + +def _require_env(name: str) -> str: + value = os.environ.get(name) + if not value: + raise SystemExit(f"{name} must be set before running this example.") + return value + + +async def main( + *, + model: str, + question: str, + shape: str, + rootfs: str | None, + pause_on_exit: bool, + stream: bool, +) -> None: + _require_env("OPENAI_API_KEY") + api_key = _require_env("CREATEOS_API_KEY") + + agent = SandboxAgent( + name="CreateOS Sandbox Assistant", + model=model, + instructions=( + "Inspect the sandbox workspace before answering. Keep the answer concise and cite " + "the file names you inspected." + ), + default_manifest=_manifest(), + capabilities=[WorkspaceShellCapability()], + model_settings=ModelSettings(tool_choice="required"), + ) + client = CreateOSSandboxClient(api_key=api_key) + run_config = RunConfig( + sandbox=SandboxRunConfig( + client=client, + options=CreateOSSandboxClientOptions( + shape=shape, + rootfs=rootfs, + pause_on_exit=pause_on_exit, + ), + ), + workflow_name="CreateOS sandbox example", + ) + + try: + if not stream: + result = await Runner.run(agent, question, run_config=run_config) + print(result.final_output) + return + + stream_result = Runner.run_streamed(agent, question, run_config=run_config) + saw_text_delta = False + async for event in stream_result.stream_events(): + if event.type == "raw_response_event" and isinstance( + event.data, ResponseTextDeltaEvent + ): + if not saw_text_delta: + print("assistant> ", end="", flush=True) + saw_text_delta = True + print(event.data.delta, end="", flush=True) + if saw_text_delta: + print() + finally: + await client.close() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--model", default="gpt-5.6-sol", help="Model name to use.") + parser.add_argument("--question", default=DEFAULT_QUESTION, help="Prompt to send.") + parser.add_argument("--shape", default="s-4vcpu-4gb", help="CreateOS sandbox shape.") + parser.add_argument("--rootfs", default="devbox:1", help="CreateOS root filesystem.") + parser.add_argument("--pause-on-exit", action="store_true") + parser.add_argument("--stream", action="store_true") + args = parser.parse_args() + asyncio.run(main(**vars(args))) diff --git a/integration_tests/_contract_surface.py b/integration_tests/_contract_surface.py index 26067ac728..3961d41dd6 100644 --- a/integration_tests/_contract_surface.py +++ b/integration_tests/_contract_surface.py @@ -36,6 +36,7 @@ class OptionalDependencyInstallation: extra: str | None = None requirement: str | None = None unsupported_platforms: tuple[str, ...] = () + distribution: str | None = None def is_supported_on_current_platform(self) -> bool: return sys.platform not in self.unsupported_platforms @@ -113,7 +114,7 @@ def load_submodule_export_policy(path: Path) -> SubmoduleExportPolicy: f"optional dependency installation for {module_name} must be an object" ) unknown_fields = sorted( - set(installation) - {"extra", "requirement", "unsupported_platforms"} + set(installation) - {"extra", "requirement", "unsupported_platforms", "distribution"} ) if unknown_fields: raise ValueError( @@ -133,6 +134,14 @@ def load_submodule_export_policy(path: Path) -> SubmoduleExportPolicy: f"optional dependency installation {field_name} for {module_name} must be a " "non-empty string" ) + distribution = installation.get("distribution") + if distribution is not None and ( + field_name != "extra" or type(distribution) is not str or not distribution + ): + raise ValueError( + f"optional dependency installation distribution for {module_name} must be a " + "non-empty string declared with an extra" + ) unsupported_platforms = installation.get("unsupported_platforms", []) if ( not isinstance(unsupported_platforms, list) @@ -149,6 +158,7 @@ def load_submodule_export_policy(path: Path) -> SubmoduleExportPolicy: extra=install_value if field_name == "extra" else None, requirement=install_value if field_name == "requirement" else None, unsupported_platforms=tuple(unsupported_platforms), + distribution=distribution, ) ) diff --git a/integration_tests/packaging/test_released_api_contract.py b/integration_tests/packaging/test_released_api_contract.py index 89300dd70e..56f0473d97 100644 --- a/integration_tests/packaging/test_released_api_contract.py +++ b/integration_tests/packaging/test_released_api_contract.py @@ -17,6 +17,7 @@ REQUIRED_OPTIONAL_DEPENDENCIES_ENV = "OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_DEPENDENCIES" OPTIONAL_DEPENDENCY_INSTALLATION_ENV = "OPENAI_AGENTS_INTEGRATION_OPTIONAL_DEPENDENCY_INSTALLATION" REQUIRED_OPTIONAL_EXTRA_ENV = "OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_EXTRA" +REQUIRED_OPTIONAL_DISTRIBUTION_ENV = "OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_DISTRIBUTION" def _distributions_declared_by_extra(requirement_strings: list[str], extra: str) -> set[str]: @@ -38,6 +39,7 @@ def _extra_metadata_error( *, extra: str, dependency_module: str, + distribution: str | None = None, provided_extras: list[str], requirement_strings: list[str], ) -> str | None: @@ -51,7 +53,7 @@ def _extra_metadata_error( "the extra under [project.optional-dependencies]." ) - distribution_name = canonicalize_name(dependency_module) + distribution_name = canonicalize_name(distribution or dependency_module) declared_distributions = _distributions_declared_by_extra(requirement_strings, extra) if distribution_name not in declared_distributions: return ( @@ -126,6 +128,7 @@ def test_artifact_extra_declares_its_policy_dependency() -> None: error = _extra_metadata_error( extra=extra, dependency_module=dependency_module, + distribution=os.environ.get(REQUIRED_OPTIONAL_DISTRIBUTION_ENV), provided_extras=metadata("openai-agents").get_all("Provides-Extra") or [], requirement_strings=requires("openai-agents") or [], ) @@ -173,6 +176,31 @@ def test_extra_metadata_provenance_rejects_unknown_extra() -> None: ) +def test_extra_metadata_provenance_uses_declared_distribution() -> None: + requirement_strings = ['createos-sandbox>=0.1.0,<0.2; extra == "createos"'] + + assert ( + _extra_metadata_error( + extra="createos", + dependency_module="createos", + distribution="createos-sandbox", + provided_extras=["createos"], + requirement_strings=requirement_strings, + ) + is None + ) + assert ( + _extra_metadata_error( + extra="createos", + dependency_module="createos", + distribution="createos-sandbox", + provided_extras=["createos"], + requirement_strings=['other-package>=1; extra == "createos"'], + ) + is not None + ) + + @pytest.mark.packaging_dependency def test_installed_distribution_preserves_released_public_api_contract() -> None: contract = load_api_contract(CONTRACT) diff --git a/pyproject.toml b/pyproject.toml index 9afb1ae6d5..77fe348129 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -60,6 +60,7 @@ dapr = [ ] mongodb = ["pymongo>=4.14"] docker = ["docker>=6.1"] +createos = ["createos-sandbox>=0.1.0,<0.2"] blaxel = ["blaxel>=0.2.50", "aiohttp>=3.14.3,<4"] daytona = [ "daytona>=0.155.0", @@ -195,6 +196,10 @@ disallow_untyped_calls = false module = "sounddevice.*" ignore_missing_imports = true +[[tool.mypy.overrides]] +module = ["createos", "createos.*"] +ignore_missing_imports = true + [[tool.mypy.overrides]] module = ["modal", "modal.*"] ignore_missing_imports = true @@ -268,10 +273,10 @@ format-command = "ruff format --stdin-filename {filename}" [tool.uv] exclude-newer = "7 days" -exclude-newer-package = { openai = false } +exclude-newer-package = { createos-sandbox = false, openai = false } index-strategy = "first-index" [tool.uv.pip] exclude-newer = "7 days" -exclude-newer-package = { openai = false } +exclude-newer-package = { createos-sandbox = false, openai = false } index-strategy = "first-index" diff --git a/src/agents/extensions/sandbox/__init__.py b/src/agents/extensions/sandbox/__init__.py index 9ea2eb8dc3..6600341235 100644 --- a/src/agents/extensions/sandbox/__init__.py +++ b/src/agents/extensions/sandbox/__init__.py @@ -1,3 +1,21 @@ +from importlib.util import find_spec + +try: + if find_spec("createos") is None: + raise ImportError("The optional CreateOS dependency is not installed") + from .createos import ( + DEFAULT_CREATEOS_WORKSPACE_ROOT as DEFAULT_CREATEOS_WORKSPACE_ROOT, + CreateOSSandboxClient as CreateOSSandboxClient, + CreateOSSandboxClientOptions as CreateOSSandboxClientOptions, + CreateOSSandboxSession as CreateOSSandboxSession, + CreateOSSandboxSessionState as CreateOSSandboxSessionState, + CreateOSSandboxTimeouts as CreateOSSandboxTimeouts, + ) + + _HAS_CREATEOS = True +except Exception: # pragma: no cover + _HAS_CREATEOS = False + try: from .e2b import ( E2BCloudBucketMountStrategy as E2BCloudBucketMountStrategy, @@ -113,6 +131,18 @@ __all__: list[str] = [] +if _HAS_CREATEOS: + __all__.extend( + [ + "DEFAULT_CREATEOS_WORKSPACE_ROOT", + "CreateOSSandboxClient", + "CreateOSSandboxClientOptions", + "CreateOSSandboxSession", + "CreateOSSandboxSessionState", + "CreateOSSandboxTimeouts", + ] + ) + if _HAS_E2B: __all__.extend( [ diff --git a/src/agents/extensions/sandbox/createos/__init__.py b/src/agents/extensions/sandbox/createos/__init__.py new file mode 100644 index 0000000000..39bc993efa --- /dev/null +++ b/src/agents/extensions/sandbox/createos/__init__.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +from .sandbox import ( + DEFAULT_CREATEOS_WORKSPACE_ROOT, + CreateOSSandboxClient, + CreateOSSandboxClientOptions, + CreateOSSandboxSession, + CreateOSSandboxSessionState, + CreateOSSandboxTimeouts, +) + +__all__ = [ + "DEFAULT_CREATEOS_WORKSPACE_ROOT", + "CreateOSSandboxClient", + "CreateOSSandboxClientOptions", + "CreateOSSandboxSession", + "CreateOSSandboxSessionState", + "CreateOSSandboxTimeouts", +] diff --git a/src/agents/extensions/sandbox/createos/sandbox.py b/src/agents/extensions/sandbox/createos/sandbox.py new file mode 100644 index 0000000000..364592485b --- /dev/null +++ b/src/agents/extensions/sandbox/createos/sandbox.py @@ -0,0 +1,1578 @@ +"""CreateOS sandbox implementation. + +This module adapts the synchronous ``createos`` Python SDK to the asynchronous +Agents SDK sandbox interfaces. The dependency is optional and imported lazily. +""" + +from __future__ import annotations + +import asyncio +import io +import logging +import shlex +import time +import uuid +from collections import deque +from collections.abc import Awaitable, Callable, MutableMapping +from contextvars import ContextVar +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Literal, NoReturn, cast +from urllib.parse import urlsplit + +import httpx +from pydantic import BaseModel, Field, PrivateAttr, field_validator + +from ....logger import log_tool_action_debug +from ....sandbox._mount_security import redact_mount_error_data +from ....sandbox.errors import ( + ErrorCode, + ExecTimeoutError, + ExecTransportError, + ExposedPortUnavailableError, + SandboxRuntimeError, + WorkspaceArchiveReadError, + WorkspaceArchiveWriteError, + WorkspaceReadNotFoundError, + WorkspaceWriteTypeError, +) +from ....sandbox.manifest import Manifest +from ....sandbox.session import SandboxSession, SandboxSessionState +from ....sandbox.session.base_sandbox_session import BaseSandboxSession +from ....sandbox.session.dependencies import Dependencies +from ....sandbox.session.manager import Instrumentation +from ....sandbox.session.mount_lifecycle import with_ephemeral_mounts_removed +from ....sandbox.session.pty_output import collect_pty_output +from ....sandbox.session.pty_types import ( + PTY_PROCESSES_MAX, + PTY_PROCESSES_WARNING, + PtyExecUpdate, + _settle_pty_cleanup, + allocate_pty_process_id, + clamp_pty_yield_time_ms, + process_id_to_prune_from_meta, + resolve_pty_write_yield_time_ms, +) +from ....sandbox.session.runtime_helpers import RESOLVE_WORKSPACE_PATH_HELPER, RuntimeHelperScript +from ....sandbox.session.sandbox_client import BaseSandboxClient, BaseSandboxClientOptions +from ....sandbox.session.tar_workspace import shell_tar_exclude_args +from ....sandbox.snapshot import SnapshotBase, SnapshotSpec, resolve_snapshot +from ....sandbox.types import ExecResult, ExposedPortEndpoint, User +from ....sandbox.util.retry import TRANSIENT_HTTP_STATUS_CODES, iter_exception_chain, retry_async +from ....sandbox.util.tar_utils import UnsafeTarMemberError, validate_tar_bytes +from ....sandbox.workspace_paths import coerce_posix_path, posix_path_as_path, sandbox_path_str + +DEFAULT_CREATEOS_WORKSPACE_ROOT = "/workspace" +logger = logging.getLogger(__name__) +_TERMINAL_CREATEOS_STATUSES = frozenset({"destroying", "destroyed", "failed", "error"}) +_PTY_STREAM_READ_TIMEOUT_S = 5.0 +_lifecycle_owner: ContextVar[asyncio.Task[None] | None] = ContextVar( + "createos_lifecycle_owner", default=None +) + + +def _import_createos_sdk() -> Any: + try: + import createos + + return createos + except ImportError as e: + raise ImportError( + "CreateOSSandboxClient requires the optional `createos-sandbox` dependency.\n" + "Install the CreateOS extra before using this sandbox backend." + ) from e + + +def _provider_status_code(error: BaseException) -> int | None: + for candidate in iter_exception_chain(error): + status = getattr(candidate, "status_code", None) + if isinstance(status, int): + return status + return None + + +def _provider_retryability(error: BaseException) -> bool | None: + status = _provider_status_code(error) + if status is None: + return True if isinstance(error, ConnectionError | TimeoutError) else None + if status in TRANSIENT_HTTP_STATUS_CODES: + return True + if 400 <= status < 500: + return False + return None + + +def _require_rootfs(rootfs: str | None) -> str: + if rootfs is None or not rootfs.strip(): + raise ValueError("CreateOS sandbox creation requires a non-empty rootfs") + return rootfs + + +def _set_pty_stream_read_timeout(stream: Any, timeout: float) -> None: + """Bound an idle stream read so process cleanup can settle its reader.""" + response = getattr(stream, "_response", None) + request = getattr(response, "request", None) + extensions = getattr(request, "extensions", None) + timeout_config = extensions.get("timeout") if isinstance(extensions, MutableMapping) else None + if not isinstance(timeout_config, MutableMapping): + raise RuntimeError( + "CreateOS managed process stream does not expose configurable HTTP timeouts" + ) + timeout_config["read"] = timeout + + +async def _to_thread_settled( + function: Any, + /, + *args: Any, + timeout: float | None = None, + interrupted_result_cleanup: Callable[[Any], Awaitable[None]] | None = None, + on_success: Callable[[Any], None] | None = None, +) -> Any: + """Keep a synchronous provider call owned until its worker thread exits.""" + task = asyncio.create_task(asyncio.to_thread(function, *args)) + caller_cancelled = False + timed_out = False + try: + if timeout is None: + await asyncio.shield(task) + else: + done, _ = await asyncio.wait({task}, timeout=timeout) + timed_out = not done + except asyncio.CancelledError: + caller_cancelled = True + + if caller_cancelled or timed_out: + while not task.done(): + try: + await asyncio.shield(task) + except asyncio.CancelledError: + caller_cancelled = True + try: + result = task.result() + except BaseException: + if timed_out: + raise TimeoutError from None + if caller_cancelled: + raise asyncio.CancelledError from None + raise + if on_success is not None: + on_success(result) + if interrupted_result_cleanup is not None: + cleanup_task: asyncio.Future[None] = asyncio.ensure_future( + interrupted_result_cleanup(result) + ) + while not cleanup_task.done(): + try: + await asyncio.shield(cleanup_task) + except asyncio.CancelledError: + caller_cancelled = True + cleanup_task.result() + if timed_out: + raise TimeoutError from None + raise asyncio.CancelledError from None + + result = task.result() + if on_success is not None: + on_success(result) + return result + + +def _raise_exec_error( + error: BaseException, + *, + command: tuple[str | Path, ...], + timeout: float, +) -> NoReturn: + sdk = _import_createos_sdk() + context: dict[str, object] = {"backend": "createos"} + status = _provider_status_code(error) + if status is not None: + context["http_status"] = status + detail = str(error).strip() + if detail: + context["provider_error"] = detail + + if isinstance(error, sdk.OperationTimeout | TimeoutError) or status in { + 408, + 504, + }: + raise ExecTimeoutError( + command=command, + timeout_s=timeout, + context=context, + cause=error, + ) from error + raise ExecTransportError( + command=command, + context=context, + cause=error, + retryable=_provider_retryability(error), + ) from error + + +class CreateOSSandboxTimeouts(BaseModel): + """Timeout configuration for CreateOS sandbox operations.""" + + model_config = {"frozen": True} + + exec_timeout_unbounded_s: float = Field(default=24 * 60 * 60, ge=1) + create_s: float = Field(default=120, ge=1) + lifecycle_s: float = Field(default=120, ge=1) + fast_op_s: float = Field(default=30, ge=1) + file_upload_s: float = Field(default=1800, ge=1) + file_download_s: float = Field(default=1800, ge=1) + workspace_tar_s: float = Field(default=300, ge=1) + cleanup_s: float = Field(default=30, ge=1) + + +class CreateOSSandboxClientOptions(BaseSandboxClientOptions): + """Creation and lifecycle settings for a CreateOS sandbox.""" + + type: Literal["createos"] = "createos" + shape: str = Field(min_length=1) + rootfs: str | None = None + env_vars: dict[str, str] | None = None + pause_on_exit: bool = False + name: str | None = None + exposed_ports: tuple[int, ...] = () + timeouts: CreateOSSandboxTimeouts | dict[str, object] | None = None + network_ids: tuple[str, ...] = () + disk_mib: int | None = Field(default=None, gt=0) + egress_rules: tuple[str, ...] = () + ssh_public_keys: tuple[str, ...] = () + host_id: str | None = None + node_selector: dict[str, str] | None = None + region: str | None = None + auto_pause_after_seconds: int | None = Field(default=None, gt=0) + + @field_validator("shape") + @classmethod + def _validate_shape(cls, value: str) -> str: + if not value.strip(): + raise ValueError("shape must not be blank") + return value + + @field_validator("exposed_ports") + @classmethod + def _validate_exposed_ports(cls, value: tuple[int, ...]) -> tuple[int, ...]: + if any(port < 1 or port > 65_535 for port in value): + raise ValueError("exposed_ports must contain valid TCP ports") + return value + + def __init__( + self, + shape: str, + rootfs: str | None = None, + env_vars: dict[str, str] | None = None, + pause_on_exit: bool = False, + name: str | None = None, + exposed_ports: tuple[int, ...] = (), + timeouts: CreateOSSandboxTimeouts | dict[str, object] | None = None, + network_ids: tuple[str, ...] = (), + disk_mib: int | None = None, + egress_rules: tuple[str, ...] = (), + ssh_public_keys: tuple[str, ...] = (), + host_id: str | None = None, + node_selector: dict[str, str] | None = None, + region: str | None = None, + auto_pause_after_seconds: int | None = None, + *, + type: Literal["createos"] = "createos", + ) -> None: + super().__init__( + type=type, + shape=shape, + rootfs=rootfs, + env_vars=env_vars, + pause_on_exit=pause_on_exit, + name=name, + exposed_ports=exposed_ports, + timeouts=timeouts, + network_ids=network_ids, + disk_mib=disk_mib, + egress_rules=egress_rules, + ssh_public_keys=ssh_public_keys, + host_id=host_id, + node_selector=node_selector, + region=region, + auto_pause_after_seconds=auto_pause_after_seconds, + ) + + +class CreateOSSandboxSessionState(SandboxSessionState): + """Serializable state for a CreateOS-backed sandbox session.""" + + type: Literal["createos"] = "createos" + sandbox_id: str + shape: str = Field(min_length=1) + rootfs: str | None = None + _base_env_vars: dict[str, str] = PrivateAttr(default_factory=dict) + pause_on_exit: bool = False + name: str | None = None + timeouts: CreateOSSandboxTimeouts = Field(default_factory=CreateOSSandboxTimeouts) + network_ids: tuple[str, ...] = () + disk_mib: int | None = Field(default=None, gt=0) + egress_rules: tuple[str, ...] = () + ssh_public_keys: tuple[str, ...] = () + host_id: str | None = None + node_selector: dict[str, str] | None = None + region: str | None = None + auto_pause_after_seconds: int | None = Field(default=None, gt=0) + + @field_validator("shape") + @classmethod + def _validate_shape(cls, value: str) -> str: + if not value.strip(): + raise ValueError("shape must not be blank") + return value + + @field_validator("exposed_ports") + @classmethod + def _validate_exposed_ports(cls, value: tuple[int, ...]) -> tuple[int, ...]: + if any(port < 1 or port > 65_535 for port in value): + raise ValueError("exposed_ports must contain valid TCP ports") + return value + + +@dataclass +class _CreateOSPtySessionEntry: + provider_process_id: str + termination_lock: asyncio.Lock = field(default_factory=asyncio.Lock) + output_chunks: deque[bytes] = field(default_factory=deque) + output_lock: asyncio.Lock = field(default_factory=asyncio.Lock) + output_notify: asyncio.Event = field(default_factory=asyncio.Event) + output_closed: asyncio.Event = field(default_factory=asyncio.Event) + last_used: float = field(default_factory=time.monotonic) + exit_code: int | None = None + reader_task: asyncio.Task[None] | None = None + stream: Any | None = None + termination_started: bool = False + provider_deleted: bool = False + + def mark_provider_deleted(self, _result: Any) -> None: + self.provider_deleted = True + + +class CreateOSSandboxSession(BaseSandboxSession): + """CreateOS-backed implementation of the provider-neutral sandbox session.""" + + state: CreateOSSandboxSessionState + + def __init__(self, *, state: CreateOSSandboxSessionState, sandbox: Any) -> None: + self.state = state + self._sandbox = sandbox + self._mount_transition_terminal = False + self._pause_issued = False + self._pause_completed = False + self._destroy_issued = False + self._destroy_completed = False + self._lifecycle_lock = asyncio.Lock() + self._active_lifecycle_owner: asyncio.Task[None] | None = None + self._pty_lock = asyncio.Lock() + self._pty_sessions: dict[int, _CreateOSPtySessionEntry] = {} + self._reserved_pty_process_ids: set[int] = set() + self._pty_admission_generation = 0 + self._pty_admission_open = True + + @classmethod + def from_state( + cls, + state: CreateOSSandboxSessionState, + *, + sandbox: Any, + ) -> CreateOSSandboxSession: + return cls(state=state, sandbox=sandbox) + + @property + def sandbox_id(self) -> str: + return self.state.sandbox_id + + def _runtime_helpers(self) -> tuple[RuntimeHelperScript, ...]: + return (RESOLVE_WORKSPACE_PATH_HELPER,) + + def _assert_session_usable(self) -> None: + if self._mount_transition_terminal: + raise SandboxRuntimeError( + message="sandbox session is unavailable after an ambiguous mount transition", + error_code=ErrorCode.MOUNT_FAILED, + op="shutdown", + context={"backend": "createos"}, + retryable=False, + ) + + def _mark_pause_issued(self, _result: Any) -> None: + self._pause_issued = True + + def _mark_destroy_issued(self, _result: Any) -> None: + self._destroy_issued = True + + def _mark_resume_completed(self, _result: Any) -> None: + # A successful resume starts a new pause lifecycle even if waiting for + # the running state subsequently fails. + self._pause_issued = False + self._pause_completed = False + + async def start(self) -> None: + admission_generation = self._pty_admission_generation + async with self._lifecycle_lock: + owner = asyncio.current_task() + assert owner is not None + self._active_lifecycle_owner = owner + token = _lifecycle_owner.set(owner) + try: + await _settle_pty_cleanup(super().start()) + if admission_generation == self._pty_admission_generation: + self._pty_admission_open = True + finally: + self._active_lifecycle_owner = None + _lifecycle_owner.reset(token) + + async def stop(self) -> None: + async with self._lifecycle_lock: + owner = asyncio.current_task() + assert owner is not None + self._active_lifecycle_owner = owner + token = _lifecycle_owner.set(owner) + try: + await _settle_pty_cleanup(super().stop()) + finally: + self._active_lifecycle_owner = None + _lifecycle_owner.reset(token) + + async def _before_stop(self) -> None: + # stop() already holds the lifecycle lock through snapshot persistence. + await self._pty_terminate_all_locked() + + async def _ensure_backend_started(self) -> None: + await self._ensure_backend_started_locked() + + async def _ensure_backend_started_locked(self) -> None: + self._assert_session_usable() + if self._destroy_issued or self._destroy_completed: + raise SandboxRuntimeError( + message="CreateOS sandbox session cannot restart after destruction", + error_code=ErrorCode.WORKSPACE_START_ERROR, + op="start", + context={"backend": "createos", "sandbox_id": self.state.sandbox_id}, + retryable=False, + ) + sdk = _import_createos_sdk() + status = str(self._sandbox.status) + resumed = False + if status == "pausing": + await _to_thread_settled( + self._sandbox.wait_until_paused, + sdk.WaitOptions(timeout=self.state.timeouts.lifecycle_s), + timeout=self.state.timeouts.lifecycle_s + 1, + ) + status = str(self._sandbox.status) + if status == "paused": + await _to_thread_settled( + self._sandbox.resume, + timeout=self.state.timeouts.lifecycle_s, + on_success=self._mark_resume_completed, + ) + resumed = True + status = str(self._sandbox.status) + if resumed or status != "running": + await _to_thread_settled( + self._sandbox.wait_until_running, + sdk.WaitOptions(timeout=self.state.timeouts.lifecycle_s), + timeout=self.state.timeouts.lifecycle_s + 1, + ) + + def _close_pty_admission(self) -> None: + self._pty_admission_open = False + self._pty_admission_generation += 1 + + def _capture_pty_admission_generation(self) -> int: + if not self._pty_admission_open: + raise SandboxRuntimeError( + message="CreateOS sandbox session is stopping", + error_code=ErrorCode.EXEC_TRANSPORT_ERROR, + op="exec", + context={"backend": "createos", "sandbox_id": self.state.sandbox_id}, + retryable=False, + ) + return self._pty_admission_generation + + def _assert_pty_admission_open(self, generation: int) -> None: + if not self._pty_admission_open or generation != self._pty_admission_generation: + raise SandboxRuntimeError( + message="CreateOS sandbox session is stopping", + error_code=ErrorCode.EXEC_TRANSPORT_ERROR, + op="exec", + context={"backend": "createos", "sandbox_id": self.state.sandbox_id}, + retryable=False, + ) + + async def _before_shutdown(self) -> None: + # shutdown() already holds the lifecycle lock. + await self._pty_terminate_all_locked() + + async def shutdown(self) -> None: + self._close_pty_admission() + async with self._lifecycle_lock: + await _settle_pty_cleanup(super().shutdown()) + + async def _cleanup_provider_process(self, process: Any) -> None: + process_id = str(getattr(process, "process_id", "")) + if not process_id: + return + entry = _CreateOSPtySessionEntry(provider_process_id=process_id) + entry.output_closed.set() + try: + await _to_thread_settled( + self._sandbox.processes.delete, + process_id, + timeout=self.state.timeouts.cleanup_s, + on_success=entry.mark_provider_deleted, + ) + except BaseException: + async with self._pty_lock: + local_process_id = allocate_pty_process_id(self._reserved_pty_process_ids) + self._reserved_pty_process_ids.add(local_process_id) + self._pty_sessions[local_process_id] = entry + raise + + async def _prepare_backend_workspace(self) -> None: + root = sandbox_path_str(self.state.manifest.root) + result = await self._exec_at_cwd( + ("mkdir", "-p", "--", root), + cwd="/", + timeout=self.state.timeouts.fast_op_s, + ) + if not result.ok(): + raise WorkspaceArchiveWriteError( + path=self._workspace_root_path(), + context={ + "reason": "workspace_root_create_failed", + "stderr": result.stderr.decode("utf-8", errors="replace"), + }, + ) + + async def _validate_path_access( + self, + path: Path | str, + *, + for_write: bool = False, + ) -> Path: + return await self._validate_remote_path_access(path, for_write=for_write) + + async def _resolved_envs(self) -> dict[str, str]: + manifest_envs = await self.state.manifest.environment.resolve() + return {**self.state._base_env_vars, **manifest_envs} + + def _coerce_exec_timeout(self, timeout: float | None) -> float: + if timeout is None: + return self.state.timeouts.exec_timeout_unbounded_s + return max(float(timeout), 0.001) + + async def _exec_internal( + self, + *command: str | Path, + timeout: float | None = None, + ) -> ExecResult: + return await self._exec_at_cwd( + tuple(str(part) for part in command), + cwd=sandbox_path_str(self.state.manifest.root), + timeout=timeout, + ) + + async def _exec_at_cwd( + self, + command: tuple[str, ...], + *, + cwd: str, + timeout: float | None, + ) -> ExecResult: + self._assert_session_usable() + sdk = _import_createos_sdk() + effective_timeout = self._coerce_exec_timeout(timeout) + request = sdk.RunCommandRequest( + command="sh", + arguments=[ + "-lc", + 'cd "$1" && shift && exec "$@"', + "sh", + cwd, + *command, + ], + ) + options = sdk.ExecOptions( + timeout=effective_timeout, + environment_variables=await self._resolved_envs() or None, + ) + try: + response = await _to_thread_settled( + self._sandbox.run_command, + request, + options, + timeout=effective_timeout + 1, + ) + except Exception as e: + _raise_exec_error( + e, + command=command, + timeout=effective_timeout, + ) + + result = response.result + return ExecResult( + stdout=str(result.standard_output or "").encode("utf-8", errors="replace"), + stderr=str(result.standard_error or result.error_message or "").encode( + "utf-8", errors="replace" + ), + exit_code=int(result.exit_code or 0), + ) + + def supports_pty(self) -> bool: + processes = getattr(self._sandbox, "processes", None) + return all( + callable(getattr(processes, name, None)) + for name in ("create", "connect", "input", "delete") + ) + + async def pty_exec_start( + self, + *command: str | Path, + timeout: float | None = None, + shell: bool | list[str] = True, + user: str | User | None = None, + tty: bool = False, + yield_time_s: float | None = None, + max_output_tokens: int | None = None, + ) -> PtyExecUpdate: + if user is not None: + raise NotImplementedError("CreateOS PTY execution does not support `user`") + if not self.supports_pty(): + raise NotImplementedError("PTY execution is not supported by this CreateOS SDK") + + generation = self._capture_pty_admission_generation() + sanitized = self._prepare_exec_command(*command, shell=shell, user=None) + command_text = shlex.join(str(part) for part in sanitized) + effective_timeout = self._coerce_exec_timeout(timeout) + sdk = _import_createos_sdk() + process: Any | None = None + entry: _CreateOSPtySessionEntry | None = None + registered = False + pruned: tuple[int, _CreateOSPtySessionEntry] | None = None + process_id = 0 + process_count = 0 + + async with self._lifecycle_lock: + try: + self._assert_pty_admission_open(generation) + request = sdk.ManagedProcessCreateRequest( + command="/bin/sh", + working_directory=sandbox_path_str(self.state.manifest.root), + environment_variables=await self._resolved_envs(), + pty=sdk.PTYSize(rows=24, cols=80) if tty else None, + ) + process = await _to_thread_settled( + self._sandbox.processes.create, + request, + timeout=effective_timeout, + interrupted_result_cleanup=self._cleanup_provider_process, + ) + provider_process_id = str(getattr(process, "process_id", "")) + if not provider_process_id: + raise RuntimeError("CreateOS managed process creation returned no process id") + entry = _CreateOSPtySessionEntry(provider_process_id=provider_process_id) + await _to_thread_settled( + self._sandbox.processes.input, + provider_process_id, + f"{command_text}\n", + timeout=self.state.timeouts.fast_op_s, + ) + + async with self._pty_lock: + process_id = allocate_pty_process_id(self._reserved_pty_process_ids) + self._reserved_pty_process_ids.add(process_id) + pruned = self._prune_pty_sessions_if_needed() + self._pty_sessions[process_id] = entry + process_count = len(self._pty_sessions) + registered = True + entry.reader_task = asyncio.create_task(self._run_pty_reader(entry)) + except BaseException: + if process is not None and not registered: + await _settle_pty_cleanup(self._cleanup_provider_process(process)) + raise + + assert entry is not None + if pruned is not None: + pruned_process_id, pruned_entry = pruned + await self._terminate_pty_entry(pruned_entry) + await self._remove_pty_entry(pruned_process_id, pruned_entry) + if process_count >= PTY_PROCESSES_WARNING: + logger.warning( + "PTY process count reached warning threshold: %s active sessions", + process_count, + ) + + yield_time_ms = 10_000 if yield_time_s is None else int(yield_time_s * 1000) + output, original_token_count, output_closed = await self._collect_pty_output( + entry=entry, + yield_time_ms=clamp_pty_yield_time_ms(yield_time_ms), + max_output_tokens=max_output_tokens, + ) + return await self._finalize_pty_update( + process_id=process_id, + entry=entry, + output=output, + original_token_count=original_token_count, + output_closed=output_closed, + ) + + async def _run_pty_reader(self, entry: _CreateOSPtySessionEntry) -> None: + sdk = _import_createos_sdk() + loop = asyncio.get_running_loop() + + def connect(after_sequence: int) -> Any: + stream = self._sandbox.processes.connect( + entry.provider_process_id, + sdk.ManagedProcessConnectOptions( + timeout=self.state.timeouts.fast_op_s, + after_sequence=after_sequence, + ), + ) + try: + _set_pty_stream_read_timeout(stream, _PTY_STREAM_READ_TIMEOUT_S) + except BaseException: + try: + stream.close() + except Exception: + pass + raise + return stream + + def append_output(data: bytes) -> None: + entry.output_chunks.append(data) + entry.output_notify.set() + + def consume(stream: Any, after_sequence: int) -> tuple[int, bool]: + sequence = after_sequence + try: + with stream: + for event in stream: + event_type = str(getattr(event, "type", "")) + if event_type == "error": + message = str(getattr(event, "error_message", "") or "") + loop.call_soon_threadsafe( + append_output, + (message or "CreateOS PTY output stream failed").encode( + "utf-8", errors="replace" + ), + ) + entry.exit_code = 1 + return sequence, False + event_sequence = int(getattr(event, "sequence", 0) or 0) + if event_sequence > 0: + if event_sequence <= sequence: + continue + sequence = event_sequence + if event_type == "data": + data = bytes(getattr(event, "data", b"") or b"") + if data: + loop.call_soon_threadsafe(append_output, data) + elif event_type == "exit": + exit_code = getattr(event, "exit_code", None) + if exit_code is not None: + entry.exit_code = int(exit_code) + return sequence, False + except (httpx.ReadTimeout, TimeoutError): + return sequence, True + return sequence, False + + try: + sequence = 0 + while not entry.termination_started: + stream = await _to_thread_settled( + connect, + sequence, + timeout=self.state.timeouts.fast_op_s + 1, + interrupted_result_cleanup=self._close_pty_stream, + ) + if entry.termination_started: + await self._close_pty_stream(stream) + break + entry.stream = stream + consume_task = asyncio.create_task(asyncio.to_thread(consume, stream, sequence)) + try: + sequence, idle_timeout = await asyncio.shield(consume_task) + except asyncio.CancelledError: + await _settle_pty_cleanup(self._close_pty_stream(stream)) + await _settle_pty_cleanup(cast(Awaitable[None], consume_task)) + raise + finally: + if entry.stream is stream: + entry.stream = None + if not idle_timeout: + break + except Exception as e: + log_tool_action_debug(logger, "CreateOS PTY output stream failed", e) + finally: + entry.output_closed.set() + entry.output_notify.set() + + async def _close_pty_stream(self, stream: Any) -> None: + await _to_thread_settled( + stream.close, + timeout=self.state.timeouts.cleanup_s, + ) + + async def pty_write_stdin( + self, + *, + session_id: int, + chars: str, + yield_time_s: float | None = None, + max_output_tokens: int | None = None, + ) -> PtyExecUpdate: + self._assert_session_usable() + generation = self._capture_pty_admission_generation() + async with self._pty_lock: + entry = self._resolve_pty_session_entry( + pty_processes=self._pty_sessions, + session_id=session_id, + ) + + if chars: + async with entry.termination_lock: + self._assert_session_usable() + self._assert_pty_admission_open(generation) + if entry.termination_started: + raise SandboxRuntimeError( + message="CreateOS PTY process is stopping", + error_code=ErrorCode.EXEC_TRANSPORT_ERROR, + op="exec", + context={"backend": "createos", "sandbox_id": self.state.sandbox_id}, + retryable=False, + ) + await _to_thread_settled( + self._sandbox.processes.input, + entry.provider_process_id, + chars, + timeout=self.state.timeouts.fast_op_s, + ) + + yield_time_ms = 250 if yield_time_s is None else int(yield_time_s * 1000) + output, original_token_count, output_closed = await self._collect_pty_output( + entry=entry, + yield_time_ms=resolve_pty_write_yield_time_ms( + yield_time_ms=yield_time_ms, + input_empty=chars == "", + ), + max_output_tokens=max_output_tokens, + ) + entry.last_used = time.monotonic() + return await self._finalize_pty_update( + process_id=session_id, + entry=entry, + output=output, + original_token_count=original_token_count, + output_closed=output_closed, + ) + + async def _collect_pty_output( + self, + *, + entry: _CreateOSPtySessionEntry, + yield_time_ms: int, + max_output_tokens: int | None, + ) -> tuple[bytes, int | None, bool]: + return await collect_pty_output( + output_chunks=entry.output_chunks, + output_lock=entry.output_lock, + output_notify=entry.output_notify, + is_done=entry.output_closed.is_set, + yield_time_ms=yield_time_ms, + max_output_tokens=max_output_tokens, + ) + + async def _finalize_pty_update( + self, + *, + process_id: int, + entry: _CreateOSPtySessionEntry, + output: bytes, + original_token_count: int | None, + output_closed: bool, + ) -> PtyExecUpdate: + exit_code = entry.exit_code if output_closed else None + live_process_id: int | None = process_id + if output_closed: + await self._terminate_pty_entry(entry) + await self._remove_pty_entry(process_id, entry) + live_process_id = None + return PtyExecUpdate( + process_id=live_process_id, + output=output, + exit_code=exit_code, + original_token_count=original_token_count, + ) + + async def pty_terminate_all(self) -> None: + async with self._lifecycle_lock: + await self._pty_terminate_all_locked() + + async def _pty_terminate_all_locked(self) -> None: + async with self._pty_lock: + entries = list(self._pty_sessions.items()) + cleanup_error: BaseException | None = None + for process_id, entry in entries: + try: + await self._terminate_pty_entry(entry) + except BaseException as e: + if cleanup_error is None: + cleanup_error = e + else: + await self._remove_pty_entry(process_id, entry) + if cleanup_error is not None: + raise cleanup_error + + async def _remove_pty_entry( + self, + process_id: int, + entry: _CreateOSPtySessionEntry, + ) -> None: + async with self._pty_lock: + if self._pty_sessions.get(process_id) is entry: + self._pty_sessions.pop(process_id) + self._reserved_pty_process_ids.discard(process_id) + + def _prune_pty_sessions_if_needed( + self, + ) -> tuple[int, _CreateOSPtySessionEntry] | None: + if len(self._pty_sessions) < PTY_PROCESSES_MAX: + return None + meta = [ + (process_id, entry.last_used, entry.output_closed.is_set()) + for process_id, entry in self._pty_sessions.items() + ] + process_id = process_id_to_prune_from_meta(meta) + if process_id is None: + return None + entry = self._pty_sessions.get(process_id) + if entry is None: + return None + return process_id, entry + + async def _terminate_pty_entry(self, entry: _CreateOSPtySessionEntry) -> None: + async with entry.termination_lock: + entry.termination_started = True + deletion_error: BaseException | None = None + if not entry.provider_deleted: + try: + await _to_thread_settled( + self._sandbox.processes.delete, + entry.provider_process_id, + timeout=self.state.timeouts.cleanup_s, + on_success=entry.mark_provider_deleted, + ) + except BaseException as e: + deletion_error = e + stream = entry.stream + entry.stream = None + if stream is not None: + try: + await _to_thread_settled( + stream.close, + timeout=self.state.timeouts.cleanup_s, + ) + except Exception as e: + log_tool_action_debug(logger, "CreateOS PTY stream close failed", e) + reader_task = entry.reader_task + entry.reader_task = None + if reader_task is not None and reader_task is not asyncio.current_task(): + await _settle_pty_cleanup(reader_task) + if deletion_error is not None: + raise deletion_error + + async def read(self, path: Path | str, *, user: str | User | None = None) -> io.IOBase: + error_path = posix_path_as_path(coerce_posix_path(path)) + if user is not None: + workspace_path = await self._check_read_with_exec(path, user=user) + else: + workspace_path = await self._validate_path_access(path) + + try: + return io.BytesIO( + await self._download_file( + sandbox_path_str(workspace_path), + timeout=self.state.timeouts.file_download_s, + ) + ) + except Exception as e: + if _provider_status_code(e) == 404: + raise WorkspaceReadNotFoundError(path=error_path, cause=e) from e + raise WorkspaceArchiveReadError(path=error_path, cause=e) from e + + async def write( + self, + path: Path | str, + data: io.IOBase, + *, + user: str | User | None = None, + ) -> None: + error_path = posix_path_as_path(coerce_posix_path(path)) + if user is not None: + await self._check_write_with_exec(path, user=user) + workspace_path = await self._validate_path_access(path, for_write=True) + payload = data.read() + if isinstance(payload, str): + payload = payload.encode("utf-8") + if not isinstance(payload, bytes | bytearray): + raise WorkspaceWriteTypeError(path=error_path, actual_type=type(payload).__name__) + + try: + await self._upload_file( + sandbox_path_str(workspace_path), + bytes(payload), + timeout=self.state.timeouts.file_upload_s, + ) + except Exception as e: + raise WorkspaceArchiveWriteError(path=workspace_path, cause=e) from e + + async def _download_file(self, path: str, *, timeout: float) -> bytes: + self._assert_session_usable() + sdk = _import_createos_sdk() + + def download() -> bytes: + with self._sandbox.files.download( + path, + sdk.RequestOptions(timeout=timeout), + ) as stream: + return bytes(stream.read()) + + return bytes(await _to_thread_settled(download, timeout=timeout + 1)) + + async def _upload_file(self, path: str, data: bytes, *, timeout: float) -> None: + self._assert_session_usable() + sdk = _import_createos_sdk() + await _to_thread_settled( + self._sandbox.files.upload, + path, + data, + sdk.RequestOptions(timeout=timeout), + timeout=timeout + 1, + ) + + async def running(self) -> bool: + self._assert_session_usable() + try: + await _to_thread_settled( + self._sandbox.refresh, + timeout=self.state.timeouts.fast_op_s, + ) + except Exception as e: + log_tool_action_debug(logger, "CreateOS sandbox health check failed", e) + return False + return str(self._sandbox.status) == "running" + + async def _resolve_exposed_port(self, port: int) -> ExposedPortEndpoint: + self._assert_session_usable() + try: + url = await _to_thread_settled( + self._sandbox.preview_url, + port, + timeout=self.state.timeouts.fast_op_s, + ) + split = urlsplit(url) + if split.hostname is None or split.scheme not in {"http", "https"}: + raise ValueError("CreateOS returned an invalid preview URL") + return ExposedPortEndpoint( + host=split.hostname, + port=split.port or (443 if split.scheme == "https" else 80), + tls=split.scheme == "https", + ) + except Exception as e: + raise ExposedPortUnavailableError( + port=port, + exposed_ports=self.state.exposed_ports, + reason="backend_unavailable", + context={"backend": "createos"}, + cause=e, + retryable=_provider_retryability(e), + ) from e + + def _tar_exclude_args(self) -> list[str]: + return shell_tar_exclude_args(self._persist_workspace_skip_relpaths()) + + @retry_async( + retry_if=lambda exc, self: ( + isinstance(exc, TimeoutError) + or _provider_status_code(exc) in TRANSIENT_HTTP_STATUS_CODES + ) + ) + async def persist_workspace(self) -> io.IOBase: + root = self._workspace_root_path() + tar_path = f"/tmp/createos-persist-{self.state.session_id.hex}.tar" + excludes = " ".join(self._tar_exclude_args()) + tar_cmd = ( + f"tar {excludes} -C {shlex.quote(root.as_posix())} -cf {shlex.quote(tar_path)} ." + ).strip() + + async def create_archive() -> bytes: + try: + result = await self._exec_internal( + "sh", "-c", tar_cmd, timeout=self.state.timeouts.workspace_tar_s + ) + if not result.ok(): + raise WorkspaceArchiveReadError( + path=root, + context={ + "reason": "tar_failed", + "stderr": result.stderr.decode("utf-8", errors="replace"), + }, + retryable=False, + ) + return await self._download_file( + tar_path, + timeout=self.state.timeouts.file_download_s, + ) + except WorkspaceArchiveReadError: + raise + except Exception as e: + raise WorkspaceArchiveReadError(path=root, cause=e) from e + finally: + try: + await self._exec_internal( + "rm", "-f", "--", tar_path, timeout=self.state.timeouts.cleanup_s + ) + except Exception as e: + log_tool_action_debug(logger, "CreateOS persist cleanup failed", e) + + raw = await with_ephemeral_mounts_removed( + self, + create_archive, + error_path=root, + error_cls=WorkspaceArchiveReadError, + operation_error_context_key="snapshot_error_before_remount_corruption", + ) + return io.BytesIO(raw) + + async def hydrate_workspace(self, data: io.IOBase) -> None: + root = self._workspace_root_path() + tar_path = f"/tmp/createos-hydrate-{self.state.session_id.hex}.tar" + payload = data.read() + if isinstance(payload, str): + payload = payload.encode("utf-8") + if not isinstance(payload, bytes | bytearray): + raise WorkspaceWriteTypeError(path=Path(tar_path), actual_type=type(payload).__name__) + raw = bytes(payload) + try: + validate_tar_bytes(raw, allow_external_symlink_targets=False) + except UnsafeTarMemberError as e: + raise WorkspaceArchiveWriteError( + path=root, + context={ + "reason": "unsafe_or_invalid_tar", + "member": e.member, + "detail": str(e), + }, + cause=e, + ) from e + + async def extract_archive() -> None: + try: + await self.mkdir(root, parents=True) + await self._upload_file( + tar_path, + raw, + timeout=self.state.timeouts.file_upload_s, + ) + result = await self._exec_internal( + "sh", + "-c", + f"tar -C {shlex.quote(root.as_posix())} -xf {shlex.quote(tar_path)}", + timeout=self.state.timeouts.workspace_tar_s, + ) + if not result.ok(): + raise WorkspaceArchiveWriteError( + path=root, + context={ + "reason": "tar_extract_failed", + "stderr": result.stderr.decode("utf-8", errors="replace"), + }, + ) + except WorkspaceArchiveWriteError: + raise + except Exception as e: + raise WorkspaceArchiveWriteError(path=root, cause=e) from e + finally: + try: + await self._exec_internal( + "rm", "-f", "--", tar_path, timeout=self.state.timeouts.cleanup_s + ) + except Exception as e: + log_tool_action_debug(logger, "CreateOS hydrate cleanup failed", e) + + await with_ephemeral_mounts_removed( + self, + extract_archive, + error_path=root, + error_cls=WorkspaceArchiveWriteError, + operation_error_context_key="hydrate_error_before_remount_corruption", + ) + + async def _shutdown_backend(self) -> None: + sdk = _import_createos_sdk() + if self._mount_transition_terminal or not self.state.pause_on_exit: + if self._destroy_completed: + return + if not self._destroy_issued: + await _to_thread_settled( + self._sandbox.destroy, + timeout=self.state.timeouts.lifecycle_s, + on_success=self._mark_destroy_issued, + ) + await _to_thread_settled( + self._sandbox.wait_until_destroyed, + sdk.WaitOptions(timeout=self.state.timeouts.lifecycle_s), + timeout=self.state.timeouts.lifecycle_s + 1, + ) + self._destroy_completed = True + else: + if self._pause_completed: + return + if not self._pause_issued: + await _to_thread_settled( + self._sandbox.pause, + timeout=self.state.timeouts.lifecycle_s, + on_success=self._mark_pause_issued, + ) + await _to_thread_settled( + self._sandbox.wait_until_paused, + sdk.WaitOptions(timeout=self.state.timeouts.lifecycle_s), + timeout=self.state.timeouts.lifecycle_s + 1, + ) + self._pause_completed = True + + async def _terminate_ambiguous_mount_transition(self) -> None: + self._mount_transition_terminal = True + self._close_pty_admission() + # Mount cleanup runs in an owned child task. During snapshot hydration, + # start or stop already holds the lifecycle lock while awaiting that child. + # Inherited ownership lets the child terminate without waiting on its + # parent; stale children cannot bypass a later owner's lock. + owner = _lifecycle_owner.get() + if ( + owner is not None + and owner is self._active_lifecycle_owner + and self._lifecycle_lock.locked() + ): + await _settle_pty_cleanup(self._terminate_ambiguous_mount_transition_locked()) + return + async with self._lifecycle_lock: + await _settle_pty_cleanup(self._terminate_ambiguous_mount_transition_locked()) + + async def _terminate_ambiguous_mount_transition_locked(self) -> None: + cleanup_error: BaseException | None = None + try: + await self._before_shutdown() + except BaseException as e: + cleanup_error = e + try: + await self._shutdown_backend() + except BaseException as e: + if cleanup_error is None: + cleanup_error = e + try: + await self._after_shutdown() + except BaseException as e: + if cleanup_error is None: + cleanup_error = e + if cleanup_error is not None: + raise cleanup_error + + +class CreateOSSandboxClient(BaseSandboxClient[CreateOSSandboxClientOptions]): + """Client that manages Agents SDK sessions through the CreateOS SDK.""" + + backend_id = "createos" + + def __init__( + self, + *, + api_key: str | None = None, + base_url: str | None = None, + timeout: float = 60, + instrumentation: Instrumentation | None = None, + dependencies: Dependencies | None = None, + env_vars: dict[str, str] | None = None, + ) -> None: + sdk = _import_createos_sdk() + self._client = sdk.Client(api_key=api_key, base_url=base_url, timeout=timeout) + self._instrumentation = instrumentation or Instrumentation() + self._dependencies = dependencies + self._env_vars = dict(env_vars or {}) + + @staticmethod + def _normalize_timeouts( + value: CreateOSSandboxTimeouts | dict[str, object] | None, + ) -> CreateOSSandboxTimeouts: + if isinstance(value, CreateOSSandboxTimeouts): + return value + if value is None: + return CreateOSSandboxTimeouts() + return CreateOSSandboxTimeouts.model_validate(value) + + def _create_request( + self, + *, + shape: str, + rootfs: str | None, + name: str | None, + env_vars: dict[str, str], + exposed_ports: tuple[int, ...], + network_ids: tuple[str, ...], + disk_mib: int | None, + egress_rules: tuple[str, ...], + ssh_public_keys: tuple[str, ...], + host_id: str | None, + node_selector: dict[str, str] | None, + region: str | None, + auto_pause_after_seconds: int | None, + ) -> Any: + sdk = _import_createos_sdk() + return sdk.CreateSandboxRequest( + shape=shape, + rootfs=rootfs or "", + name=name or "", + networks=[sdk.NetworkEntry(id=network_id) for network_id in network_ids], + disk_mib=disk_mib or 0, + egress_rules=list(egress_rules), + environment_variables=env_vars, + ssh_public_keys=list(ssh_public_keys), + host_id=host_id or "", + node_selector=dict(node_selector or {}), + ingress_enabled=bool(exposed_ports), + region=region or "", + auto_pause_after_seconds=auto_pause_after_seconds or 0, + ) + + async def _create_sandbox( + self, + *, + shape: str, + rootfs: str | None, + name: str | None, + env_vars: dict[str, str], + exposed_ports: tuple[int, ...], + network_ids: tuple[str, ...], + disk_mib: int | None, + egress_rules: tuple[str, ...], + ssh_public_keys: tuple[str, ...], + host_id: str | None, + node_selector: dict[str, str] | None, + region: str | None, + auto_pause_after_seconds: int | None, + timeout: float, + cleanup_timeout: float, + ) -> Any: + sdk = _import_createos_sdk() + request = self._create_request( + shape=shape, + rootfs=rootfs, + name=name, + env_vars=env_vars, + exposed_ports=exposed_ports, + network_ids=network_ids, + disk_mib=disk_mib, + egress_rules=egress_rules, + ssh_public_keys=ssh_public_keys, + host_id=host_id, + node_selector=node_selector, + region=region, + auto_pause_after_seconds=auto_pause_after_seconds, + ) + + async def destroy_interrupted_creation(sandbox: Any) -> None: + try: + await _to_thread_settled( + sandbox.destroy, + timeout=cleanup_timeout, + ) + await _to_thread_settled( + sandbox.wait_until_destroyed, + sdk.WaitOptions(timeout=cleanup_timeout), + timeout=cleanup_timeout + 1, + ) + except BaseException as e: + sandbox_id = str(getattr(sandbox, "id", "unknown")) + raise SandboxRuntimeError( + message=( + "CreateOS interrupted sandbox creation cleanup failed for " + f"sandbox {sandbox_id}" + ), + error_code=ErrorCode.WORKSPACE_START_ERROR, + op="start", + context={"backend": "createos", "sandbox_id": sandbox_id}, + cause=e, + retryable=_provider_retryability(e), + ) from e + + return await _to_thread_settled( + self._client.create_sandbox, + request, + sdk.RequestOptions(timeout=timeout), + timeout=timeout + 1, + interrupted_result_cleanup=destroy_interrupted_creation, + ) + + @redact_mount_error_data + async def create( + self, + *, + snapshot: SnapshotSpec | SnapshotBase | None = None, + manifest: Manifest | None = None, + options: CreateOSSandboxClientOptions, + ) -> SandboxSession: + manifest = manifest or Manifest(root=DEFAULT_CREATEOS_WORKSPACE_ROOT) + self._validate_manifest_for_create(manifest) + rootfs = _require_rootfs(options.rootfs) + timeouts = self._normalize_timeouts(options.timeouts) + session_id = uuid.uuid4() + name = session_id.hex[:22] if options.name is None else options.name + env_vars = dict(self._env_vars if options.env_vars is None else options.env_vars) + sandbox = await self._create_sandbox( + shape=options.shape, + rootfs=rootfs, + name=name, + env_vars=env_vars, + exposed_ports=options.exposed_ports, + network_ids=options.network_ids, + disk_mib=options.disk_mib, + egress_rules=options.egress_rules, + ssh_public_keys=options.ssh_public_keys, + host_id=options.host_id, + node_selector=options.node_selector, + region=options.region, + auto_pause_after_seconds=options.auto_pause_after_seconds, + timeout=timeouts.create_s, + cleanup_timeout=timeouts.lifecycle_s, + ) + state = CreateOSSandboxSessionState( + session_id=session_id, + snapshot=resolve_snapshot(snapshot, str(session_id)), + manifest=manifest, + exposed_ports=options.exposed_ports, + sandbox_id=sandbox.id, + shape=options.shape, + rootfs=rootfs, + pause_on_exit=options.pause_on_exit, + name=name, + timeouts=timeouts, + network_ids=options.network_ids, + disk_mib=options.disk_mib, + egress_rules=options.egress_rules, + ssh_public_keys=options.ssh_public_keys, + host_id=options.host_id, + node_selector=( + dict(options.node_selector) if options.node_selector is not None else None + ), + region=options.region, + auto_pause_after_seconds=options.auto_pause_after_seconds, + ) + state._base_env_vars = env_vars + return self._wrap_session( + CreateOSSandboxSession.from_state(state, sandbox=sandbox), + instrumentation=self._instrumentation, + ) + + async def delete(self, session: SandboxSession) -> SandboxSession: + inner = session._inner + if not isinstance(inner, CreateOSSandboxSession): + raise TypeError("CreateOSSandboxClient.delete expects a CreateOSSandboxSession") + inner.state.pause_on_exit = False + await inner.shutdown() + return session + + @redact_mount_error_data + async def resume(self, state: SandboxSessionState) -> SandboxSession: + if not isinstance(state, CreateOSSandboxSessionState): + raise TypeError("CreateOSSandboxClient.resume expects a CreateOSSandboxSessionState") + state.assert_path_grants_rebound() + state = state.model_copy() + state._base_env_vars = dict(self._env_vars) + sdk = _import_createos_sdk() + sandbox: Any | None = None + try: + sandbox = await _to_thread_settled( + self._client.get_sandbox, + state.sandbox_id, + timeout=state.timeouts.fast_op_s, + ) + except Exception as e: + if _provider_status_code(e) != 404: + raise + log_tool_action_debug(logger, "CreateOS sandbox no longer exists; recreating", e) + + reconnected = False + if sandbox is not None: + status = str(sandbox.status) + if status == "destroying": + try: + await _to_thread_settled( + sandbox.wait_until_destroyed, + sdk.WaitOptions(timeout=state.timeouts.lifecycle_s), + timeout=state.timeouts.lifecycle_s + 1, + ) + except Exception: + if str(sandbox.status) not in _TERMINAL_CREATEOS_STATUSES - {"destroying"}: + raise + sandbox = None + elif status in _TERMINAL_CREATEOS_STATUSES: + sandbox = None + else: + reconnected = True + + if sandbox is None: + rootfs = _require_rootfs(state.rootfs) + sandbox = await self._create_sandbox( + shape=state.shape, + rootfs=rootfs, + name=state.name, + env_vars=state._base_env_vars, + exposed_ports=state.exposed_ports, + network_ids=state.network_ids, + disk_mib=state.disk_mib, + egress_rules=state.egress_rules, + ssh_public_keys=state.ssh_public_keys, + host_id=state.host_id, + node_selector=state.node_selector, + region=state.region, + auto_pause_after_seconds=state.auto_pause_after_seconds, + timeout=state.timeouts.create_s, + cleanup_timeout=state.timeouts.lifecycle_s, + ) + state.sandbox_id = cast(Any, sandbox).id + state.workspace_root_ready = False + + inner = CreateOSSandboxSession.from_state(state, sandbox=sandbox) + inner._set_start_state_preserved(reconnected, system=reconnected) + return self._wrap_session(inner, instrumentation=self._instrumentation) + + async def close(self) -> None: + await _to_thread_settled(self._client.close) + + async def __aenter__(self) -> CreateOSSandboxClient: + return self + + async def __aexit__(self, *_: object) -> None: + await self.close() + + def deserialize_session_state(self, payload: dict[str, object]) -> SandboxSessionState: + return self._deserialize_session_state_payload(payload, CreateOSSandboxSessionState) + + +__all__ = [ + "DEFAULT_CREATEOS_WORKSPACE_ROOT", + "CreateOSSandboxClient", + "CreateOSSandboxClientOptions", + "CreateOSSandboxSession", + "CreateOSSandboxSessionState", + "CreateOSSandboxTimeouts", +] diff --git a/tests/extensions/sandbox/test_createos.py b/tests/extensions/sandbox/test_createos.py new file mode 100644 index 0000000000..8b394a71e6 --- /dev/null +++ b/tests/extensions/sandbox/test_createos.py @@ -0,0 +1,2100 @@ +from __future__ import annotations + +import asyncio +import io +import json +import queue +import subprocess +import sys +import tarfile +import threading +import time +import types +from dataclasses import dataclass +from pathlib import Path +from typing import Any, cast +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +from agents.extensions.sandbox.createos import ( + DEFAULT_CREATEOS_WORKSPACE_ROOT, + CreateOSSandboxClient, + CreateOSSandboxClientOptions, + CreateOSSandboxSession, + CreateOSSandboxSessionState, + CreateOSSandboxTimeouts, +) +from agents.sandbox import Manifest +from agents.sandbox.errors import ( + ErrorCode, + ExecTimeoutError, + ExecTransportError, + ExposedPortUnavailableError, + SandboxRuntimeError, + WorkspaceArchiveReadError, + WorkspaceArchiveWriteError, + WorkspaceReadNotFoundError, +) +from agents.sandbox.snapshot import LocalSnapshot, NoopSnapshot +from agents.sandbox.types import ExecResult + + +@dataclass +class _Options: + timeout: float | None = None + environment_variables: dict[str, str] | None = None + after_sequence: int = 0 + + +class _Request: + def __init__(self, **kwargs: Any) -> None: + self.__dict__.update(kwargs) + + +@dataclass +class _CommandResult: + standard_output: str = "" + standard_error: str = "" + exit_code: int = 0 + error_message: str = "" + + +@dataclass +class _CommandResponse: + result: _CommandResult + + +class _APIError(Exception): + def __init__(self, status_code: int, message: str = "provider error") -> None: + super().__init__(message) + self.status_code = status_code + + +class _OperationTimeout(TimeoutError): + pass + + +class _Download(io.BytesIO): + def __enter__(self) -> _Download: + return self + + def __exit__(self, *_: object) -> None: + self.close() + + +class _FakeFiles: + def __init__(self) -> None: + self.data: dict[str, bytes] = {} + self.uploads: list[tuple[str, bytes, object]] = [] + self.downloads: list[tuple[str, object]] = [] + + def upload(self, path: str, data: bytes, options: object) -> None: + self.uploads.append((path, data, options)) + self.data[path] = data + + def download(self, path: str, options: object) -> _Download: + self.downloads.append((path, options)) + if path not in self.data: + raise _APIError(404, "not found") + return _Download(self.data[path]) + + +@dataclass +class _ProcessEvent: + type: str + data: bytes = b"" + exit_code: int | None = None + error_message: str = "" + sequence: int = 0 + + +class _FakeProcessStream: + _END = object() + + def __init__(self, events: queue.Queue[object], *, timeout: float | None = None) -> None: + self._events = events + self._closed = False + self._response = types.SimpleNamespace( + request=types.SimpleNamespace( + extensions={ + "timeout": { + "connect": timeout, + "read": timeout, + "write": timeout, + "pool": timeout, + } + } + ) + ) + + def __iter__(self) -> _FakeProcessStream: + return self + + def __next__(self) -> _ProcessEvent: + event = self._events.get(timeout=5) + if event is self._END: + self._closed = True + raise StopIteration + assert isinstance(event, _ProcessEvent) + return event + + def close(self) -> None: + if not self._closed: + self._closed = True + self._events.put(self._END) + + def __enter__(self) -> _FakeProcessStream: + return self + + def __exit__(self, *_: object) -> None: + self.close() + + +class _FakeProcesses: + def __init__(self, lifecycle_events: list[str]) -> None: + self._lifecycle_events = lifecycle_events + self.created_requests: list[object] = [] + self.inputs: list[tuple[str, str]] = [] + self.deleted: list[str] = [] + self._events: dict[str, queue.Queue[object]] = {} + + def create(self, request: object) -> _Request: + self.created_requests.append(request) + process_id = f"process-{len(self.created_requests)}" + self._events[process_id] = queue.Queue() + return _Request(process_id=process_id) + + def connect(self, process_id: str, _options: object) -> _FakeProcessStream: + return _FakeProcessStream( + self._events[process_id], + timeout=getattr(_options, "timeout", None), + ) + + def input(self, process_id: str, data: str) -> int: + self.inputs.append((process_id, data)) + events = self._events[process_id] + if data == "exit\n": + events.put(_ProcessEvent(type="data", data=b"terminal-done\n")) + events.put(_ProcessEvent(type="exit", exit_code=0)) + events.put(_FakeProcessStream._END) + else: + events.put(_ProcessEvent(type="data", data=b"terminal-ready\n")) + return len(self.inputs) + + def delete(self, process_id: str) -> _Request: + self._lifecycle_events.append("process_delete") + self.deleted.append(process_id) + events = self._events.get(process_id) + if events is not None: + events.put(_ProcessEvent(type="exit", exit_code=0)) + events.put(_FakeProcessStream._END) + return _Request(process_id=process_id, exit_code=0) + + +def _tar_bytes() -> bytes: + output = io.BytesIO() + with tarfile.open(fileobj=output, mode="w") as archive: + payload = b"persisted\n" + info = tarfile.TarInfo("state.txt") + info.size = len(payload) + archive.addfile(info, io.BytesIO(payload)) + return output.getvalue() + + +class _FakeSandbox: + def __init__(self, sandbox_id: str = "sb-1", *, status: str = "running") -> None: + self.id = sandbox_id + self.status = status + self.files = _FakeFiles() + self.lifecycle_events: list[str] = [] + self.processes = _FakeProcesses(self.lifecycle_events) + self.requests: list[tuple[object, object]] = [] + self.next_result: _CommandResult | None = None + self.next_error: Exception | None = None + self.paused = False + self.resumed = False + self.destroyed = False + self.pause_calls = 0 + self.destroy_calls = 0 + self.wait_until_paused_calls = 0 + self.wait_until_destroyed_calls = 0 + self.pause_errors: list[Exception] = [] + self.destroy_errors: list[Exception] = [] + self.wait_until_paused_errors: list[Exception] = [] + self.wait_until_destroyed_errors: list[Exception] = [] + self.waited_until_destroyed = False + self.destroy_wait_error_status: str | None = None + + def run_command(self, request: object, options: object) -> _CommandResponse: + self.requests.append((request, options)) + if self.next_error is not None: + error = self.next_error + self.next_error = None + raise error + if self.next_result is not None: + result = self.next_result + self.next_result = None + return _CommandResponse(result) + + command = tuple(request.arguments[4:]) + if command and "resolve-workspace-path-" in command[0]: + return _CommandResponse(_CommandResult(standard_output=f"{command[2]}\n")) + command_text = " ".join(command) + if "createos-persist-" in command_text and " -cf " in command_text: + tar_path = next(part for part in command_text.split() if "createos-persist-" in part) + self.files.data[tar_path] = _tar_bytes() + return _CommandResponse(_CommandResult()) + + def refresh(self) -> _FakeSandbox: + return self + + def preview_url(self, port: int) -> str: + return f"https://sandbox.example.test/{port}" + + def pause(self) -> _FakeSandbox: + self.lifecycle_events.append("pause") + self.pause_calls += 1 + if self.pause_errors: + raise self.pause_errors.pop(0) + self.paused = True + self.status = "paused" + return self + + def resume(self) -> _FakeSandbox: + self.resumed = True + self.status = "running" + return self + + def destroy(self) -> None: + self.lifecycle_events.append("destroy") + self.destroy_calls += 1 + if self.destroy_errors: + raise self.destroy_errors.pop(0) + self.destroyed = True + self.status = "destroyed" + + def wait_until_paused(self, _options: object) -> _FakeSandbox: + self.wait_until_paused_calls += 1 + if self.wait_until_paused_errors: + raise self.wait_until_paused_errors.pop(0) + self.status = "paused" + return self + + def wait_until_running(self, _options: object) -> _FakeSandbox: + self.status = "running" + return self + + def wait_until_destroyed(self, _options: object) -> _FakeSandbox: + self.wait_until_destroyed_calls += 1 + self.waited_until_destroyed = True + if self.wait_until_destroyed_errors: + raise self.wait_until_destroyed_errors.pop(0) + if self.destroy_wait_error_status is not None: + self.status = self.destroy_wait_error_status + raise RuntimeError(f"destruction settled as {self.status}") + self.status = "destroyed" + return self + + +class _FakeClient: + current: _FakeClient | None = None + + def __init__(self, **kwargs: Any) -> None: + self.kwargs = kwargs + self.created_requests: list[tuple[object, object]] = [] + self.sandboxes: dict[str, _FakeSandbox] = {} + self.closed = False + type(self).current = self + + def create_sandbox(self, request: object, options: object) -> _FakeSandbox: + self.created_requests.append((request, options)) + sandbox = _FakeSandbox(f"sb-{len(self.created_requests)}") + self.sandboxes[sandbox.id] = sandbox + return sandbox + + def get_sandbox(self, sandbox_id: str) -> _FakeSandbox: + if sandbox_id not in self.sandboxes: + raise _APIError(404, "not found") + return self.sandboxes[sandbox_id] + + def close(self) -> None: + self.closed = True + + +@pytest.fixture(autouse=True) +def fake_createos(monkeypatch: pytest.MonkeyPatch) -> types.SimpleNamespace: + sdk = types.SimpleNamespace( + Client=_FakeClient, + CreateSandboxRequest=_Request, + RunCommandRequest=_Request, + RequestOptions=_Options, + ExecOptions=_Options, + WaitOptions=_Options, + ManagedProcessConnectOptions=_Options, + ManagedProcessCreateRequest=_Request, + NetworkEntry=_Request, + PTYSize=_Request, + OperationTimeout=_OperationTimeout, + APIError=_APIError, + ) + monkeypatch.setitem(sys.modules, "createos", sdk) + return sdk + + +def _state(*, pause_on_exit: bool = False) -> CreateOSSandboxSessionState: + return CreateOSSandboxSessionState( + snapshot=NoopSnapshot(id="test-snapshot"), + manifest=Manifest(root=DEFAULT_CREATEOS_WORKSPACE_ROOT), + sandbox_id="sb-1", + shape="s-4vcpu-4gb", + rootfs="devbox:1", + pause_on_exit=pause_on_exit, + ) + + +def test_package_re_exports_createos_symbols() -> None: + from agents.extensions import sandbox as package + + assert package.CreateOSSandboxClient is CreateOSSandboxClient + assert package.CreateOSSandboxSession is CreateOSSandboxSession + + +def test_package_omits_createos_exports_without_optional_dependency() -> None: + subprocess.run( + [ + sys.executable, + "-c", + "import sys; sys.modules['createos'] = None; " + "from agents.extensions import sandbox; " + "assert not any('CreateOS' in name or 'CREATEOS' in name " + "for name in sandbox.__all__)", + ], + check=True, + capture_output=True, + text=True, + ) + + +def test_options_are_positional_and_round_trip() -> None: + options = CreateOSSandboxClientOptions( + "s-4vcpu-4gb", + "devbox:1", + {"BASE": "1"}, + True, + "demo", + (8080,), + network_ids=("network-1",), + disk_mib=8192, + egress_rules=("tcp:443",), + ssh_public_keys=("ssh-ed25519 test",), + host_id="host-1", + node_selector={"pool": "gpu"}, + region="us-east", + auto_pause_after_seconds=300, + ) + restored = CreateOSSandboxClientOptions.model_validate(options.model_dump()) + assert restored == options + assert restored.type == "createos" + + +@pytest.mark.parametrize( + "kwargs", + [ + {"shape": " "}, + {"shape": "s-4vcpu-4gb", "disk_mib": 0}, + {"shape": "s-4vcpu-4gb", "auto_pause_after_seconds": 0}, + {"shape": "s-4vcpu-4gb", "exposed_ports": (0,)}, + {"shape": "s-4vcpu-4gb", "exposed_ports": (65_536,)}, + ], +) +def test_options_reject_invalid_js_parity_values(kwargs: dict[str, object]) -> None: + with pytest.raises(ValueError): + CreateOSSandboxClientOptions(**kwargs) # type: ignore[arg-type] + + +@pytest.mark.parametrize("name", [None, "caller-name", ""]) +async def test_create_uses_provider_compatible_name(name: str | None) -> None: + client = CreateOSSandboxClient() + session = await client.create( + options=CreateOSSandboxClientOptions("s-1vcpu-2gb", "devbox:1", name=name) + ) + + sdk_client = _FakeClient.current + assert sdk_client is not None + request, _ = sdk_client.created_requests[0] + assert request.name == session.state.name + if name is None: + assert request.name + assert len(request.name) <= 22 + else: + assert request.name == name + + +@pytest.mark.parametrize("rootfs", [None, "", " "]) +async def test_create_requires_rootfs_before_provider_effects(rootfs: str | None) -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + + with pytest.raises(ValueError, match="non-empty rootfs"): + await client.create(options=CreateOSSandboxClientOptions("s-1vcpu-2gb", rootfs)) + + assert sdk_client.created_requests == [] + + +async def test_example_passes_current_createos_api_key_explicitly( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from examples.sandbox.extensions import createos_runner + + monkeypatch.setenv("OPENAI_API_KEY", "model-key") + monkeypatch.setenv("CREATEOS_API_KEY", "provider-key") + monkeypatch.delenv("CREATEOS_SANDBOX_API_KEY", raising=False) + run = AsyncMock(return_value=types.SimpleNamespace(final_output="ok")) + monkeypatch.setattr(createos_runner.Runner, "run", run) + + await createos_runner.main( + model="test-model", + question="Inspect the workspace", + shape="s-1vcpu-2gb", + rootfs="devbox:1", + pause_on_exit=False, + stream=False, + ) + + sdk_client = _FakeClient.current + assert sdk_client is not None + assert sdk_client.kwargs["api_key"] == "provider-key" + assert sdk_client.closed is True + assert run.await_count == 1 + + +async def test_create_environment_is_live_only_and_resume_rebinds_trusted_values() -> None: + creator = CreateOSSandboxClient(env_vars={"API_TOKEN": "initial-secret"}) + session = await creator.create(options=CreateOSSandboxClientOptions("s-1vcpu-2gb", "devbox:1")) + original_sandbox = session._inner._sandbox + sdk_client = _FakeClient.current + assert sdk_client is not None + request, _ = sdk_client.created_requests[0] + assert request.environment_variables == {"API_TOKEN": "initial-secret"} + + payload = creator.serialize_session_state(session.state) + assert "initial-secret" not in json.dumps(payload) + assert "base_env_vars" not in payload + restored = creator.deserialize_session_state( + {**payload, "base_env_vars": {"API_TOKEN": "untrusted-secret"}} + ) + assert isinstance(restored, CreateOSSandboxSessionState) + assert restored._base_env_vars == {} + + resumed_client = CreateOSSandboxClient(env_vars={"API_TOKEN": "rotated-secret"}) + resumed_sdk_client = _FakeClient.current + assert resumed_sdk_client is not None + resumed_sdk_client.sandboxes[original_sandbox.id] = original_sandbox + resumed = await resumed_client.resume(restored) + assert resumed.state is not restored + assert restored._base_env_vars == {} + result = await resumed.exec("printenv", "API_TOKEN", shell=False) + assert result.ok() + _, exec_options = original_sandbox.requests[-1] + assert exec_options.environment_variables == {"API_TOKEN": "rotated-secret"} + + recreated_state = resumed_client.deserialize_session_state(payload) + resumed_sdk_client.sandboxes.pop(original_sandbox.id) + recreated = await resumed_client.resume(recreated_state) + recreated_request, _ = resumed_sdk_client.created_requests[-1] + assert recreated_request.environment_variables == {"API_TOKEN": "rotated-secret"} + assert recreated._inner._sandbox is not original_sandbox + + +async def test_create_start_exec_and_file_transfer() -> None: + client = CreateOSSandboxClient(api_key="secret", base_url="https://api.example.test") + session = await client.create( + manifest=Manifest(root=DEFAULT_CREATEOS_WORKSPACE_ROOT), + options=CreateOSSandboxClientOptions( + "s-4vcpu-4gb", + "devbox:1", + {"BASE": "1"}, + exposed_ports=(8080,), + network_ids=("network-1", "network-2"), + disk_mib=8192, + egress_rules=("tcp:443",), + ssh_public_keys=("ssh-ed25519 test",), + host_id="host-1", + node_selector={"pool": "gpu"}, + region="us-east", + auto_pause_after_seconds=300, + ), + ) + inner = session._inner + assert isinstance(inner, CreateOSSandboxSession) + sandbox = inner._sandbox + + await session.start() + sandbox.next_result = _CommandResult( + standard_output="hello\n", + standard_error="warning\n", + exit_code=3, + ) + result = await session.exec("example", "argument", shell=False) + assert result.stdout == b"hello\n" + assert result.stderr == b"warning\n" + assert result.exit_code == 3 + request, options = sandbox.requests[-1] + assert request.arguments[-2:] == ["example", "argument"] + assert options.environment_variables == {"BASE": "1"} + + await session.write(Path("note.txt"), io.BytesIO(b"contents")) + assert sandbox.files.data["/workspace/note.txt"] == b"contents" + restored = await session.read(Path("note.txt")) + assert restored.read() == b"contents" + + endpoint = await session.resolve_exposed_port(8080) + assert (endpoint.host, endpoint.port, endpoint.tls) == ( + "sandbox.example.test", + 443, + True, + ) + + sdk_client = _FakeClient.current + assert sdk_client is not None + create_request, _ = sdk_client.created_requests[0] + assert create_request.shape == "s-4vcpu-4gb" + assert create_request.rootfs == "devbox:1" + assert create_request.ingress_enabled is True + assert [network.id for network in create_request.networks] == ["network-1", "network-2"] + assert create_request.disk_mib == 8192 + assert create_request.egress_rules == ["tcp:443"] + assert create_request.ssh_public_keys == ["ssh-ed25519 test"] + assert create_request.host_id == "host-1" + assert create_request.node_selector == {"pool": "gpu"} + assert create_request.region == "us-east" + assert create_request.auto_pause_after_seconds == 300 + + +async def test_managed_process_pty_exec_and_stdin() -> None: + client = CreateOSSandboxClient() + session = await client.create(options=CreateOSSandboxClientOptions("s-4vcpu-4gb", "devbox:1")) + sandbox = session._inner._sandbox + + assert session.supports_pty() is True + started = await session.pty_exec_start( + "echo", + "ready", + shell=False, + tty=True, + yield_time_s=0.25, + ) + assert started.process_id is not None + assert started.output == b"terminal-ready\n" + assert started.exit_code is None + + finished = await session.pty_write_stdin( + session_id=started.process_id, + chars="exit\n", + yield_time_s=0.25, + ) + assert finished.process_id is None + assert finished.output == b"terminal-done\n" + assert finished.exit_code == 0 + assert sandbox.processes.inputs == [ + ("process-1", "echo ready\n"), + ("process-1", "exit\n"), + ] + assert sandbox.processes.deleted == ["process-1"] + + +async def test_shutdown_waits_for_inflight_pty_input_before_process_delete() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(pause_on_exit=True), sandbox=sandbox) + started = await session.pty_exec_start("sh", shell=False, tty=True, yield_time_s=0) + assert started.process_id is not None + + input_started = threading.Event() + release_input = threading.Event() + delete_started = threading.Event() + original_input = sandbox.processes.input + original_delete = sandbox.processes.delete + + def blocking_input(process_id: str, data: str) -> int: + if data == "echo writing\n": + input_started.set() + if not release_input.wait(timeout=5): + raise TimeoutError("test did not release PTY input") + sandbox.lifecycle_events.append("process_input") + return original_input(process_id, data) + + def observed_delete(process_id: str) -> _Request: + delete_started.set() + return original_delete(process_id) + + sandbox.processes.input = blocking_input # type: ignore[method-assign] + sandbox.processes.delete = observed_delete # type: ignore[method-assign] + write_task = asyncio.create_task( + session.pty_write_stdin( + session_id=started.process_id, + chars="echo writing\n", + yield_time_s=0, + ) + ) + assert await asyncio.to_thread(input_started.wait, 2) + shutdown_task = asyncio.create_task(session.shutdown()) + try: + await asyncio.sleep(0) + assert shutdown_task.done() is False + assert not await asyncio.to_thread(delete_started.wait, 0.1) + assert sandbox.processes.deleted == [] + finally: + release_input.set() + results = await asyncio.gather(write_task, shutdown_task, return_exceptions=True) + + assert all(not isinstance(result, BaseException) for result in results) + assert sandbox.lifecycle_events.index("process_input") < sandbox.lifecycle_events.index( + "process_delete" + ) + assert sandbox.lifecycle_events.index("process_delete") < sandbox.lifecycle_events.index( + "pause" + ) + + +async def test_pty_rejects_user_before_provider_effects() -> None: + session = CreateOSSandboxSession.from_state(_state(), sandbox=_FakeSandbox()) + + with pytest.raises(NotImplementedError, match="does not support `user`"): + await session.pty_exec_start("id", shell=False, user="sandbox") + + assert session._sandbox.processes.created_requests == [] + + +async def test_pty_support_requires_complete_process_api() -> None: + sandbox = _FakeSandbox() + sandbox.processes.connect = None # type: ignore[method-assign] + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + + assert session.supports_pty() is False + with pytest.raises(NotImplementedError, match="not supported"): + await session.pty_exec_start("echo", "ready", shell=False) + assert sandbox.processes.created_requests == [] + + +async def test_pty_initial_input_failure_cleans_up_process() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + + def fail_input(_process_id: str, _data: str) -> int: + raise RuntimeError("input failed") + + sandbox.processes.input = fail_input # type: ignore[method-assign] + with pytest.raises(RuntimeError, match="input failed"): + await session.pty_exec_start("echo", "ready", shell=False) + + assert sandbox.processes.deleted == ["process-1"] + assert session._pty_sessions == {} + + +async def test_shutdown_waits_for_failed_pty_setup_cleanup_before_pause() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + delete_started = threading.Event() + release_delete = threading.Event() + original_delete = sandbox.processes.delete + + def fail_input(_process_id: str, _data: str) -> int: + raise RuntimeError("input failed") + + def blocking_delete(process_id: str) -> _Request: + delete_started.set() + if not release_delete.wait(timeout=5): + raise TimeoutError("test did not release PTY deletion") + return original_delete(process_id) + + sandbox.processes.input = fail_input # type: ignore[method-assign] + sandbox.processes.delete = blocking_delete # type: ignore[method-assign] + pty = asyncio.create_task(session.pty_exec_start("echo", "ready", shell=False)) + assert await asyncio.to_thread(delete_started.wait, 2) + shutdown = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + assert shutdown.done() is False + release_delete.set() + + with pytest.raises(RuntimeError, match="input failed"): + await pty + await shutdown + assert sandbox.lifecycle_events.index("process_delete") < sandbox.lifecycle_events.index( + "pause" + ) + + +async def test_failed_pty_setup_delete_is_visible_to_queued_shutdown_retry() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + delete_started = threading.Event() + release_delete = threading.Event() + original_delete = sandbox.processes.delete + delete_calls = 0 + + def fail_input(_process_id: str, _data: str) -> int: + raise RuntimeError("input failed") + + def fail_delete_once(process_id: str) -> _Request: + nonlocal delete_calls + delete_calls += 1 + if delete_calls == 1: + delete_started.set() + if not release_delete.wait(timeout=5): + raise TimeoutError("test did not release PTY deletion") + raise RuntimeError("delete failed") + return original_delete(process_id) + + sandbox.processes.input = fail_input # type: ignore[method-assign] + sandbox.processes.delete = fail_delete_once # type: ignore[method-assign] + pty = asyncio.create_task(session.pty_exec_start("echo", "ready", shell=False)) + assert await asyncio.to_thread(delete_started.wait, 2) + shutdown = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + assert shutdown.done() is False + release_delete.set() + + with pytest.raises(RuntimeError, match="delete failed"): + await pty + await shutdown + assert delete_calls == 2 + assert sandbox.lifecycle_events.index("process_delete") < sandbox.lifecycle_events.index( + "pause" + ) + + +async def test_cancelled_pty_setup_cleanup_blocks_queued_shutdown() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + input_started = threading.Event() + release_input = threading.Event() + delete_started = threading.Event() + release_delete = threading.Event() + original_input = sandbox.processes.input + original_delete = sandbox.processes.delete + + def blocking_input(process_id: str, data: str) -> int: + input_started.set() + if not release_input.wait(timeout=5): + raise TimeoutError("test did not release PTY input") + return original_input(process_id, data) + + def blocking_delete(process_id: str) -> _Request: + delete_started.set() + if not release_delete.wait(timeout=5): + raise TimeoutError("test did not release PTY deletion") + return original_delete(process_id) + + sandbox.processes.input = blocking_input # type: ignore[method-assign] + sandbox.processes.delete = blocking_delete # type: ignore[method-assign] + pty = asyncio.create_task(session.pty_exec_start("echo", "ready", shell=False)) + assert await asyncio.to_thread(input_started.wait, 2) + pty.cancel() + release_input.set() + assert await asyncio.to_thread(delete_started.wait, 2) + shutdown = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + assert shutdown.done() is False + release_delete.set() + + with pytest.raises(asyncio.CancelledError): + await pty + await shutdown + assert sandbox.lifecycle_events.index("process_delete") < sandbox.lifecycle_events.index( + "pause" + ) + + +async def test_shutdown_waits_for_inflight_pty_creation_before_pause() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + creation_started = threading.Event() + release_creation = threading.Event() + original_create = sandbox.processes.create + + def blocking_create(request: object) -> _Request: + creation_started.set() + if not release_creation.wait(timeout=5): + raise TimeoutError("test did not release PTY creation") + return original_create(request) + + sandbox.processes.create = blocking_create # type: ignore[method-assign] + exec_task = asyncio.create_task( + session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + ) + assert await asyncio.to_thread(creation_started.wait, 2) + shutdown_task = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + release_creation.set() + + await exec_task + await shutdown_task + + assert sandbox.processes.deleted == ["process-1"] + assert sandbox.lifecycle_events.index("process_delete") < sandbox.lifecycle_events.index( + "pause" + ) + + +@pytest.mark.parametrize("operation", ["stop", "terminate"]) +async def test_pty_cleanup_waits_for_inflight_creation( + monkeypatch: pytest.MonkeyPatch, + operation: str, +) -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + creation_started = threading.Event() + release_creation = threading.Event() + original_create = sandbox.processes.create + + def blocking_create(request: object) -> _Request: + creation_started.set() + if not release_creation.wait(timeout=5): + raise TimeoutError("test did not release PTY creation") + return original_create(request) + + async def record_snapshot() -> None: + sandbox.lifecycle_events.append("persist_snapshot") + + sandbox.processes.create = blocking_create # type: ignore[method-assign] + monkeypatch.setattr(session, "_persist_snapshot", record_snapshot) + exec_task = asyncio.create_task( + session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + ) + cleanup_task: asyncio.Task[None] | None = None + try: + assert await asyncio.to_thread(creation_started.wait, 2) + cleanup_task = asyncio.create_task( + session.stop() if operation == "stop" else session.pty_terminate_all() + ) + await asyncio.sleep(0) + assert not cleanup_task.done() + finally: + release_creation.set() + if cleanup_task is not None: + await asyncio.gather(exec_task, cleanup_task, return_exceptions=True) + else: + await asyncio.gather(exec_task, return_exceptions=True) + + exec_task.result() + assert cleanup_task is not None + cleanup_task.result() + assert sandbox.processes.deleted == ["process-1"] + if operation == "stop": + assert sandbox.lifecycle_events.index("process_delete") < sandbox.lifecycle_events.index( + "persist_snapshot" + ) + else: + assert "persist_snapshot" not in sandbox.lifecycle_events + + +async def test_stop_keeps_queued_pty_creation_out_of_snapshot( + monkeypatch: pytest.MonkeyPatch, +) -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + snapshot_started = asyncio.Event() + release_snapshot = asyncio.Event() + admission_captured = asyncio.Event() + original_create = sandbox.processes.create + original_capture = session._capture_pty_admission_generation + + def record_create(request: object) -> _Request: + sandbox.lifecycle_events.append("process_create") + return original_create(request) + + def record_admission() -> int: + admission_captured.set() + return original_capture() + + async def blocking_snapshot() -> None: + sandbox.lifecycle_events.append("snapshot_start") + snapshot_started.set() + await release_snapshot.wait() + sandbox.lifecycle_events.append("snapshot_end") + + sandbox.processes.create = record_create # type: ignore[method-assign] + monkeypatch.setattr(session, "_capture_pty_admission_generation", record_admission) + monkeypatch.setattr(session, "_persist_snapshot", blocking_snapshot) + stop_task: asyncio.Task[None] | None = None + exec_task: asyncio.Task[object] | None = None + try: + async with session._lifecycle_lock: + stop_task = asyncio.create_task(session.stop()) + await asyncio.sleep(0) + exec_task = asyncio.create_task( + session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + ) + await asyncio.wait_for(admission_captured.wait(), timeout=2) + + await asyncio.wait_for(snapshot_started.wait(), timeout=2) + assert sandbox.processes.created_requests == [] + finally: + release_snapshot.set() + if stop_task is not None and exec_task is not None: + await asyncio.gather(stop_task, exec_task, return_exceptions=True) + + assert stop_task is not None + assert exec_task is not None + stop_task.result() + exec_task.result() + assert sandbox.lifecycle_events.index("snapshot_end") < sandbox.lifecycle_events.index( + "process_create" + ) + await session.pty_terminate_all() + + +async def test_shutdown_closes_late_pty_stream_and_waits_for_reader() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + connect_started = threading.Event() + release_connect = threading.Event() + connected_streams: list[_FakeProcessStream] = [] + original_connect = sandbox.processes.connect + + def blocking_connect(process_id: str, options: object) -> _FakeProcessStream: + connect_started.set() + if not release_connect.wait(timeout=5): + raise TimeoutError("test did not release PTY connection") + stream = original_connect(process_id, options) + connected_streams.append(stream) + return stream + + sandbox.processes.connect = blocking_connect # type: ignore[method-assign] + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + assert update.process_id is not None + assert await asyncio.to_thread(connect_started.wait, 2) + entry = session._pty_sessions[update.process_id] + shutdown = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + assert shutdown.done() is False + release_connect.set() + + await shutdown + assert len(connected_streams) == 1 + assert connected_streams[0]._closed is True + assert entry.reader_task is None + assert session._pty_sessions == {} + assert sandbox.status == "paused" + + +async def test_quiet_pty_reconnects_after_idle_timeout_without_losing_input( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from agents.extensions.sandbox.createos import sandbox as createos_adapter + + monkeypatch.setattr(createos_adapter, "_PTY_STREAM_READ_TIMEOUT_S", 0.02) + sandbox = _FakeSandbox() + state = _state(pause_on_exit=True) + session = CreateOSSandboxSession.from_state(state, sandbox=sandbox) + connected_streams: list[_FakeProcessStream] = [] + reconnected = threading.Event() + + class QuietProcessStream(_FakeProcessStream): + def close(self) -> None: + self._closed = True + + def __next__(self) -> _ProcessEvent: + read_timeout = self._response.request.extensions["timeout"]["read"] + try: + event = self._events.get_nowait() + except queue.Empty: + threading.Event().wait(timeout=read_timeout) + raise httpx.ReadTimeout("quiet stream") from None + if event is self._END: + raise StopIteration + assert isinstance(event, _ProcessEvent) + event.sequence = 1 + return event + + def connect(process_id: str, options: object) -> _FakeProcessStream: + if connected_streams: + assert cast(_Options, options).after_sequence == 1 + stream = _FakeProcessStream(sandbox.processes._events[process_id]) + reconnected.set() + else: + stream = QuietProcessStream( + sandbox.processes._events[process_id], + timeout=getattr(options, "timeout", None), + ) + connected_streams.append(stream) + return stream + + sandbox.processes.connect = connect # type: ignore[method-assign] + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.05) + + assert update.process_id is not None + assert await asyncio.to_thread(reconnected.wait, 2) + assert len(connected_streams) == 2 + timeout_config = connected_streams[0]._response.request.extensions["timeout"] + assert timeout_config["connect"] == state.timeouts.fast_op_s + assert timeout_config["read"] == 0.02 + assert session._pty_sessions[update.process_id].output_closed.is_set() is False + await session.pty_write_stdin( + session_id=update.process_id, + chars="exit\n", + yield_time_s=0.05, + ) + + await session.shutdown() + + assert connected_streams[0]._closed is True + assert session._pty_sessions == {} + assert sandbox.status == "paused" + + +async def test_pty_replay_gap_error_stops_reconnecting_and_reports_failure() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + cursors: list[int] = [] + + class GapStream(_FakeProcessStream): + def __init__(self, event: _ProcessEvent) -> None: + super().__init__(queue.Queue()) + self._event: _ProcessEvent | None = event + + def __next__(self) -> _ProcessEvent: + if self._event is not None: + event, self._event = self._event, None + return event + threading.Event().wait(timeout=0.02) + raise httpx.ReadTimeout("idle after event") + + def close(self) -> None: + self._closed = True + + def connect(_process_id: str, options: object) -> GapStream: + cursors.append(cast(_Options, options).after_sequence) + event = ( + _ProcessEvent("data", data=b"ready\n", sequence=1) + if len(cursors) == 1 + else _ProcessEvent("error", error_message="output replay gap", sequence=1) + ) + return GapStream(event) + + sandbox.processes.connect = connect # type: ignore[method-assign] + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + + assert update.process_id is None + assert update.exit_code == 1 + assert b"ready\n" in update.output + assert b"output replay gap" in update.output + assert cursors == [0, 1] + assert sandbox.processes.deleted == ["process-1"] + + +async def test_pty_termination_settles_when_closing_stream_does_not_wake_read( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from agents.extensions.sandbox.createos import sandbox as createos_adapter + + monkeypatch.setattr(createos_adapter, "_PTY_STREAM_READ_TIMEOUT_S", 0.1) + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + read_blocked = threading.Event() + + class BlockingProcessStream(_FakeProcessStream): + def __init__(self, events: queue.Queue[object], *, timeout: float | None = None) -> None: + super().__init__(events, timeout=timeout) + self._initial_output_read = False + + def __next__(self) -> _ProcessEvent: + if not self._initial_output_read: + self._initial_output_read = True + return super().__next__() + read_blocked.set() + read_timeout = self._response.request.extensions["timeout"]["read"] + threading.Event().wait(timeout=read_timeout) + raise httpx.ReadTimeout("provider stream stayed idle after process deletion") + + def close(self) -> None: + self._closed = True + + def connect(process_id: str, options: object) -> BlockingProcessStream: + return BlockingProcessStream( + sandbox.processes._events[process_id], + timeout=getattr(options, "timeout", None), + ) + + sandbox.processes.connect = connect # type: ignore[method-assign] + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + assert update.process_id is not None + assert await asyncio.to_thread(read_blocked.wait, 2) + + await asyncio.wait_for(session.pty_terminate_all(), timeout=2) + + assert sandbox.processes.deleted == ["process-1"] + assert session._pty_sessions == {} + + +async def test_incompatible_pty_stream_closes_before_iteration_and_cleans_up() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + connected_streams: list[_FakeProcessStream] = [] + iteration_started = False + + class IncompatibleProcessStream(_FakeProcessStream): + def __init__(self, events: queue.Queue[object]) -> None: + super().__init__(events) + del self._response + + def __next__(self) -> _ProcessEvent: + nonlocal iteration_started + iteration_started = True + return super().__next__() + + def connect(process_id: str, _options: object) -> _FakeProcessStream: + stream = IncompatibleProcessStream(sandbox.processes._events[process_id]) + connected_streams.append(stream) + return stream + + sandbox.processes.connect = connect # type: ignore[method-assign] + + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + + assert update.process_id is None + assert update.output == b"" + assert len(connected_streams) == 1 + assert connected_streams[0]._closed is True + assert iteration_started is False + assert sandbox.processes.deleted == ["process-1"] + assert session._pty_sessions == {} + + +async def test_shutdown_rejects_pty_queued_before_lifecycle_admission() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + start_blocked = asyncio.Event() + release_start = asyncio.Event() + pty_generation_captured = asyncio.Event() + original_ensure_started = session._ensure_backend_started + original_capture_generation = session._capture_pty_admission_generation + + async def blocking_ensure_started() -> None: + start_blocked.set() + await release_start.wait() + await original_ensure_started() + + def capture_generation() -> int: + generation = original_capture_generation() + pty_generation_captured.set() + return generation + + session._ensure_backend_started = blocking_ensure_started # type: ignore[method-assign] + session._capture_pty_admission_generation = capture_generation # type: ignore[method-assign] + start_task = asyncio.create_task(session.start()) + await start_blocked.wait() + pty_task = asyncio.create_task( + session.pty_exec_start("echo", "queued", shell=False, yield_time_s=0.25) + ) + await pty_generation_captured.wait() + shutdown_task = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + release_start.set() + + await start_task + with pytest.raises(SandboxRuntimeError, match="stopping"): + await pty_task + await shutdown_task + + assert sandbox.processes.created_requests == [] + assert sandbox.status == "paused" + + +async def test_concurrent_shutdown_keeps_pty_admission_closed_after_start() -> None: + sandbox = _FakeSandbox(status="paused") + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + wait_started = threading.Event() + release_wait = threading.Event() + original_wait = sandbox.wait_until_running + + def blocking_wait(options: object) -> _FakeSandbox: + wait_started.set() + if not release_wait.wait(timeout=5): + raise TimeoutError("test did not release running wait") + return original_wait(options) + + sandbox.wait_until_running = blocking_wait # type: ignore[method-assign] + start_task = asyncio.create_task(session.start()) + assert await asyncio.to_thread(wait_started.wait, 2) + shutdown_task = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + release_wait.set() + + await start_task + await shutdown_task + + with pytest.raises(SandboxRuntimeError, match="stopping"): + await session.pty_exec_start("echo", "ready", shell=False) + assert sandbox.processes.created_requests == [] + assert sandbox.status == "paused" + + +async def test_pty_delete_failure_retains_retry_ownership_before_pause() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + assert update.process_id is not None + original_delete = sandbox.processes.delete + delete_calls = 0 + + def fail_once(process_id: str) -> _Request: + nonlocal delete_calls + delete_calls += 1 + if delete_calls == 1: + raise RuntimeError("delete failed") + return original_delete(process_id) + + sandbox.processes.delete = fail_once # type: ignore[method-assign] + + with pytest.raises(RuntimeError, match="delete failed"): + await session.shutdown() + assert update.process_id in session._pty_sessions + assert sandbox.pause_calls == 0 + + await session.shutdown() + assert session._pty_sessions == {} + assert sandbox.processes.deleted == ["process-1"] + assert sandbox.status == "paused" + + +async def test_cancelled_pty_delete_records_success_before_shutdown_retry() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + update = await session.pty_exec_start("sh", shell=False, yield_time_s=0.25) + assert update.process_id is not None + delete_started = threading.Event() + release_delete = threading.Event() + original_delete = sandbox.processes.delete + delete_calls = 0 + + def blocking_delete(process_id: str) -> _Request: + nonlocal delete_calls + delete_calls += 1 + delete_started.set() + if not release_delete.wait(timeout=5): + raise TimeoutError("test did not release PTY deletion") + return original_delete(process_id) + + sandbox.processes.delete = blocking_delete # type: ignore[method-assign] + write = asyncio.create_task( + session.pty_write_stdin( + session_id=update.process_id, + chars="exit\n", + yield_time_s=0.25, + ) + ) + assert await asyncio.to_thread(delete_started.wait, 2) + write.cancel() + release_delete.set() + + with pytest.raises(asyncio.CancelledError): + await write + assert session._pty_sessions[update.process_id].provider_deleted is True + + await session.shutdown() + assert delete_calls == 1 + assert session._pty_sessions == {} + assert sandbox.status == "paused" + + +async def test_unregistered_pty_delete_failure_retains_retry_ownership() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + original_delete = sandbox.processes.delete + + def fail_input(_process_id: str, _data: str) -> int: + raise RuntimeError("input failed") + + def fail_delete(_process_id: str) -> _Request: + raise RuntimeError("delete failed") + + sandbox.processes.input = fail_input # type: ignore[method-assign] + sandbox.processes.delete = fail_delete # type: ignore[method-assign] + with pytest.raises(RuntimeError, match="delete failed"): + await session.pty_exec_start("echo", "ready", shell=False) + + assert len(session._pty_sessions) == 1 + sandbox.processes.delete = original_delete # type: ignore[method-assign] + await session.pty_terminate_all() + assert session._pty_sessions == {} + assert sandbox.processes.deleted == ["process-1"] + + +async def test_unregistered_pty_cleanup_records_delete_success_after_timeout() -> None: + sandbox = _FakeSandbox() + state = _state(pause_on_exit=True) + state.timeouts = CreateOSSandboxTimeouts(cleanup_s=1) + session = CreateOSSandboxSession.from_state(state, sandbox=sandbox) + original_delete = sandbox.processes.delete + delete_calls = 0 + delete_started = threading.Event() + + def fail_input(_process_id: str, _data: str) -> int: + raise RuntimeError("input failed") + + def slow_delete(process_id: str) -> _Request: + nonlocal delete_calls + delete_calls += 1 + delete_started.set() + time.sleep(1.05) + return original_delete(process_id) + + sandbox.processes.input = fail_input # type: ignore[method-assign] + sandbox.processes.delete = slow_delete # type: ignore[method-assign] + pty = asyncio.create_task(session.pty_exec_start("echo", "ready", shell=False)) + assert await asyncio.to_thread(delete_started.wait, 2) + shutdown = asyncio.create_task(session.shutdown()) + await asyncio.sleep(0) + assert shutdown.done() is False + with pytest.raises(TimeoutError): + await pty + + assert len(session._pty_sessions) == 1 + entry = next(iter(session._pty_sessions.values())) + assert entry.provider_deleted is True + + await shutdown + assert delete_calls == 1 + assert session._pty_sessions == {} + assert sandbox.status == "paused" + + +async def test_terminal_mount_cleanup_rejects_existing_pty_input() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + update = await session.pty_exec_start("echo", "ready", shell=False, yield_time_s=0.25) + assert update.process_id is not None + cleanup_started = asyncio.Event() + release_cleanup = asyncio.Event() + original_terminate_all = session._pty_terminate_all_locked + + async def blocking_terminate_all() -> None: + cleanup_started.set() + await release_cleanup.wait() + await original_terminate_all() + + session._pty_terminate_all_locked = blocking_terminate_all # type: ignore[method-assign] + termination = asyncio.create_task(session._terminate_ambiguous_mount_transition()) + await cleanup_started.wait() + + with pytest.raises(SandboxRuntimeError, match="unavailable"): + await session.pty_write_stdin( + session_id=update.process_id, + chars="blocked", + yield_time_s=0.25, + ) + + release_cleanup.set() + await termination + assert ("process-1", "blocked") not in sandbox.processes.inputs + assert sandbox.status == "destroyed" + + +async def test_cancelled_pty_creation_cleans_up_returned_process() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + creation_started = threading.Event() + release_creation = threading.Event() + original_create = sandbox.processes.create + + def blocking_create(request: object) -> _Request: + creation_started.set() + if not release_creation.wait(timeout=5): + raise TimeoutError("test did not release PTY creation") + return original_create(request) + + sandbox.processes.create = blocking_create # type: ignore[method-assign] + exec_task = asyncio.create_task(session.pty_exec_start("echo", "ready", shell=False)) + assert await asyncio.to_thread(creation_started.wait, 2) + exec_task.cancel() + release_creation.set() + + with pytest.raises(asyncio.CancelledError): + await exec_task + assert sandbox.processes.deleted == ["process-1"] + assert session._pty_sessions == {} + + +async def test_cancelled_create_destroys_provider_resource() -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + create_started = threading.Event() + release_create = threading.Event() + original_create = sdk_client.create_sandbox + + def blocking_create(request: object, options: object) -> _FakeSandbox: + create_started.set() + if not release_create.wait(timeout=5): + raise TimeoutError("test did not release sandbox creation") + return original_create(request, options) + + sdk_client.create_sandbox = blocking_create # type: ignore[method-assign] + create = asyncio.create_task( + client.create(options=CreateOSSandboxClientOptions("s-4vcpu-4gb", "devbox:1")) + ) + assert await asyncio.to_thread(create_started.wait, 2) + create.cancel() + release_create.set() + + with pytest.raises(asyncio.CancelledError): + await create + assert len(sdk_client.created_requests) == 1 + created = next(iter(sdk_client.sandboxes.values())) + assert created.destroyed is True + assert created.waited_until_destroyed is True + + +@pytest.mark.parametrize( + ("failure_attr", "message"), + [ + ("destroy_errors", "destroy cleanup failed"), + ("wait_until_destroyed_errors", "destroy wait cleanup failed"), + ], +) +async def test_cancelled_create_surfaces_cleanup_failure_with_sandbox_id( + failure_attr: str, + message: str, +) -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + sandbox = _FakeSandbox("sb-interrupted") + getattr(sandbox, failure_attr).append(RuntimeError(message)) + create_started = threading.Event() + release_create = threading.Event() + + def blocking_create(request: object, options: object) -> _FakeSandbox: + sdk_client.created_requests.append((request, options)) + create_started.set() + if not release_create.wait(timeout=5): + raise TimeoutError("test did not release sandbox creation") + sdk_client.sandboxes[sandbox.id] = sandbox + return sandbox + + sdk_client.create_sandbox = blocking_create # type: ignore[method-assign] + create = asyncio.create_task( + client.create(options=CreateOSSandboxClientOptions("s-4vcpu-4gb", "devbox:1")) + ) + assert await asyncio.to_thread(create_started.wait, 2) + create.cancel() + release_create.set() + + with pytest.raises(SandboxRuntimeError, match="sb-interrupted") as exc_info: + await create + assert exc_info.value.context["sandbox_id"] == "sb-interrupted" + assert isinstance(exc_info.value.cause, RuntimeError) + + +async def test_exec_maps_timeout_and_transport_errors() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + provider_timeout = _OperationTimeout("slow") + sandbox.next_error = provider_timeout + with pytest.raises(ExecTimeoutError) as timeout_info: + await session.exec("slow", shell=False, timeout=1) + assert timeout_info.value.context["provider_error"] == "slow" + assert timeout_info.value.cause is provider_timeout + + sandbox.next_error = _APIError(401, "unauthorized") + with pytest.raises(ExecTransportError) as exc_info: + await session.exec("private", shell=False) + assert exc_info.value.retryable is False + assert exc_info.value.context["http_status"] == 401 + + +async def test_read_maps_provider_not_found() -> None: + session = CreateOSSandboxSession.from_state(_state(), sandbox=_FakeSandbox()) + with pytest.raises(WorkspaceReadNotFoundError): + await session.read("missing.txt") + + +async def test_exposed_port_must_be_configured() -> None: + session = CreateOSSandboxSession.from_state(_state(), sandbox=_FakeSandbox()) + with pytest.raises(ExposedPortUnavailableError) as exc_info: + await session.resolve_exposed_port(8080) + assert exc_info.value.context["reason"] == "not_configured" + + +async def test_port_provider_failure_is_normalized() -> None: + state = _state() + state.exposed_ports = (8080,) + sandbox = _FakeSandbox() + + def fail(_port: int) -> str: + raise _APIError(503) + + sandbox.preview_url = fail # type: ignore[method-assign] + session = CreateOSSandboxSession.from_state(state, sandbox=sandbox) + with pytest.raises(ExposedPortUnavailableError) as exc_info: + await session.resolve_exposed_port(8080) + assert exc_info.value.retryable is True + + +async def test_portable_workspace_persist_and_hydrate() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + persisted = await session.persist_workspace() + with tarfile.open(fileobj=persisted, mode="r:") as archive: + assert archive.extractfile("state.txt").read() == b"persisted\n" # type: ignore[union-attr] + + await session.hydrate_workspace(io.BytesIO(_tar_bytes())) + assert any("createos-hydrate-" in path for path, _, _ in sandbox.files.uploads) + assert any( + "createos-hydrate-" in " ".join(request.arguments) for request, _ in sandbox.requests + ) + + +async def test_hydrate_detaches_and_restores_ephemeral_mount() -> None: + session = CreateOSSandboxSession.from_state(_state(), sandbox=_FakeSandbox()) + strategy = MagicMock() + strategy.teardown_for_snapshot = AsyncMock() + strategy.restore_after_snapshot = AsyncMock() + mount = MagicMock(mount_strategy=strategy) + mount_path = Path("/workspace/mounted") + manifest = MagicMock(wraps=session.state.manifest) + manifest.root = session.state.manifest.root + manifest.environment = session.state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [(mount, mount_path)] + session.state.manifest = manifest + + await session.hydrate_workspace(io.BytesIO(_tar_bytes())) + + strategy.teardown_for_snapshot.assert_awaited_once_with(mount, session, mount_path) + strategy.restore_after_snapshot.assert_awaited_once_with(mount, session, mount_path) + + +async def test_hydrate_restores_mount_after_extraction_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + session = CreateOSSandboxSession.from_state(_state(), sandbox=_FakeSandbox()) + strategy = MagicMock() + strategy.teardown_for_snapshot = AsyncMock() + strategy.restore_after_snapshot = AsyncMock() + mount = MagicMock(mount_strategy=strategy) + mount_path = Path("/workspace/mounted") + manifest = MagicMock(wraps=session.state.manifest) + manifest.root = session.state.manifest.root + manifest.environment = session.state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [(mount, mount_path)] + session.state.manifest = manifest + + async def fail_extract(*command: str | Path, timeout: float | None = None) -> ExecResult: + _ = timeout + if command[:2] == ("sh", "-c") and "tar -C" in str(command[2]): + return ExecResult(stdout=b"", stderr=b"extract failed", exit_code=1) + return ExecResult(stdout=b"", stderr=b"", exit_code=0) + + monkeypatch.setattr(session, "_exec_internal", fail_extract) + with pytest.raises(WorkspaceArchiveWriteError): + await session.hydrate_workspace(io.BytesIO(_tar_bytes())) + + strategy.restore_after_snapshot.assert_awaited_once_with(mount, session, mount_path) + + +async def test_hydrate_restores_prior_mount_after_partial_teardown_failure() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + first_strategy = MagicMock() + first_strategy.teardown_for_snapshot = AsyncMock() + first_strategy.restore_after_snapshot = AsyncMock() + second_strategy = MagicMock() + second_strategy.teardown_for_snapshot = AsyncMock(side_effect=RuntimeError("detach failed")) + second_strategy.restore_after_snapshot = AsyncMock() + first_mount = MagicMock(mount_strategy=first_strategy) + second_mount = MagicMock(mount_strategy=second_strategy) + first_path = Path("/workspace/first") + second_path = Path("/workspace/second") + manifest = MagicMock(wraps=session.state.manifest) + manifest.root = session.state.manifest.root + manifest.environment = session.state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [ + (first_mount, first_path), + (second_mount, second_path), + ] + session.state.manifest = manifest + + with pytest.raises(WorkspaceArchiveWriteError): + await session.hydrate_workspace(io.BytesIO(_tar_bytes())) + + first_strategy.restore_after_snapshot.assert_awaited_once_with(first_mount, session, first_path) + second_strategy.restore_after_snapshot.assert_not_awaited() + assert sandbox.destroyed is True + assert sandbox.paused is False + + +async def test_start_snapshot_mount_failure_destroys_without_waiting_on_own_lock( + tmp_path: Path, +) -> None: + sandbox = _FakeSandbox() + state = _state(pause_on_exit=True) + snapshot = LocalSnapshot(id="restorable-workspace", base_path=tmp_path) + await snapshot.persist(io.BytesIO(_tar_bytes())) + state.snapshot = snapshot + session = CreateOSSandboxSession.from_state(state, sandbox=sandbox) + strategy = MagicMock() + strategy.teardown_for_snapshot = AsyncMock(side_effect=RuntimeError("detach failed")) + mount = MagicMock(mount_strategy=strategy) + manifest = MagicMock(wraps=state.manifest) + manifest.root = state.manifest.root + manifest.environment = state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [(mount, Path("/workspace/mounted"))] + state.manifest = manifest + + with pytest.raises(WorkspaceArchiveWriteError): + await asyncio.wait_for(session.start(), timeout=2) + + assert sandbox.destroyed is True + assert sandbox.waited_until_destroyed is True + assert sandbox.paused is False + assert session._active_lifecycle_owner is None + + +async def test_stop_snapshot_mount_failure_destroys_without_waiting_on_own_lock( + tmp_path: Path, +) -> None: + sandbox = _FakeSandbox() + state = _state(pause_on_exit=True) + state.snapshot = LocalSnapshot(id="persisted-workspace", base_path=tmp_path) + session = CreateOSSandboxSession.from_state(state, sandbox=sandbox) + strategy = MagicMock() + strategy.teardown_for_snapshot = AsyncMock(side_effect=RuntimeError("detach failed")) + mount = MagicMock(mount_strategy=strategy) + manifest = MagicMock(wraps=state.manifest) + manifest.root = state.manifest.root + manifest.environment = state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [(mount, Path("/workspace/mounted"))] + state.manifest = manifest + + with pytest.raises(WorkspaceArchiveReadError): + await asyncio.wait_for(session.stop(), timeout=2) + + assert sandbox.destroyed is True + assert sandbox.waited_until_destroyed is True + assert sandbox.paused is False + assert session._active_lifecycle_owner is None + + +async def test_persist_restore_ambiguity_force_destroys_paused_session() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + strategy = MagicMock() + strategy.teardown_for_snapshot = AsyncMock() + strategy.restore_after_snapshot = AsyncMock(side_effect=RuntimeError("restore failed")) + mount = MagicMock(mount_strategy=strategy) + mount_path = Path("/workspace/mounted") + manifest = MagicMock(wraps=session.state.manifest) + manifest.root = session.state.manifest.root + manifest.environment = session.state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [(mount, mount_path)] + session.state.manifest = manifest + + with pytest.raises(WorkspaceArchiveReadError): + await session.persist_workspace() + + assert sandbox.destroyed is True + assert sandbox.waited_until_destroyed is True + assert sandbox.paused is False + with pytest.raises(SandboxRuntimeError) as exc_info: + await session.exec("true", shell=False) + assert exc_info.value.error_code == ErrorCode.MOUNT_FAILED + + +async def test_hydrate_waits_for_cancelled_provider_call_before_restoring_mount() -> None: + sandbox = _FakeSandbox() + original_run_command = sandbox.run_command + extract_started = threading.Event() + release_extract = threading.Event() + + def blocking_run_command(request: object, options: object) -> _CommandResponse: + command_text = " ".join(request.arguments[4:]) + if "createos-hydrate-" in command_text and " -xf " in command_text: + extract_started.set() + if not release_extract.wait(timeout=5): + raise TimeoutError("test did not release archive extraction") + return original_run_command(request, options) + + sandbox.run_command = blocking_run_command # type: ignore[method-assign] + session = CreateOSSandboxSession.from_state(_state(), sandbox=sandbox) + strategy = MagicMock() + strategy.teardown_for_snapshot = AsyncMock() + strategy.restore_after_snapshot = AsyncMock() + mount = MagicMock(mount_strategy=strategy) + mount_path = Path("/workspace/mounted") + manifest = MagicMock(wraps=session.state.manifest) + manifest.root = session.state.manifest.root + manifest.environment = session.state.manifest.environment + manifest.ephemeral_mount_targets.return_value = [(mount, mount_path)] + session.state.manifest = manifest + + hydrate = asyncio.create_task(session.hydrate_workspace(io.BytesIO(_tar_bytes()))) + assert await asyncio.to_thread(extract_started.wait, 2) + hydrate.cancel() + await asyncio.sleep(0) + strategy.restore_after_snapshot.assert_not_awaited() + + release_extract.set() + with pytest.raises(asyncio.CancelledError): + await hydrate + strategy.restore_after_snapshot.assert_awaited_once_with(mount, session, mount_path) + + +@pytest.mark.parametrize("pause_on_exit", [False, True]) +async def test_shutdown_obeys_pause_on_exit(pause_on_exit: bool) -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=pause_on_exit), + sandbox=sandbox, + ) + await session.shutdown() + assert sandbox.paused is pause_on_exit + assert sandbox.destroyed is not pause_on_exit + assert sandbox.wait_until_paused_calls == int(pause_on_exit) + assert sandbox.wait_until_destroyed_calls == int(not pause_on_exit) + + +@pytest.mark.parametrize("shutdown_first", [False, True]) +async def test_explicit_delete_destroys_pause_on_exit_session(shutdown_first: bool) -> None: + client = CreateOSSandboxClient() + session = await client.create( + options=CreateOSSandboxClientOptions("s-1vcpu-2gb", "devbox:1", pause_on_exit=True) + ) + sandbox = session._inner._sandbox + if shutdown_first: + await session.shutdown() + assert sandbox.status == "paused" + + await client.delete(session) + + assert sandbox.status == "destroyed" + assert sandbox.destroy_calls == 1 + assert sandbox.wait_until_destroyed_calls == 1 + + +async def test_restarted_paused_session_pauses_again_on_next_shutdown() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + + await session.shutdown() + await session.start() + await session.shutdown() + + assert sandbox.status == "paused" + assert sandbox.pause_calls == 2 + assert sandbox.wait_until_paused_calls == 2 + + +async def test_cancelled_restart_finishes_activation_before_next_pause() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + await session.shutdown() + wait_started = threading.Event() + release_wait = threading.Event() + + def blocking_wait(_options: object) -> _FakeSandbox: + wait_started.set() + if not release_wait.wait(timeout=5): + raise TimeoutError("test did not release running wait") + sandbox.status = "running" + return sandbox + + sandbox.wait_until_running = blocking_wait # type: ignore[method-assign] + start = asyncio.create_task(session.start()) + assert await asyncio.to_thread(wait_started.wait, 2) + start.cancel() + release_wait.set() + + with pytest.raises(asyncio.CancelledError): + await start + await session.shutdown() + + assert sandbox.status == "paused" + assert sandbox.pause_calls == 2 + assert sandbox.wait_until_paused_calls == 2 + + +@pytest.mark.parametrize( + ("pause_on_exit", "failure_point", "error_attr", "message"), + [ + (True, "pause", "pause_errors", "pause failed"), + (True, "pause_wait", "wait_until_paused_errors", "pause wait failed"), + (False, "destroy", "destroy_errors", "destroy failed"), + (False, "destroy_wait", "wait_until_destroyed_errors", "destroy wait failed"), + ], +) +async def test_shutdown_propagates_cleanup_failures_and_retries( + pause_on_exit: bool, + failure_point: str, + error_attr: str, + message: str, +) -> None: + sandbox = _FakeSandbox() + getattr(sandbox, error_attr).append(RuntimeError(message)) + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=pause_on_exit), + sandbox=sandbox, + ) + + with pytest.raises(RuntimeError, match=message): + await session.shutdown() + await session.shutdown() + await session.shutdown() + + if pause_on_exit: + assert sandbox.pause_calls == (2 if failure_point == "pause" else 1) + assert sandbox.wait_until_paused_calls == (2 if failure_point == "pause_wait" else 1) + assert sandbox.destroy_calls == 0 + else: + assert sandbox.destroy_calls == (2 if failure_point == "destroy" else 1) + assert sandbox.wait_until_destroyed_calls == (2 if failure_point == "destroy_wait" else 1) + assert sandbox.pause_calls == 0 + + +@pytest.mark.parametrize( + ("pause_on_exit", "mutation_name"), + [(True, "pause"), (False, "destroy")], +) +async def test_cancelled_cleanup_records_success_before_retry( + pause_on_exit: bool, + mutation_name: str, +) -> None: + client = CreateOSSandboxClient() + session = await client.create( + options=CreateOSSandboxClientOptions( + "s-4vcpu-4gb", + "devbox:1", + pause_on_exit=pause_on_exit, + ) + ) + sandbox = session._inner._sandbox + mutation_started = threading.Event() + release_mutation = threading.Event() + original_mutation = getattr(sandbox, mutation_name) + + def blocking_mutation() -> object: + mutation_started.set() + if not release_mutation.wait(timeout=5): + raise TimeoutError("test did not release lifecycle mutation") + return original_mutation() + + setattr(sandbox, mutation_name, blocking_mutation) + shutdown = asyncio.create_task(session.shutdown()) + assert await asyncio.to_thread(mutation_started.wait, 2) + shutdown.cancel() + release_mutation.set() + + with pytest.raises(asyncio.CancelledError): + await shutdown + await session.shutdown() + + assert getattr(sandbox, f"{mutation_name}_calls") == 1 + if pause_on_exit: + assert sandbox.wait_until_paused_calls == 1 + else: + assert sandbox.wait_until_destroyed_calls == 1 + + +async def test_failed_running_wait_still_starts_a_new_pause_cycle() -> None: + sandbox = _FakeSandbox() + session = CreateOSSandboxSession.from_state( + _state(pause_on_exit=True), + sandbox=sandbox, + ) + await session.shutdown() + + def fail_running_wait(_options: object) -> _FakeSandbox: + raise RuntimeError("activation wait failed") + + sandbox.wait_until_running = fail_running_wait # type: ignore[method-assign] + with pytest.raises(RuntimeError, match="activation wait failed"): + await session.start() + + await session.shutdown() + assert sandbox.status == "paused" + assert sandbox.pause_calls == 2 + assert sandbox.wait_until_paused_calls == 2 + + +async def test_resume_reconnects_paused_sandbox() -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + sandbox = _FakeSandbox(status="paused") + sdk_client.sandboxes[sandbox.id] = sandbox + + resumed = await client.resume(_state(pause_on_exit=True)) + assert resumed._inner._sandbox is sandbox + assert sandbox.resumed is False + assert resumed._inner._workspace_state_preserved_on_start() is True + + await resumed.start() + assert sandbox.resumed is True + + +async def test_paused_resume_is_owned_before_cancellable_activation() -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + sandbox = _FakeSandbox(status="paused") + wait_started = threading.Event() + release_wait = threading.Event() + + def blocking_wait(_options: object) -> _FakeSandbox: + wait_started.set() + if not release_wait.wait(timeout=5): + raise TimeoutError("test did not release running wait") + return sandbox + + sandbox.wait_until_running = blocking_wait # type: ignore[method-assign] + sdk_client.sandboxes[sandbox.id] = sandbox + resumed = await client.resume(_state(pause_on_exit=True)) + assert sandbox.resumed is False + + start = asyncio.create_task(resumed.start()) + assert await asyncio.to_thread(wait_started.wait, 2) + start.cancel() + release_wait.set() + with pytest.raises(asyncio.CancelledError): + await start + + assert sandbox.resumed is True + await resumed.shutdown() + assert sandbox.paused is True + + +async def test_pausing_resume_settles_pause_before_owned_activation() -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + sandbox = _FakeSandbox(status="pausing") + sdk_client.sandboxes[sandbox.id] = sandbox + + resumed = await client.resume(_state(pause_on_exit=True)) + assert sandbox.resumed is False + + await resumed.start() + + assert sandbox.resumed is True + assert sandbox.status == "running" + + +@pytest.mark.parametrize("status", ["destroying", "destroyed", "failed", "error"]) +async def test_resume_recreates_terminal_sandbox(status: str) -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + terminal = _FakeSandbox(status=status) + sdk_client.sandboxes[terminal.id] = terminal + + resumed = await client.resume(_state()) + + assert resumed._inner._sandbox is not terminal + assert resumed._inner._workspace_state_preserved_on_start() is False + assert len(sdk_client.created_requests) == 1 + assert terminal.waited_until_destroyed is (status == "destroying") + + +@pytest.mark.parametrize("settled_status", ["failed", "error"]) +async def test_resume_recreates_destroying_sandbox_that_settles_terminally( + settled_status: str, +) -> None: + client = CreateOSSandboxClient() + sdk_client = _FakeClient.current + assert sdk_client is not None + terminal = _FakeSandbox(status="destroying") + terminal.destroy_wait_error_status = settled_status + sdk_client.sandboxes[terminal.id] = terminal + + resumed = await client.resume(_state()) + + assert resumed._inner._sandbox is not terminal + assert resumed._inner._workspace_state_preserved_on_start() is False + assert len(sdk_client.created_requests) == 1 + assert terminal.waited_until_destroyed is True + + +async def test_resume_recreates_only_missing_sandbox() -> None: + client = CreateOSSandboxClient() + state = _state() + state.network_ids = ("network-1",) + state.disk_mib = 8192 + state.egress_rules = ("tcp:443",) + state.ssh_public_keys = ("ssh-ed25519 test",) + state.host_id = "host-1" + state.node_selector = {"pool": "gpu"} + state.region = "us-east" + state.auto_pause_after_seconds = 300 + resumed = await client.resume(state) + assert resumed.state.sandbox_id == "sb-1" + assert resumed._inner._workspace_state_preserved_on_start() is False + + sdk_client = _FakeClient.current + assert sdk_client is not None + create_request, _ = sdk_client.created_requests[0] + assert [network.id for network in create_request.networks] == ["network-1"] + assert create_request.disk_mib == 8192 + assert create_request.egress_rules == ["tcp:443"] + assert create_request.ssh_public_keys == ["ssh-ed25519 test"] + assert create_request.host_id == "host-1" + assert create_request.node_selector == {"pool": "gpu"} + assert create_request.region == "us-east" + assert create_request.auto_pause_after_seconds == 300 + sdk_client.get_sandbox = lambda _sandbox_id: (_ for _ in ()).throw(_APIError(503)) # type: ignore[method-assign] + with pytest.raises(_APIError, match="provider error"): + await client.resume(_state()) + + +async def test_client_state_serialization_and_close() -> None: + client = CreateOSSandboxClient() + payload = client.serialize_session_state(_state()) + restored = client.deserialize_session_state(payload) + assert isinstance(restored, CreateOSSandboxSessionState) + assert restored.shape == "s-4vcpu-4gb" + + await client.close() + sdk_client = _FakeClient.current + assert sdk_client is not None + assert sdk_client.closed is True diff --git a/tests/fixtures/released_api_contract_policy.json b/tests/fixtures/released_api_contract_policy.json index 471e6ca3f2..59fc7e250c 100644 --- a/tests/fixtures/released_api_contract_policy.json +++ b/tests/fixtures/released_api_contract_policy.json @@ -368,6 +368,10 @@ "aiosqlite": { "requirement": "aiosqlite>=0.21.0" }, + "createos": { + "extra": "createos", + "distribution": "createos-sandbox" + }, "cryptography": { "extra": "encrypt" }, @@ -422,6 +426,12 @@ "CloudflareSandboxClientOptions": "aiohttp", "CloudflareSandboxSession": "aiohttp", "CloudflareSandboxSessionState": "aiohttp", + "CreateOSSandboxClient": "createos", + "CreateOSSandboxClientOptions": "createos", + "CreateOSSandboxSession": "createos", + "CreateOSSandboxSessionState": "createos", + "CreateOSSandboxTimeouts": "createos", + "DEFAULT_CREATEOS_WORKSPACE_ROOT": "createos", "DEFAULT_RUNLOOP_ROOT_WORKSPACE_ROOT": "runloop_api_client", "DEFAULT_RUNLOOP_WORKSPACE_ROOT": "runloop_api_client", "ModalCloudBucketMountStrategy": "modal", diff --git a/tests/test_integration_runner.py b/tests/test_integration_runner.py index f5ec2bdfde..ea7aca282f 100644 --- a/tests/test_integration_runner.py +++ b/tests/test_integration_runner.py @@ -847,6 +847,18 @@ def fake_run_suite(*args: object, **kwargs: Any) -> None: (installation.dependency_module, installation.extra) for installation in supported_installations } + createos_suites = [ + suite + for suite in isolated_suites + if suite["additional_env"]["OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_DEPENDENCIES"] + == "createos" + ] + assert len(createos_suites) == 2 + assert all( + suite["additional_env"]["OPENAI_AGENTS_INTEGRATION_REQUIRED_OPTIONAL_DISTRIBUTION"] + == "createos-sandbox" + for suite in createos_suites + ) @pytest.mark.parametrize("platform", ["linux", "win32"]) diff --git a/tests/test_released_api_contract.py b/tests/test_released_api_contract.py index 8523637fd9..8086e0334b 100644 --- a/tests/test_released_api_contract.py +++ b/tests/test_released_api_contract.py @@ -3508,7 +3508,8 @@ def test_load_submodule_export_policy_collects_artifact_installations(tmp_path: '{"LazyBinding": "binding_dependency"}, "optional_exports": ' '{"ConditionalExport": "export_dependency"}}}, "optional_dependencies": ' '{"binding_dependency": {"requirement": "binding-package>=1"}, ' - '"export_dependency": {"extra": "export-extra"}}, "public_properties": ' + '"export_dependency": {"extra": "export-extra", ' + '"distribution": "provider-wheel"}}, "public_properties": ' '[{"class_name": "ConditionalExport", ' '"module": "agents.submodule", "names": ["status"]}, ' '{"factory_name": "create_client", ' @@ -3536,12 +3537,14 @@ def test_load_submodule_export_policy_collects_artifact_installations(tmp_path: "extra": None, "requirement": "binding-package>=1", "unsupported_platforms": (), + "distribution": None, }, { "dependency_module": "export_dependency", "extra": "export-extra", "requirement": None, "unsupported_platforms": (), + "distribution": "provider-wheel", }, ] assert policy.canonical_imports == ( diff --git a/uv.lock b/uv.lock index 2a3989398b..35f8abbd8f 100644 --- a/uv.lock +++ b/uv.lock @@ -13,6 +13,7 @@ exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for exclude-newer-span = "P7D" [options.exclude-newer-package] +createos-sandbox = false openai = false [[package]] @@ -801,6 +802,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/84/19/e67f4ae24e232c7f713337f3f4f7c9c58afd0c02866fb07c7b9255a19ed7/coverage-7.10.3-py3-none-any.whl", hash = "sha256:416a8d74dc0adfd33944ba2f405897bab87b7e9e84a391e09d241956bd953ce1", size = 207921, upload-time = "2025-08-10T21:27:38.254Z" }, ] +[[package]] +name = "createos-sandbox" +version = "0.1.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "httpx" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d4/79/c5a9d65ea24b258865f37db9c1b9d48afaba1eb372ed098be7c873c9badd/createos_sandbox-0.1.1.tar.gz", hash = "sha256:cfa544bde569e2c8f840bd23a5733d503ddf95e916e37c3623460a75d5dc8c0e", size = 42607, upload-time = "2026-09-11T12:06:41.982Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/2a/d5/1f024854ef4a5342fec124a6e7469f32b0eb2cbdaf65a5f2d42d074c46c3/createos_sandbox-0.1.1-py3-none-any.whl", hash = "sha256:faee37ccf62b22e4bda09d3c8819b38075e1cc4ccb4651724d72a4c48a4874a0", size = 29040, upload-time = "2026-09-11T12:06:40.486Z" }, +] + [[package]] name = "cryptography" version = "50.0.0" @@ -2644,6 +2657,9 @@ blaxel = [ cloudflare = [ { name = "aiohttp" }, ] +createos = [ + { name = "createos-sandbox" }, +] dapr = [ { name = "aiohttp" }, { name = "dapr" }, @@ -2770,6 +2786,7 @@ requires-dist = [ { name = "boto3", marker = "extra == 's3'", specifier = ">=1.34" }, { name = "cbor2", marker = "extra == 'modal'", specifier = ">=5.9.0" }, { name = "cbor2", marker = "extra == 'vercel'", specifier = ">=5.9.0" }, + { name = "createos-sandbox", marker = "extra == 'createos'", specifier = ">=0.1.0,<0.2" }, { name = "cryptography", marker = "extra == 'encrypt'", specifier = ">=45.0,<51" }, { name = "dapr", marker = "extra == 'dapr'", specifier = ">=1.16.0" }, { name = "daytona", marker = "extra == 'daytona'", specifier = ">=0.155.0" }, @@ -2815,7 +2832,7 @@ requires-dist = [ { name = "websockets", marker = "extra == 'realtime'", specifier = ">=15.0,<17" }, { name = "websockets", marker = "extra == 'voice'", specifier = ">=15.0,<17" }, ] -provides-extras = ["voice", "viz", "litellm", "any-llm", "realtime", "sqlalchemy", "encrypt", "redis", "dapr", "mongodb", "docker", "blaxel", "daytona", "cloudflare", "e2b", "modal", "runloop", "vercel", "s3", "temporal"] +provides-extras = ["voice", "viz", "litellm", "any-llm", "realtime", "sqlalchemy", "encrypt", "redis", "dapr", "mongodb", "docker", "createos", "blaxel", "daytona", "cloudflare", "e2b", "modal", "runloop", "vercel", "s3", "temporal"] [package.metadata.requires-dev] dev = [