mirror of
https://github.com/usestrix/strix.git
synced 2026-08-18 01:39:19 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d31d6fb6b8 |
+130
-3
@@ -2,7 +2,9 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
@@ -16,6 +18,127 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
SandboxBackend = Callable[..., Awaitable[tuple[Any, Any]]]
|
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(
|
async def _docker_backend(
|
||||||
*,
|
*,
|
||||||
@@ -50,8 +173,10 @@ async def _docker_backend(
|
|||||||
client = StrixDockerSandboxClient(docker.from_env())
|
client = StrixDockerSandboxClient(docker.from_env())
|
||||||
client.strix_bind_mounts = bind_mounts or []
|
client.strix_bind_mounts = bind_mounts or []
|
||||||
options = DockerSandboxClientOptions(image=image, exposed_ports=exposed_ports)
|
options = DockerSandboxClientOptions(image=image, exposed_ports=exposed_ports)
|
||||||
session = await client.create(options=options, manifest=manifest)
|
session = await start_session_with_retry(
|
||||||
await session.start()
|
client,
|
||||||
|
lambda: client.create(options=options, manifest=manifest),
|
||||||
|
)
|
||||||
return client, session
|
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
|
Intended for downstream users who ship their own runtime — register
|
||||||
before any ``session_manager.create_or_reuse`` call. Re-registering
|
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
|
_BACKENDS[name] = backend
|
||||||
logger.info("Registered sandbox backend: %s", name)
|
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