mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 01:16:40 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d31d6fb6b8 |
+130
-3
@@ -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,127 @@ 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_E2B_BOOTSTRAP_ATTEMPTS")
|
||||
if raw is None:
|
||||
return _DEFAULT_START_ATTEMPTS
|
||||
try:
|
||||
attempts = int(raw)
|
||||
except ValueError:
|
||||
logger.warning(
|
||||
"Invalid STRIX_E2B_BOOTSTRAP_ATTEMPTS=%r; using %d",
|
||||
raw,
|
||||
_DEFAULT_START_ATTEMPTS,
|
||||
)
|
||||
return _DEFAULT_START_ATTEMPTS
|
||||
if attempts < 1:
|
||||
logger.warning(
|
||||
"STRIX_E2B_BOOTSTRAP_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", "e2b", "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: # noqa: BLE001
|
||||
logger.warning(
|
||||
"Failed to tear down sandbox after start failure",
|
||||
exc_info=True,
|
||||
)
|
||||
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 +173,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 +208,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)
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""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) -> None:
|
||||
self.created = 0
|
||||
self.deleted: list[_FakeSession] = []
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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_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()
|
||||
Reference in New Issue
Block a user