Retry the Caido setup connection

This commit is contained in:
Alex Schapiro 2026-08-31 14:15:13 +00:00
parent 0e3d8e506e
commit 5ea5129b67
2 changed files with 111 additions and 53 deletions

View file

@ -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(

View file

@ -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)