diff --git a/strix/runtime/caido_bootstrap.py b/strix/runtime/caido_bootstrap.py index 9f93cece..ca143ce4 100644 --- a/strix/runtime/caido_bootstrap.py +++ b/strix/runtime/caido_bootstrap.py @@ -80,52 +80,67 @@ async def _login_as_guest( raise RuntimeError(f"loginAsGuest failed after {attempts} attempts: {last_err}") +async def _find_sandbox_project(client: Client) -> str | None: + """Look for a project a create that timed out client-side may have left behind.""" + with contextlib.suppress(Exception): + projects = [item for item in await client.project.list() if item.name == "sandbox"] + if projects: + return str(max(projects, key=lambda item: item.id).id) + return None + + async def _setup_project(host_url: str, access_token: str) -> None: - """Connect with a longer deadline and select the sandbox project.""" + """Select the sandbox project, retrying the whole connect/create/select sequence. + + Each attempt gets a fresh client with a deadline well past the SDK default: + these mutations are slow on a cold Caido, while the long-lived client the + scan uses keeps the short default so a traffic poll cannot stall on it. + Until a project is selected Caido answers every proxied request with a 500, + so giving up here costs the whole run, not just the traffic capture. + """ from caido_sdk_client import Client, TokenAuthOptions from caido_sdk_client.types import CreateProjectOptions - client = Client( - host_url, - auth=TokenAuthOptions(token=access_token), - timeout_ms=_PROJECT_SETUP_TIMEOUT_MS, - ) - try: - await client.connect() - project = None - last_exc: Exception | None = None - for i in range(1, _PROJECT_SETUP_ATTEMPTS + 1): - try: - if project is None: - project = await client.project.create( + project_id: str | None = None + last_exc: Exception | None = None + for attempt in range(1, _PROJECT_SETUP_ATTEMPTS + 1): + client = Client( + host_url, + auth=TokenAuthOptions(token=access_token), + timeout_ms=_PROJECT_SETUP_TIMEOUT_MS, + ) + try: + await client.connect() + if project_id is None: + try: + created = await client.project.create( CreateProjectOptions(name="sandbox", temporary=True), ) - await client.project.select(project.id) - except Exception as exc: # noqa: BLE001 - last_exc = exc - if project is None: - with contextlib.suppress(Exception): - projects = await client.project.list() - sandbox_projects = [item for item in projects if item.name == "sandbox"] - if sandbox_projects: - project = max(sandbox_projects, key=lambda item: item.id) - logger.warning( - "Caido project setup attempt %d/%d failed: %s", - i, - _PROJECT_SETUP_ATTEMPTS, - exc, - ) - if i < _PROJECT_SETUP_ATTEMPTS: - await asyncio.sleep(min(2.0 * i, 8.0)) - else: - logger.info("Caido project selected: %s", project.id) - return - raise RuntimeError( - f"Caido project setup failed after {_PROJECT_SETUP_ATTEMPTS} attempts" - ) from last_exc - finally: - with contextlib.suppress(Exception): - await client.aclose() + except Exception: + # A create that timed out client-side may still have landed. + project_id = await _find_sandbox_project(client) + raise + project_id = created.id + await client.project.select(project_id) + except Exception as exc: # noqa: BLE001 + last_exc = exc + logger.warning( + "Caido project setup attempt %d/%d failed: %s", + attempt, + _PROJECT_SETUP_ATTEMPTS, + exc, + ) + if attempt < _PROJECT_SETUP_ATTEMPTS: + await asyncio.sleep(min(2.0 * attempt, 8.0)) + else: + logger.info("Caido project selected: %s", project_id) + return + finally: + with contextlib.suppress(Exception): + await client.aclose() + raise RuntimeError( + f"Caido project setup failed after {_PROJECT_SETUP_ATTEMPTS} attempts" + ) from last_exc async def bootstrap_caido( diff --git a/tests/test_caido_bootstrap.py b/tests/test_caido_bootstrap.py index f4c68c88..0a4680e6 100644 --- a/tests/test_caido_bootstrap.py +++ b/tests/test_caido_bootstrap.py @@ -116,6 +116,18 @@ def _install_sdk( return factory +def _setup_clients( + project: _FakeProjectSDK, + count: int, + *, + connect_errors: list[BaseException | None] | None = None, +) -> list[_FakeClient]: + """One client per setup attempt, all sharing the same server-side project state.""" + errors: list[BaseException | None] = list(connect_errors or []) + errors += [None] * (count - len(errors)) + return [_FakeClient(errors[i], project=project) for i in range(count)] + + async def _bootstrap_expecting( monkeypatch: pytest.MonkeyPatch, error: BaseException ) -> _FakeClient: @@ -150,9 +162,9 @@ async def test_select_retries_without_creating_another_project( monkeypatch: pytest.MonkeyPatch, ) -> None: setup_project = _FakeProjectSDK(select_errors=[RuntimeError("not ready")]) - setup_client = _FakeClient(project=setup_project) + setup_clients = _setup_clients(setup_project, 2) returned_client = _FakeClient() - factory = _install_sdk(monkeypatch, [setup_client, returned_client]) + factory = _install_sdk(monkeypatch, [*setup_clients, returned_client]) sleep_calls: list[float] = [] async def _sleep(delay: float) -> None: @@ -172,9 +184,10 @@ async def test_select_retries_without_creating_another_project( assert result is returned_client assert setup_project.create_calls == 1 assert setup_project.selected_ids == ["created", "created"] - assert setup_client.closed + assert all(client.closed for client in setup_clients) assert factory.calls[0][1]["timeout_ms"] == 45_000 - assert factory.calls[1][1].get("timeout_ms") is None + assert factory.calls[1][1]["timeout_ms"] == 45_000 + assert factory.calls[2][1].get("timeout_ms") is None assert sleep_calls == [2.0] @@ -189,9 +202,9 @@ async def test_create_failure_reuses_the_most_recent_sandbox_project( _FakeProject("other", name="other"), ], ) - setup_client = _FakeClient(project=setup_project) + setup_clients = _setup_clients(setup_project, 2) returned_client = _FakeClient() - _install_sdk(monkeypatch, [setup_client, returned_client]) + _install_sdk(monkeypatch, [*setup_clients, returned_client]) async def _sleep(_delay: float) -> None: pass @@ -208,7 +221,7 @@ async def test_create_failure_reuses_the_most_recent_sandbox_project( assert setup_project.create_calls == 1 assert setup_project.list_calls == 1 assert setup_project.selected_ids == ["project-2"] - assert setup_client.closed + assert all(client.closed for client in setup_clients) async def test_project_setup_failure_chains_last_error_and_closes_setup_client( @@ -220,8 +233,8 @@ async def test_project_setup_failure_chains_last_error_and_closes_setup_client( RuntimeError("last"), ] setup_project = _FakeProjectSDK(select_errors=errors) - setup_client = _FakeClient(project=setup_project) - _install_sdk(monkeypatch, [setup_client]) + setup_clients = _setup_clients(setup_project, 3) + _install_sdk(monkeypatch, setup_clients) async def _sleep(_delay: float) -> None: pass @@ -239,7 +252,7 @@ async def test_project_setup_failure_chains_last_error_and_closes_setup_client( ) assert exc_info.value.__cause__ is errors[-1] - assert setup_client.closed + assert all(client.closed for client in setup_clients) assert setup_project.create_calls == 1 assert setup_project.selected_ids == ["created", "created", "created"] @@ -248,8 +261,8 @@ async def test_cancelled_project_setup_is_not_retried( monkeypatch: pytest.MonkeyPatch, ) -> None: setup_project = _FakeProjectSDK(select_errors=[asyncio.CancelledError()]) - setup_client = _FakeClient(project=setup_project) - _install_sdk(monkeypatch, [setup_client]) + setup_clients = _setup_clients(setup_project, 1) + _install_sdk(monkeypatch, setup_clients) sleep_calls: list[float] = [] async def _sleep(delay: float) -> None: @@ -270,4 +283,34 @@ async def test_cancelled_project_setup_is_not_retried( assert setup_project.create_calls == 1 assert setup_project.selected_ids == ["created"] assert sleep_calls == [] - assert setup_client.closed + assert all(client.closed for client in setup_clients) + + +async def test_setup_connect_failure_is_retried(monkeypatch: pytest.MonkeyPatch) -> None: + setup_project = _FakeProjectSDK() + setup_clients = _setup_clients( + setup_project, + 2, + connect_errors=[RuntimeError("gateway not ready")], + ) + returned_client = _FakeClient() + _install_sdk(monkeypatch, [*setup_clients, returned_client]) + sleep_calls: list[float] = [] + + async def _sleep(delay: float) -> None: + sleep_calls.append(delay) + + monkeypatch.setattr("strix.runtime.caido_bootstrap.asyncio.sleep", _sleep) + + result = await bootstrap_caido( + _FakeSession(), # type: ignore[arg-type] + host_url="http://host", + container_url="http://container", + ) + + assert result is returned_client + assert setup_project.create_calls == 1 + assert setup_project.selected_ids == ["created"] + assert setup_project.list_calls == 0 + assert sleep_calls == [2.0] + assert all(client.closed for client in setup_clients)