From 40f4e67320bf229bee064a33fe4dbd300c6b8e1b Mon Sep 17 00:00:00 2001 From: alex s <46074070+bearsyankees@users.noreply.github.com> Date: Tue, 14 Jul 2026 23:27:06 -0400 Subject: [PATCH] fix(runtime): retry transient sandbox startup failures (#768) * fix(runtime): retry transient sandbox startup failures * fix(runtime): fail closed when sandbox teardown fails --- strix/runtime/backends.py | 134 ++++++++++++++++++++- tests/test_backend_retries.py | 218 ++++++++++++++++++++++++++++++++++ 2 files changed, 349 insertions(+), 3 deletions(-) create mode 100644 tests/test_backend_retries.py diff --git a/strix/runtime/backends.py b/strix/runtime/backends.py index d7eba335..b73fdbad 100644 --- a/strix/runtime/backends.py +++ b/strix/runtime/backends.py @@ -2,7 +2,9 @@ from __future__ import annotations +import asyncio import logging +import os from collections.abc import Awaitable, Callable from typing import TYPE_CHECKING, Any @@ -16,6 +18,128 @@ logger = logging.getLogger(__name__) SandboxBackend = Callable[..., Awaitable[tuple[Any, Any]]] +_DEFAULT_START_ATTEMPTS = 3 +_START_BACKOFF_SECONDS = 2.0 +_TRANSIENT_TIMEOUT_NAMES = { + "ConnectTimeout", + "PoolTimeout", + "ReadTimeout", + "TimeoutError", + "TimeoutException", + "WriteTimeout", +} +_TRANSIENT_CONNECTION_NAMES = { + "ConnectError", + "ConnectionError", + "ConnectionResetError", + "ReadError", + "WriteError", +} + + +def _start_attempts() -> int: + raw = os.environ.get("STRIX_SANDBOX_START_ATTEMPTS") + if raw is None: + return _DEFAULT_START_ATTEMPTS + try: + attempts = int(raw) + except ValueError: + logger.warning( + "Invalid STRIX_SANDBOX_START_ATTEMPTS=%r; using %d", + raw, + _DEFAULT_START_ATTEMPTS, + ) + return _DEFAULT_START_ATTEMPTS + if attempts < 1: + logger.warning( + "STRIX_SANDBOX_START_ATTEMPTS must be positive; using %d", + _DEFAULT_START_ATTEMPTS, + ) + return _DEFAULT_START_ATTEMPTS + return attempts + + +def _exception_chain(error: BaseException) -> list[BaseException]: + chain: list[BaseException] = [] + pending: list[BaseException | None] = [error] + seen: set[int] = set() + while pending: + current = pending.pop() + if current is None or id(current) in seen: + continue + seen.add(id(current)) + chain.append(current) + pending.extend( + ( + current.__cause__, + current.__context__, + getattr(current, "cause", None), + ) + ) + return chain + + +def _is_transient_start_error(error: BaseException) -> bool: + for cause in _exception_chain(error): + name = type(cause).__name__ + module = type(cause).__module__ + if isinstance(cause, TimeoutError | ConnectionError | ConnectionResetError): + return True + if name in _TRANSIENT_TIMEOUT_NAMES: + return True + if name in _TRANSIENT_CONNECTION_NAMES and ( + module.startswith(("httpcore", "httpx", "agents")) + or name in {"ConnectionError", "ConnectionResetError"} + ): + return True + return False + + +async def start_session_with_retry( + client: Any, + create_session: Callable[[], Awaitable[Any]], + *, + attempts: int | None = None, +) -> Any: + """Start a sandbox session, retrying transient transport failures. + + Backend implementations should use this helper when they own both session + creation and ``session.start()`` so failed starts can be torn down before a + retry. The caller owns the manifest and any temporary source directories + until this helper returns. + """ + max_attempts = attempts if attempts is not None else _start_attempts() + for attempt in range(1, max_attempts + 1): + session: Any | None = None + try: + session = await create_session() + assert session is not None + await session.start() + except Exception as exc: + if session is not None: + try: + await client.delete(session) + except Exception as teardown_error: + logger.warning( + "Failed to tear down sandbox after start failure; aborting retry", + exc_info=True, + ) + raise exc from teardown_error + transient = _is_transient_start_error(exc) + if not transient or attempt == max_attempts: + raise + delay = _START_BACKOFF_SECONDS * (2 ** (attempt - 1)) + logger.warning( + "Transient sandbox start failure; retrying attempt %d/%d in %.1fs", + attempt + 1, + max_attempts, + delay, + ) + await asyncio.sleep(delay) + else: + return session + raise AssertionError("sandbox start retry loop completed without returning or raising") + async def _docker_backend( *, @@ -50,8 +174,10 @@ async def _docker_backend( client = StrixDockerSandboxClient(docker.from_env()) client.strix_bind_mounts = bind_mounts or [] options = DockerSandboxClientOptions(image=image, exposed_ports=exposed_ports) - session = await client.create(options=options, manifest=manifest) - await session.start() + session = await start_session_with_retry( + client, + lambda: client.create(options=options, manifest=manifest), + ) return client, session @@ -83,7 +209,9 @@ def register_backend(name: str, backend: SandboxBackend) -> None: Intended for downstream users who ship their own runtime — register before any ``session_manager.create_or_reuse`` call. Re-registering - an existing name overwrites the prior entry. + an existing name overwrites the prior entry. Backends that own both + session creation and ``session.start()`` should use + :func:`start_session_with_retry`. """ _BACKENDS[name] = backend logger.info("Registered sandbox backend: %s", name) diff --git a/tests/test_backend_retries.py b/tests/test_backend_retries.py new file mode 100644 index 00000000..b18d4a0b --- /dev/null +++ b/tests/test_backend_retries.py @@ -0,0 +1,218 @@ +"""Tests for transient sandbox start retries.""" + +from __future__ import annotations + +import shutil +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +import pytest +from agents.sandbox.errors import ( + LocalDirReadError, + WorkspaceArchiveWriteError, + WorkspaceStartError, +) + +from strix.runtime import session_manager +from strix.runtime.backends import start_session_with_retry + + +class _FakeSession: + def __init__(self, failures: list[BaseException]) -> None: + self._failures = iter(failures) + + async def start(self) -> None: + try: + raise next(self._failures) + except StopIteration: + return + + async def resolve_exposed_port(self, _port: int) -> SimpleNamespace: + return SimpleNamespace(tls=False, host="127.0.0.1", port=48080) + + +class _FakeClient: + def __init__(self, *, delete_error: BaseException | None = None) -> None: + self.created = 0 + self.deleted: list[_FakeSession] = [] + self.delete_error = delete_error + + async def create(self) -> _FakeSession: + self.created += 1 + failures: list[BaseException] = [] + if self.created == 1: + failures = [ + WorkspaceStartError( + path=Path("/workspace"), + cause=WorkspaceArchiveWriteError( + path=Path("/workspace"), + cause=TimeoutError("transient transport timeout"), + ), + ) + ] + return _FakeSession(failures) + + async def delete(self, session: _FakeSession) -> None: + self.deleted.append(session) + if self.delete_error is not None: + raise self.delete_error + + +async def test_transient_workspace_failure_retries_and_tears_down( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client = _FakeClient() + sleeps: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + monkeypatch.setattr("strix.runtime.backends.asyncio.sleep", record_sleep) + + session = await start_session_with_retry(client, client.create, attempts=3) + + assert isinstance(session, _FakeSession) + assert client.created == 2 + assert len(client.deleted) == 1 + assert sleeps == [2.0] + + +async def test_non_transient_workspace_failure_does_not_retry() -> None: + client = _FakeClient() + session = _FakeSession([LocalDirReadError(src=Path("/workspace/repo"))]) + + async def create_session() -> _FakeSession: + client.created += 1 + return session + + with pytest.raises(LocalDirReadError): + await start_session_with_retry(client, create_session, attempts=3) + + assert client.created == 1 + assert client.deleted == [session] + + +async def test_teardown_failure_raises_original_error_without_retry() -> None: + start_error = WorkspaceStartError( + path=Path("/workspace"), + cause=TimeoutError("transient transport timeout"), + ) + teardown_error = RuntimeError("teardown failed") + client = _FakeClient(delete_error=teardown_error) + session = _FakeSession([start_error]) + + async def create_session() -> _FakeSession: + client.created += 1 + return session + + with pytest.raises(WorkspaceStartError) as caught: + await start_session_with_retry(client, create_session, attempts=3) + + assert caught.value is start_error + assert caught.value.__cause__ is teardown_error + assert client.created == 1 + assert client.deleted == [session] + + +async def test_each_transient_attempt_is_torn_down(monkeypatch: pytest.MonkeyPatch) -> None: + client = _FakeClient() + client.created = 0 + sessions: list[_FakeSession] = [] + sleeps: list[float] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + monkeypatch.setattr("strix.runtime.backends.asyncio.sleep", record_sleep) + + async def create_session() -> _FakeSession: + client.created += 1 + failures: list[BaseException] = [] + if client.created < 3: + failures = [ + WorkspaceStartError( + path=Path("/workspace"), + cause=TimeoutError("transient transport timeout"), + ) + ] + session = _FakeSession(failures) + sessions.append(session) + return session + + result = await start_session_with_retry(client, create_session, attempts=3) + + assert result is sessions[2] + assert client.deleted == sessions[:2] + assert sleeps == [2.0, 4.0] + + +async def test_staged_dirs_survive_retries_and_cleanup_once( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + repo = tmp_path / "repo" + repo.mkdir() + (repo / "real.txt").write_text("content") + (repo / "link.txt").symlink_to(repo / "real.txt") + + client = _FakeClient() + observed_paths: list[Path] = [] + sleeps: list[float] = [] + original_rmtree = shutil.rmtree # pyright: ignore[reportDeprecated] + removed_paths: list[Path] = [] + + async def record_sleep(delay: float) -> None: + sleeps.append(delay) + + monkeypatch.setattr("strix.runtime.backends.asyncio.sleep", record_sleep) + + def record_rmtree(path: str | Path, **kwargs: Any) -> None: + removed_paths.append(Path(path)) + original_rmtree(path, **kwargs) # pyright: ignore[reportDeprecated] + + monkeypatch.setattr("strix.runtime.session_manager.shutil.rmtree", record_rmtree) + monkeypatch.setattr( + session_manager, + "load_settings", + lambda: SimpleNamespace(runtime=SimpleNamespace(backend="fake")), + ) + monkeypatch.setattr(session_manager, "bootstrap_caido", _bootstrap_caido) + + async def fake_backend(**kwargs: Any) -> tuple[_FakeClient, _FakeSession]: + staged_path = kwargs["manifest"].entries["repo"].src + + async def create_session() -> _FakeSession: + observed_paths.append(Path(staged_path)) + return await client.create() + + session = await start_session_with_retry(client, create_session, attempts=3) + return client, session + + def fake_get_backend(_name: str) -> Any: + return fake_backend + + monkeypatch.setattr(session_manager, "get_backend", fake_get_backend) + + try: + await session_manager.create_or_reuse( + "retry-test", + image="test-image", + local_sources=[ + { + "source_path": str(repo), + "workspace_subdir": "repo", + } + ], + ) + finally: + await session_manager.cleanup("retry-test") + + assert len(observed_paths) == 2 + assert observed_paths[0] == observed_paths[1] + assert observed_paths[0] in removed_paths + assert not observed_paths[0].exists() + + +async def _bootstrap_caido(*_args: Any, **_kwargs: Any) -> object: + return object()