diff --git a/containers/Dockerfile b/containers/Dockerfile index 16c1c540..4c54ddc2 100644 --- a/containers/Dockerfile +++ b/containers/Dockerfile @@ -24,6 +24,7 @@ RUN apt-get update && \ python3 python3-pip python3-dev python3-venv python3-setuptools \ golang-go \ net-tools dnsutils whois \ + file xxd \ jq parallel ripgrep grep \ less man-db procps htop \ iproute2 iputils-ping netcat-traditional \ @@ -192,6 +193,8 @@ RUN mkdir -p /workspace && chown -R pentester:pentester /workspace /app USER pentester RUN python3 -m venv /app/.venv && \ /app/.venv/bin/pip install --no-cache-dir caido-sdk-client && \ + /app/.venv/bin/pip install --no-cache-dir \ + requests httpx beautifulsoup4 lxml pyjwt cryptography && \ /app/.venv/bin/pip install --no-cache-dir -r /home/pentester/tools/jwt_tool/requirements.txt && \ printf '%s\n' \ '#!/bin/bash' \ diff --git a/containers/docker-entrypoint.sh b/containers/docker-entrypoint.sh index a12390df..d22cefdf 100644 --- a/containers/docker-entrypoint.sh +++ b/containers/docker-entrypoint.sh @@ -91,10 +91,13 @@ http_proxy=http://127.0.0.1:${CAIDO_PORT} https_proxy=http://127.0.0.1:${CAIDO_PORT} EOF -echo "source /etc/profile.d/proxy.sh" >> ~/.bashrc -echo "source /etc/profile.d/proxy.sh" >> ~/.zshrc +# Use POSIX `.` (not the bashism `source`) so these lines are safe when the rc +# files are read by a POSIX shell (e.g. `sh -lc`), which otherwise fails with +# "source: not found". `.` is understood by bash, zsh, and dash alike. +echo ". /etc/profile.d/proxy.sh" >> ~/.bashrc +echo ". /etc/profile.d/proxy.sh" >> ~/.zshrc -source /etc/profile.d/proxy.sh +. /etc/profile.d/proxy.sh echo "✅ System-wide proxy configuration complete" diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index b3d8e047..bb6f5f2d 100644 --- a/strix/agents/prompts/system_prompt.jinja +++ b/strix/agents/prompts/system_prompt.jinja @@ -169,8 +169,23 @@ EFFICIENCY TACTICS: - Run multiple scans in parallel when possible - Load the most relevant skill before starting a specialized testing workflow if doing so will improve accuracy, speed, or tool usage - Use `exec_command` for Python code: write reusable scripts to a file and - run them with `python3`. For one-off snippets, `python3 -c` or a - here-document is acceptable. + run them with `python3 script.py`. For one-off snippets, `python3 -c` or a + here-document is acceptable, but avoid deeply nested quotes/parentheses — if + a snippet needs complex quoting or is more than a few lines, write it to a + file first to prevent syntax errors. +- Before importing a third-party Python library, make sure it is installed. The + sandbox's `python3` runs inside a preconfigured virtualenv that ships + `requests`, `httpx`, `beautifulsoup4` (bs4), `lxml`, `pyjwt`, and + `cryptography`; for anything else prefer the stdlib or run `pip install ` + (it installs into that active venv) before importing, rather than letting the + script fail with `ModuleNotFoundError`. +- `exec_command` runs each command in a fresh non-interactive shell (plain + pipes, no TTY). To drive an interactive or long-running process with + `write_stdin` — REPLs, `ssh`/`nc`/`ftp`, `msfconsole`, or to send Ctrl-C — + you MUST start it with `exec_command(cmd="...", tty=true)` and then + `write_stdin(session_id=, chars="...")`. Calling `write_stdin` on a + default (non-TTY) command or on a process that has already exited fails with + "stdin is not available". - For Caido proxy automation inside Python, explicitly import from `caido_api`: `from caido_api import list_requests, view_request, repeat_request, list_sitemap, view_sitemap_entry, scope_rules` diff --git a/strix/core/runner.py b/strix/core/runner.py index e6a15dc3..9f1be995 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import contextlib import json import logging @@ -283,6 +284,11 @@ async def run_strix_scan( "coordinator": coordinator, "sandbox_session": bundle["session"], "caido_client": bundle["caido_client"], + # One shared Caido client is reused by every agent in the scan; its + # GraphQL transport is not safe for concurrent use. Child contexts + # are shallow copies (``dict(parent_ctx)``) so they inherit this + # same lock object, serializing all proxy calls scan-wide. + "caido_lock": asyncio.Lock(), "agent_id": root_id, "parent_id": None, "interactive": interactive, diff --git a/strix/skills/tooling/agent_browser.md b/strix/skills/tooling/agent_browser.md index 9ef810d6..db074e86 100644 --- a/strix/skills/tooling/agent_browser.md +++ b/strix/skills/tooling/agent_browser.md @@ -365,6 +365,23 @@ agent-browser dialog accept "text" # accept with prompt input agent-browser dialog dismiss # cancel ``` +## Readiness & recovery + +The first `agent-browser open` in a session launches the headless-Chrome +daemon; later commands reuse it. Distinguish the two failure modes and react +differently — do **not** blindly re-run the same failing command in a loop: + +- **Daemon / connection failure** (`Failed to connect`, `connection refused`, + socket missing, `browser not running`): the daemon isn't up or has died. Run + `agent-browser doctor` (add `--fix` if it reports repairable problems), then + re-open the page. Retrying the original command unchanged will keep failing. +- **Malformed command** (`Unknown command`, `Ref not found`, bad flag): fix the + command itself — re-snapshot for fresh refs, or correct the syntax. + +Invoke `agent-browser` directly through `exec_command`; there is no need to wrap +it in an extra `sh -c "..."` / `bash -lc "..."` layer, which only adds shell +quoting and startup-file pitfalls. + ## Diagnosing install issues If a command fails unexpectedly (`Unknown command`, `Failed to connect`, diff --git a/strix/skills/tooling/python.md b/strix/skills/tooling/python.md index d80aa69f..85a53d9f 100644 --- a/strix/skills/tooling/python.md +++ b/strix/skills/tooling/python.md @@ -92,10 +92,18 @@ For iterative exploit work, put code in a file: ## Installing extra packages -The sandbox's Python lives in `/app/.venv`. To add a one-off dependency -for an exploit script, use `uv` (already in the image and much faster -than pip): +The sandbox's Python lives in `/app/.venv`, and it is the active virtualenv +(`python3` / `pip` already resolve to it). The following common libraries are +**pre-installed** — import them directly, no install step needed: +`requests`, `httpx`, `beautifulsoup4` (`bs4`), `lxml`, `pyjwt` (`jwt`), +`cryptography`. + +To add a one-off dependency for an exploit script, use `uv` (already in the +image and much faster than pip): ```bash uv pip install --python /app/.venv/bin/python ``` + +Plain `pip install ` also works because the venv is active. Install +before you import, so scripts don't fail with `ModuleNotFoundError`. diff --git a/strix/tools/proxy/caido_api.py b/strix/tools/proxy/caido_api.py index 926d5c74..cf5aef23 100644 --- a/strix/tools/proxy/caido_api.py +++ b/strix/tools/proxy/caido_api.py @@ -21,6 +21,8 @@ from caido_sdk_client.types import ( if TYPE_CHECKING: + from collections.abc import Awaitable, Callable + from caido_sdk_client import Client as CaidoClient @@ -42,6 +44,19 @@ _SITEMAP_PAGE_SIZE = 30 _DEFAULT_CAIDO_URL = "http://127.0.0.1:48080" _CLIENT_CACHE: dict[str, Client] = {} +_CLIENT_LOCK = asyncio.Lock() + +# Substrings that mean the shared client's transport has died or is being used +# concurrently — recoverable by rebuilding the client and retrying once. +_CONNECTION_ERROR_MARKERS = ( + "transport is already connected", + "connector is closed", + "server disconnected", + "session is closed", + "cannot write to closing transport", + "connection reset", + "connection closed", +) _REQ_FIELD_MAP: dict[SortBy, tuple[str, str]] = { "timestamp": ("req", "created_at"), "host": ("req", "host"), @@ -81,19 +96,64 @@ def _login_as_guest() -> str: return str(payload["data"]["loginAsGuest"]["token"]["accessToken"]) -async def get_client() -> Client: - if client := _CLIENT_CACHE.get("default"): - return client - +async def _new_client() -> Client: token = await asyncio.to_thread(_login_as_guest) client = Client(caido_url(), auth=TokenAuthOptions(token=token)) await client.connect() - _CLIENT_CACHE["default"] = client return client +def _is_connection_error(exc: BaseException) -> bool: + message = str(exc).lower() + if any(marker in message for marker in _CONNECTION_ERROR_MARKERS): + return True + cause = exc.__cause__ or exc.__context__ + return cause is not None and cause is not exc and _is_connection_error(cause) + + +async def get_client() -> Client: + """Return the shared Caido client, creating it under a lock if needed. + + The lock prevents two concurrent callers from each building a client and + racing ``connect()`` on the same transport ("Transport is already + connected"). + """ + async with _CLIENT_LOCK: + client = _CLIENT_CACHE.get("default") + if client is None: + client = await _new_client() + _CLIENT_CACHE["default"] = client + return client + + +async def call_with_client[T](fn: Callable[[Client], Awaitable[T]]) -> T: + """Run ``fn`` against the shared client, serialized and reconnect-safe. + + The Caido GraphQL transport is not safe for concurrent use: two in-flight + requests race and raise "Transport is already connected". All proxy calls + are therefore serialized through ``_CLIENT_LOCK``. If the cached client's + transport has since died ("Connector is closed" / "Server disconnected"), + the stale client is rebuilt and the call retried once, instead of every + subsequent call in the run failing against a dead client. + """ + async with _CLIENT_LOCK: + client = _CLIENT_CACHE.get("default") + if client is None: + client = await _new_client() + _CLIENT_CACHE["default"] = client + try: + return await fn(client) + except Exception as exc: + if not _is_connection_error(exc): + raise + client = await _new_client() + _CLIENT_CACHE["default"] = client + return await fn(client) + + async def close_client() -> None: - client = _CLIENT_CACHE.pop("default", None) + async with _CLIENT_LOCK: + client = _CLIENT_CACHE.pop("default", None) if client is None: return await client.aclose() @@ -385,19 +445,23 @@ async def list_requests( sort_order: SortOrder = "desc", scope_id: str | None = None, ) -> Any: - return await list_requests_with_client( - await get_client(), - httpql_filter=httpql_filter, - first=first, - after=after, - sort_by=sort_by, - sort_order=sort_order, - scope_id=scope_id, + return await call_with_client( + lambda client: list_requests_with_client( + client, + httpql_filter=httpql_filter, + first=first, + after=after, + sort_by=sort_by, + sort_order=sort_order, + scope_id=scope_id, + ) ) async def view_request(request_id: str, *, part: RequestPart = "request") -> Any: - return await get_request_with_client(await get_client(), request_id, part=part) + return await call_with_client( + lambda client: get_request_with_client(client, request_id, part=part) + ) async def repeat_request( @@ -406,22 +470,26 @@ async def repeat_request( modifications: dict[str, Any] | None = None, ) -> dict[str, Any]: mods = modifications or {} - result = await get_request_with_client(await get_client(), request_id, part="request") - if result is None or result.request.raw is None: - raise ValueError(f"Request {request_id} not found") - original = result.request - raw_str = result.request.raw.decode("utf-8", errors="replace") - components = parse_raw_request(raw_str) - full_url = full_url_from_components(original, components, mods) - modified = apply_modifications(components, mods, full_url) - connection, raw = build_raw_request( - method=modified["method"], - url=modified["url"], - headers=modified["headers"], - body=modified["body"], - ) - return await replay_send_raw(await get_client(), raw=raw, connection=connection) + async def _run(client: CaidoClient) -> dict[str, Any]: + result = await get_request_with_client(client, request_id, part="request") + if result is None or result.request.raw is None: + raise ValueError(f"Request {request_id} not found") + + original = result.request + raw_str = result.request.raw.decode("utf-8", errors="replace") + components = parse_raw_request(raw_str) + full_url = full_url_from_components(original, components, mods) + modified = apply_modifications(components, mods, full_url) + connection, raw = build_raw_request( + method=modified["method"], + url=modified["url"], + headers=modified["headers"], + body=modified["body"], + ) + return await replay_send_raw(client, raw=raw, connection=connection) + + return await call_with_client(_run) async def scope_rules( @@ -432,7 +500,28 @@ async def scope_rules( scope_id: str | None = None, scope_name: str | None = None, ) -> Any: - client = await get_client() + async def _run(client: CaidoClient) -> Any: + return await _scope_rules_with_client( + client, + action, + allowlist=allowlist, + denylist=denylist, + scope_id=scope_id, + scope_name=scope_name, + ) + + return await call_with_client(_run) + + +async def _scope_rules_with_client( + client: CaidoClient, + action: ScopeAction, + *, + allowlist: list[str] | None = None, + denylist: list[str] | None = None, + scope_id: str | None = None, + scope_name: str | None = None, +) -> Any: if action == "list": result = await scope_list(client) elif action == "get": @@ -651,18 +740,20 @@ async def list_sitemap( page: int = 1, page_size: int = _SITEMAP_PAGE_SIZE, ) -> dict[str, Any]: - return await list_sitemap_with_client( - await get_client(), - scope_id=scope_id, - parent_id=parent_id, - depth=depth, - page=page, - page_size=page_size, + return await call_with_client( + lambda client: list_sitemap_with_client( + client, + scope_id=scope_id, + parent_id=parent_id, + depth=depth, + page=page, + page_size=page_size, + ) ) async def view_sitemap_entry(entry_id: str) -> dict[str, Any]: - return await view_sitemap_entry_with_client(await get_client(), entry_id) + return await call_with_client(lambda client: view_sitemap_entry_with_client(client, entry_id)) __all__ = [ @@ -671,6 +762,7 @@ __all__ = [ "SitemapDepth", "SortBy", "SortOrder", + "call_with_client", "close_client", "get_client", "list_requests", diff --git a/strix/tools/proxy/tools.py b/strix/tools/proxy/tools.py index db8a6b24..1610dfde 100644 --- a/strix/tools/proxy/tools.py +++ b/strix/tools/proxy/tools.py @@ -2,6 +2,8 @@ from __future__ import annotations +import asyncio +import contextlib import dataclasses import json import logging @@ -44,6 +46,23 @@ def _ctx_client(ctx: RunContextWrapper) -> Client | None: return inner.get("caido_client") +def _ctx_lock(ctx: RunContextWrapper) -> contextlib.AbstractAsyncContextManager[None]: + """Return the scan-wide lock serializing access to the shared Caido client. + + All agents in a scan share one ``caido_client`` whose GraphQL transport is + not concurrency-safe (parallel calls raise "Transport is already + connected", and racing session teardown yields "Connector is closed" / + "Server disconnected"). Holding this lock around every proxy call serializes + them. Falls back to a no-op context when no lock is present (e.g. standalone + tool invocation outside a scan run). + """ + inner = ctx.context if isinstance(ctx.context, dict) else {} + lock = inner.get("caido_lock") + if isinstance(lock, asyncio.Lock): + return lock + return contextlib.nullcontext() + + def _to_tool_json(value: Any) -> Any: """Recursively convert SDK dataclasses/Pydantic objects to tool JSON values.""" if value is None or isinstance(value, str | int | float | bool): @@ -83,6 +102,39 @@ def _err(name: str, exc: Exception) -> str: ) +_HTTPQL_HINT = ( + "HTTPQL syntax: quote string values and leave integers unquoted; combine " + "terms with AND / OR (there is no NOT). Numeric fields (resp.code, req.port, " + "id, roundtrip) use eq/ne/gt/gte/lt/lte; text/byte fields (req.host, req.path, " + "req.method, req.raw, resp.raw) use cont/ncont/eq/ne/like/nlike/regex/nregex. " + "Example: 'resp.code.gte:200 AND resp.code.lt:300 AND req.host.cont:\"api\"'." +) + + +def _is_httpql_error(exc: Exception) -> bool: + message = str(exc).lower() + return "httpql" in message or ("filter" in message and "pars" in message) + + +def _httpql_error(exc: Exception, httpql_filter: str | None) -> str: + """Return an actionable error for a rejected HTTPQL filter. + + Preserves Caido's exact parser message and echoes the offending query so + the agent can self-correct instead of retrying the same broken filter. + """ + logger.info("list_requests rejected HTTPQL filter %r: %s", httpql_filter, exc) + return json.dumps( + { + "success": False, + "error": f"Invalid HTTPQL filter: {exc}", + "httpql_filter": httpql_filter, + "hint": _HTTPQL_HINT, + }, + ensure_ascii=False, + default=str, + ) + + @function_tool(timeout=120) async def list_requests( ctx: RunContextWrapper, @@ -146,15 +198,16 @@ async def list_requests( return _no_client() try: - connection = await caido_api.list_requests_with_client( - client, - httpql_filter=httpql_filter, - first=first, - after=after, - sort_by=sort_by, - sort_order=sort_order, - scope_id=scope_id, - ) + async with _ctx_lock(ctx): + connection = await caido_api.list_requests_with_client( + client, + httpql_filter=httpql_filter, + first=first, + after=after, + sort_by=sort_by, + sort_order=sort_order, + scope_id=scope_id, + ) entries = [] for edge in connection.edges: @@ -207,6 +260,8 @@ async def list_requests( default=str, ) except Exception as exc: # noqa: BLE001 + if httpql_filter and _is_httpql_error(exc): + return _httpql_error(exc, httpql_filter) return _err("list_requests", exc) @@ -249,7 +304,8 @@ async def view_request( return _no_client() try: - result = await caido_api.get_request_with_client(client, request_id, part=part) + async with _ctx_lock(ctx): + result = await caido_api.get_request_with_client(client, request_id, part=part) if result is None: return json.dumps( {"success": False, "error": f"Request {request_id} not found"}, @@ -365,26 +421,27 @@ async def repeat_request( mods = modifications or {} try: - result = await caido_api.get_request_with_client(client, request_id, part="request") - if result is None or result.request.raw is None: - return json.dumps( - {"success": False, "error": f"Request {request_id} not found"}, - ensure_ascii=False, - default=str, - ) + async with _ctx_lock(ctx): + result = await caido_api.get_request_with_client(client, request_id, part="request") + if result is None or result.request.raw is None: + return json.dumps( + {"success": False, "error": f"Request {request_id} not found"}, + ensure_ascii=False, + default=str, + ) - original = result.request - raw_str = result.request.raw.decode("utf-8", errors="replace") - components = caido_api.parse_raw_request(raw_str) - full_url = caido_api.full_url_from_components(original, components, mods) - modified = caido_api.apply_modifications(components, mods, full_url) - connection, raw = caido_api.build_raw_request( - method=modified["method"], - url=modified["url"], - headers=modified["headers"], - body=modified["body"], - ) - replay = await caido_api.replay_send_raw(client, raw=raw, connection=connection) + original = result.request + raw_str = result.request.raw.decode("utf-8", errors="replace") + components = caido_api.parse_raw_request(raw_str) + full_url = caido_api.full_url_from_components(original, components, mods) + modified = caido_api.apply_modifications(components, mods, full_url) + connection, raw = caido_api.build_raw_request( + method=modified["method"], + url=modified["url"], + headers=modified["headers"], + body=modified["body"], + ) + replay = await caido_api.replay_send_raw(client, raw=raw, connection=connection) return _format_replay_tool_result(replay) except Exception as exc: # noqa: BLE001 return _err("repeat_request", exc) @@ -441,13 +498,14 @@ async def list_sitemap( if client is None: return _no_client() try: - payload = await caido_api.list_sitemap_with_client( - client, - scope_id=scope_id, - parent_id=parent_id, - depth=depth, - page=page, - ) + async with _ctx_lock(ctx): + payload = await caido_api.list_sitemap_with_client( + client, + scope_id=scope_id, + parent_id=parent_id, + depth=depth, + page=page, + ) return json.dumps(payload, ensure_ascii=False, default=str) except Exception as exc: # noqa: BLE001 return _err("list_sitemap", exc) @@ -472,7 +530,8 @@ async def view_sitemap_entry( if client is None: return _no_client() try: - payload = await caido_api.view_sitemap_entry_with_client(client, entry_id) + async with _ctx_lock(ctx): + payload = await caido_api.view_sitemap_entry_with_client(client, entry_id) return json.dumps(payload, ensure_ascii=False, default=str) except Exception as exc: # noqa: BLE001 return _err("view_sitemap_entry", exc) @@ -529,68 +588,75 @@ async def scope_rules( return _no_client() try: - if action == "list": - scopes = await caido_api.scope_list(client) - return json.dumps( - {"success": True, "scopes": [_to_tool_json(s) for s in scopes]}, - ensure_ascii=False, - default=str, - ) - if action == "get": + async with _ctx_lock(ctx): + if action == "list": + scopes = await caido_api.scope_list(client) + return json.dumps( + {"success": True, "scopes": [_to_tool_json(s) for s in scopes]}, + ensure_ascii=False, + default=str, + ) + if action == "get": + if not scope_id: + return json.dumps( + {"success": False, "error": "Scope_id is required for action='get'"}, + ensure_ascii=False, + default=str, + ) + scope = await caido_api.scope_get(client, scope_id) + return json.dumps( + {"success": True, "scope": _to_tool_json(scope)}, + ensure_ascii=False, + default=str, + ) + if action == "create": + if not scope_name: + return json.dumps( + {"success": False, "error": "Scope_name is required for action='create'"}, + ensure_ascii=False, + default=str, + ) + scope = await caido_api.scope_create( + client, name=scope_name, allowlist=allowlist, denylist=denylist + ) + return json.dumps( + {"success": True, "scope": _to_tool_json(scope)}, + ensure_ascii=False, + default=str, + ) + if action == "update": + if not scope_id or not scope_name: + return json.dumps( + { + "success": False, + "error": "Scope_id and scope_name are required for action='update'", + }, + ensure_ascii=False, + default=str, + ) + scope = await caido_api.scope_update( + client, scope_id, name=scope_name, allowlist=allowlist, denylist=denylist + ) + return json.dumps( + {"success": True, "scope": _to_tool_json(scope)}, + ensure_ascii=False, + default=str, + ) if not scope_id: return json.dumps( - {"success": False, "error": "Scope_id is required for action='get'"}, + {"success": False, "error": "Scope_id is required for action='delete'"}, ensure_ascii=False, default=str, ) - scope = await caido_api.scope_get(client, scope_id) + await caido_api.scope_delete(client, scope_id) return json.dumps( - {"success": True, "scope": _to_tool_json(scope)}, ensure_ascii=False, default=str - ) - if action == "create": - if not scope_name: - return json.dumps( - {"success": False, "error": "Scope_name is required for action='create'"}, - ensure_ascii=False, - default=str, - ) - scope = await caido_api.scope_create( - client, name=scope_name, allowlist=allowlist, denylist=denylist - ) - return json.dumps( - {"success": True, "scope": _to_tool_json(scope)}, ensure_ascii=False, default=str - ) - if action == "update": - if not scope_id or not scope_name: - return json.dumps( - { - "success": False, - "error": "Scope_id and scope_name are required for action='update'", - }, - ensure_ascii=False, - default=str, - ) - scope = await caido_api.scope_update( - client, scope_id, name=scope_name, allowlist=allowlist, denylist=denylist - ) - return json.dumps( - {"success": True, "scope": _to_tool_json(scope)}, ensure_ascii=False, default=str - ) - if not scope_id: - return json.dumps( - {"success": False, "error": "Scope_id is required for action='delete'"}, + { + "success": True, + "deleted": scope_id, + "message": f"Scope {scope_id} deleted", + }, ensure_ascii=False, default=str, ) - await caido_api.scope_delete(client, scope_id) - return json.dumps( - { - "success": True, - "deleted": scope_id, - "message": f"Scope {scope_id} deleted", - }, - ensure_ascii=False, - default=str, - ) except Exception as exc: # noqa: BLE001 return _err("scope_rules", exc) diff --git a/strix/tools/shell/README.md b/strix/tools/shell/README.md index 204c9256..b842e317 100644 --- a/strix/tools/shell/README.md +++ b/strix/tools/shell/README.md @@ -5,6 +5,23 @@ invocation the agent makes (nmap, ffuf, agent-browser, python3, …) goes through `exec_command`. `write_stdin` streams input to a still-running process started by an earlier `exec_command` (for interactive prompts). +## `write_stdin` requires a TTY-backed process + +`exec_command` runs each command in a fresh **non-interactive** shell (plain +pipes, no TTY) by default. `write_stdin` only works against a process that is +still running **and** was started with a PTY. The canonical sequence is: + +```text +exec_command(cmd="python3", tty=true) # start a PTY-backed process +write_stdin(session_id=, chars="print(1)\n") +``` + +Calling `write_stdin` on a command started with the default `tty=false`, or on +a process that has already exited, fails with +`stdin is not available for this process. Start the command with 'tty=true' in +'exec_command' before using 'write_stdin'.` Use `tty=true` for REPLs, +`ssh`/`nc`/`ftp`, `msfconsole`, or to deliver a Ctrl-C to a long-running job. + - **Implementation:** `agents.sandbox.capabilities.tools.shell_tool.ShellTool` (in the upstream `agents` SDK) - **Wired in:** `strix/agents/factory.py` — added per-run via the SDK diff --git a/tests/test_proxy_client.py b/tests/test_proxy_client.py new file mode 100644 index 00000000..6022f5a2 --- /dev/null +++ b/tests/test_proxy_client.py @@ -0,0 +1,206 @@ +"""Tests for the shared Caido client lifecycle and proxy error handling. + +Covers the concurrency/reconnect guarantees of ``caido_api.call_with_client`` +(the sandbox-imported path) and the host-side helpers in ``proxy.tools`` +(scan-wide lock + actionable HTTPQL errors). +""" + +from __future__ import annotations + +import asyncio +import contextlib +import json +from typing import TYPE_CHECKING, Any, cast + +import pytest + +from strix.tools.proxy import caido_api, tools + + +if TYPE_CHECKING: + from collections.abc import Iterator + + +class _FakeClient: + def __init__(self, name: str) -> None: + self.name = name + + +@pytest.fixture(autouse=True) +def _clear_cache() -> Iterator[None]: + caido_api._CLIENT_CACHE.clear() + yield + caido_api._CLIENT_CACHE.clear() + + +async def test_call_with_client_reuses_cached_client(monkeypatch: pytest.MonkeyPatch) -> None: + cached = _FakeClient("cached") + caido_api._CLIENT_CACHE["default"] = cached + + async def _new() -> Any: + raise AssertionError("_new_client must not run when a client is cached") + + monkeypatch.setattr(caido_api, "_new_client", _new) + + seen: dict[str, Any] = {} + + async def fn(client: Any) -> str: + seen["client"] = client + return "ok" + + assert await caido_api.call_with_client(fn) == "ok" + assert seen["client"] is cached + + +async def test_call_with_client_creates_and_caches_when_empty( + monkeypatch: pytest.MonkeyPatch, +) -> None: + created = _FakeClient("fresh") + + async def _new() -> Any: + return created + + monkeypatch.setattr(caido_api, "_new_client", _new) + + seen: dict[str, Any] = {} + + async def fn(client: Any) -> str: + seen["client"] = client + return "ok" + + assert await caido_api.call_with_client(fn) == "ok" + assert seen["client"] is created + assert caido_api._CLIENT_CACHE["default"] is created + + +async def test_failed_init_does_not_poison_cache(monkeypatch: pytest.MonkeyPatch) -> None: + async def _new() -> Any: + raise ConnectionRefusedError("caido not up yet") + + monkeypatch.setattr(caido_api, "_new_client", _new) + + async def fn(_client: Any) -> str: + return "unreachable" + + with pytest.raises(ConnectionRefusedError): + await caido_api.call_with_client(fn) + assert "default" not in caido_api._CLIENT_CACHE + + +async def test_call_with_client_reconnects_on_dead_transport( + monkeypatch: pytest.MonkeyPatch, +) -> None: + dead = _FakeClient("dead") + fresh = _FakeClient("fresh") + caido_api._CLIENT_CACHE["default"] = dead + + new_calls = {"n": 0} + + async def _new() -> Any: + new_calls["n"] += 1 + return fresh + + monkeypatch.setattr(caido_api, "_new_client", _new) + + attempts: list[Any] = [] + + async def fn(client: Any) -> str: + attempts.append(client) + if len(attempts) == 1: + raise RuntimeError("Transport is already connected") + return "ok" + + assert await caido_api.call_with_client(fn) == "ok" + assert attempts == [dead, fresh] + assert new_calls["n"] == 1 + assert caido_api._CLIENT_CACHE["default"] is fresh + + +async def test_call_with_client_does_not_retry_application_errors( + monkeypatch: pytest.MonkeyPatch, +) -> None: + cached = _FakeClient("cached") + caido_api._CLIENT_CACHE["default"] = cached + + async def _new() -> Any: + raise AssertionError("deterministic errors must not trigger a reconnect") + + monkeypatch.setattr(caido_api, "_new_client", _new) + + calls = {"n": 0} + + async def fn(_client: Any) -> str: + calls["n"] += 1 + raise ValueError("Invalid HTTPQL filter") + + with pytest.raises(ValueError, match="Invalid HTTPQL"): + await caido_api.call_with_client(fn) + assert calls["n"] == 1 + assert caido_api._CLIENT_CACHE["default"] is cached + + +async def test_call_with_client_serializes_concurrent_calls( + monkeypatch: pytest.MonkeyPatch, +) -> None: + caido_api._CLIENT_CACHE["default"] = _FakeClient("shared") + + async def _new() -> Any: + raise AssertionError("no reconnect expected") + + monkeypatch.setattr(caido_api, "_new_client", _new) + + state = {"active": 0, "max": 0} + + async def fn(_client: Any) -> str: + state["active"] += 1 + state["max"] = max(state["max"], state["active"]) + await asyncio.sleep(0.01) + state["active"] -= 1 + return "ok" + + await asyncio.gather(*(caido_api.call_with_client(fn) for _ in range(6))) + assert state["max"] == 1 + + +def test_is_connection_error_matches_markers_and_causes() -> None: + assert caido_api._is_connection_error(RuntimeError("Transport is already connected")) + assert caido_api._is_connection_error(RuntimeError("Connector is closed")) + assert caido_api._is_connection_error(RuntimeError("Server disconnected")) + assert not caido_api._is_connection_error(ValueError("Invalid HTTPQL filter")) + + nested = RuntimeError("wrapper") + nested.__cause__ = RuntimeError("connection reset by peer") + assert caido_api._is_connection_error(nested) + + +class _Ctx: + def __init__(self, context: Any) -> None: + self.context = context + + +def test_ctx_lock_returns_lock_when_present() -> None: + lock = asyncio.Lock() + got = tools._ctx_lock(cast("Any", _Ctx({"caido_lock": lock}))) + assert got is lock + + +def test_ctx_lock_falls_back_to_noop_without_lock() -> None: + got = tools._ctx_lock(cast("Any", _Ctx({}))) + assert isinstance(got, contextlib.nullcontext) + got_non_dict = tools._ctx_lock(cast("Any", _Ctx(None))) + assert isinstance(got_non_dict, contextlib.nullcontext) + + +def test_is_httpql_error_detection() -> None: + assert tools._is_httpql_error(RuntimeError("HTTPQL parse error at column 4")) + assert tools._is_httpql_error(RuntimeError("failed to parse filter")) + assert not tools._is_httpql_error(RuntimeError("Transport is already connected")) + + +def test_httpql_error_preserves_message_and_query() -> None: + exc = RuntimeError("HTTPQL parse error: unexpected token at column 12") + payload = json.loads(tools._httpql_error(exc, 'resp.code.eq:"200"')) + assert payload["success"] is False + assert "unexpected token at column 12" in payload["error"] + assert payload["httpql_filter"] == 'resp.code.eq:"200"' + assert "AND / OR" in payload["hint"]