diff --git a/strix/runtime/backends.py b/strix/runtime/backends.py index f85c99c4..b73fdbad 100644 --- a/strix/runtime/backends.py +++ b/strix/runtime/backends.py @@ -119,11 +119,12 @@ async def start_session_with_retry( if session is not None: try: await client.delete(session) - except Exception: # noqa: BLE001 + except Exception as teardown_error: logger.warning( - "Failed to tear down sandbox after start failure", + "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 diff --git a/tests/test_backend_retries.py b/tests/test_backend_retries.py index 94c0636b..b18d4a0b 100644 --- a/tests/test_backend_retries.py +++ b/tests/test_backend_retries.py @@ -33,9 +33,10 @@ class _FakeSession: class _FakeClient: - def __init__(self) -> None: + 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 @@ -54,6 +55,8 @@ class _FakeClient: 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( @@ -90,6 +93,28 @@ async def test_non_transient_workspace_failure_does_not_retry() -> None: 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