From 0e4307d6257a3674bf6c2d68912af76f26e47ef7 Mon Sep 17 00:00:00 2001 From: Shanyu Thibaut Juneja Date: Wed, 11 Mar 2026 16:33:14 -0700 Subject: [PATCH] harden: authentication for cdp --- containers/Dockerfile | 3 +- containers/cdp-auth-proxy.py | 193 ++++++++++ containers/docker-entrypoint.sh | 32 +- pyproject.toml | 8 + .../tool_components/browser_renderer.py | 333 ++++++++++++++---- strix/runtime/tool_server.py | 48 ++- strix/tools/browser/browser_actions.py | 99 +++++- .../tools/browser/browser_actions_schema.xml | 32 ++ strix/tools/browser/browser_manager.py | 143 ++++++-- 9 files changed, 744 insertions(+), 147 deletions(-) create mode 100644 containers/cdp-auth-proxy.py diff --git a/containers/Dockerfile b/containers/Dockerfile index d385c7e7..721ec8e7 100644 --- a/containers/Dockerfile +++ b/containers/Dockerfile @@ -216,7 +216,8 @@ RUN echo 'export PATH="/home/pentester/go/bin:/home/pentester/.local/bin:/home/p USER root COPY containers/docker-entrypoint.sh /usr/local/bin/docker-entrypoint.sh COPY containers/healthcheck.sh /usr/local/bin/healthcheck.sh -RUN chmod +x /usr/local/bin/docker-entrypoint.sh /usr/local/bin/healthcheck.sh +COPY containers/cdp-auth-proxy.py /usr/local/bin/cdp-auth-proxy.py +RUN chmod +x /usr/local/bin/docker-entrypoint.sh /usr/local/bin/healthcheck.sh /usr/local/bin/cdp-auth-proxy.py HEALTHCHECK --interval=15s --timeout=5s --start-period=60s --retries=3 \ CMD healthcheck.sh diff --git a/containers/cdp-auth-proxy.py b/containers/cdp-auth-proxy.py new file mode 100644 index 00000000..f11e6327 --- /dev/null +++ b/containers/cdp-auth-proxy.py @@ -0,0 +1,193 @@ +#!/usr/bin/env python3 +"""CDP authentication proxy. + +Replaces the raw socat forwarder with an auth-aware TCP proxy. Every +connection (HTTP *and* WebSocket upgrade) must carry a valid token — +either as an ``Authorization: Bearer `` header or a +``?token=`` query parameter. The latter is needed for WebSocket +clients (e.g. Playwright) that don't support custom headers on the +upgrade request. + +Auth credentials are stripped before the request is forwarded to +Chromium so the upstream sees a clean, standard CDP request. + +Environment variables: + TOOL_SERVER_TOKEN — required, shared secret + CDP_PORT — listen port (default 9222) + CDP_INTERNAL_PORT — upstream Chromium port (default 19222) +""" + +from __future__ import annotations + +import asyncio +import os +import re +import sys + + +TOKEN: str = os.environ["TOOL_SERVER_TOKEN"] +LISTEN_PORT: int = int(os.environ.get("CDP_PORT", "9222")) +UPSTREAM_PORT: int = int(os.environ.get("CDP_INTERNAL_PORT", "19222")) + +# Maximum bytes to buffer while looking for the end-of-headers marker. +# Prevents memory exhaustion from slow-loris / oversized-header attacks. +_MAX_HEADER_SIZE: int = 16 * 1024 # 16 KiB — plenty for CDP requests + + +# --------------------------------------------------------------------------- +# Auth helpers +# --------------------------------------------------------------------------- + + +def _check_auth(header_bytes: bytes) -> bool: + text = header_bytes.decode("latin-1", errors="replace") + + # 1) Authorization: Bearer + m = re.search(r"(?i)Authorization:\s*Bearer\s+(\S+)", text) + if m and m.group(1) == TOKEN: + return True + + # 2) ?token= or &token= query param + m = re.search(r"[?&]token=([^&\s]+)", text) + return bool(m and m.group(1) == TOKEN) + + +def _sanitize(header_bytes: bytes) -> bytes: + """Strip auth credentials before forwarding upstream.""" + text = header_bytes.decode("latin-1", errors="replace") + + # Remove Authorization header line + text = re.sub(r"(?im)^Authorization:[^\r\n]*\r\n", "", text) + + # Remove token query param — handle all positions: + # ?token=val&rest → ?rest (first of many) + # ?token=val → (nothing) (sole param) + # &token=val → (nothing) (not first) + text = re.sub(r"\?token=[^&\s]+&", "?", text) + text = re.sub(r"\?token=[^&\s]+", "", text) + text = re.sub(r"&token=[^&\s]+", "", text) + + return text.encode("latin-1") + + +_REJECT = ( + b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 12\r\nConnection: close\r\n\r\nUnauthorized" +) + + +# --------------------------------------------------------------------------- +# TCP proxy +# --------------------------------------------------------------------------- + + +async def _pipe( + reader: asyncio.StreamReader, + writer: asyncio.StreamWriter, +) -> None: + """Forward bytes until EOF or error.""" + try: + while True: + data = await reader.read(65536) + if not data: + break + writer.write(data) + await writer.drain() + except (ConnectionResetError, BrokenPipeError, OSError): + pass + finally: + try: + if writer.can_write_eof(): + writer.write_eof() + except OSError: + pass + + +async def _handle( + client_reader: asyncio.StreamReader, + client_writer: asyncio.StreamWriter, +) -> None: + up_writer: asyncio.StreamWriter | None = None + try: + # --- Read HTTP headers (up to the blank line) --- + buf = b"" + while b"\r\n\r\n" not in buf: + chunk = await asyncio.wait_for(client_reader.read(8192), timeout=10) + if not chunk: + return + buf += chunk + if len(buf) > _MAX_HEADER_SIZE: + client_writer.write( + b"HTTP/1.1 431 Request Header Fields Too Large\r\nConnection: close\r\n\r\n" + ) + await client_writer.drain() + return + + sep = buf.index(b"\r\n\r\n") + 4 + headers = buf[:sep] + remainder = buf[sep:] + + # --- Authenticate --- + if not _check_auth(headers): + client_writer.write(_REJECT) + await client_writer.drain() + return + + # --- Strip credentials & forward --- + headers = _sanitize(headers) + + up_reader, up_writer = await asyncio.open_connection("127.0.0.1", UPSTREAM_PORT) + up_writer.write(headers) + if remainder: + up_writer.write(remainder) + await up_writer.drain() + + # --- Bidirectional pipe --- + done, pending = await asyncio.wait( + [ + asyncio.create_task(_pipe(client_reader, up_writer)), + asyncio.create_task(_pipe(up_reader, client_writer)), + ], + return_when=asyncio.FIRST_COMPLETED, + ) + for t in pending: + t.cancel() + + except (ConnectionResetError, BrokenPipeError, OSError, TimeoutError): + pass + finally: + for w in (client_writer, up_writer): + if w is not None: + with _suppress_os(): + w.close() + + +class _suppress_os: # noqa: N801 + """Tiny context manager — cheaper than contextlib.suppress(OSError).""" + + def __enter__(self) -> None: + pass + + def __exit__(self, *exc: object) -> bool: + return isinstance(exc[1], OSError) if exc[1] is not None else False + + +# --------------------------------------------------------------------------- +# Main +# --------------------------------------------------------------------------- + + +async def _main() -> None: + server = await asyncio.start_server(_handle, "0.0.0.0", LISTEN_PORT) # nosec B104 + print( + f"CDP auth proxy: 0.0.0.0:{LISTEN_PORT} -> 127.0.0.1:{UPSTREAM_PORT}", + flush=True, + ) + async with server: + await server.serve_forever() + + +if __name__ == "__main__": + try: + asyncio.run(_main()) + except KeyboardInterrupt: + sys.exit(0) diff --git a/containers/docker-entrypoint.sh b/containers/docker-entrypoint.sh index 4fdfb90d..548e5560 100644 --- a/containers/docker-entrypoint.sh +++ b/containers/docker-entrypoint.sh @@ -153,7 +153,7 @@ echo "✅ CA added to browser trust store" CDP_PORT="${BROWSER_CDP_PORT:-9222}" # Chromium always binds CDP to 127.0.0.1 regardless of --remote-debugging-address. -# We launch it on an internal port and use socat to expose it on 0.0.0.0. +# We launch it on an internal port and use an auth proxy to expose it on 0.0.0.0. CDP_INTERNAL_PORT=19222 CHROMIUM_BIN="" CHROMIUM_RESTART_COUNT=0 @@ -167,8 +167,8 @@ if [ -z "$CHROMIUM_BIN" ]; then fi # --------------------------------------------------------------------------- -# start_chromium: launches Chromium + socat, waits for CDP readiness. -# Sets CHROMIUM_PID and SOCAT_PID on success. +# start_chromium: launches Chromium + auth proxy, waits for CDP readiness. +# Sets CHROMIUM_PID and CDP_PROXY_PID on success. # --------------------------------------------------------------------------- start_chromium() { if [ -z "$CHROMIUM_BIN" ]; then @@ -215,30 +215,30 @@ start_chromium() { return 1 fi - # Kill any leftover socat from a previous run - if [ -n "${SOCAT_PID:-}" ] && kill -0 "$SOCAT_PID" 2>/dev/null; then - kill "$SOCAT_PID" 2>/dev/null || true - wait "$SOCAT_PID" 2>/dev/null || true + # Kill any leftover CDP auth proxy from a previous run + if [ -n "${CDP_PROXY_PID:-}" ] && kill -0 "$CDP_PROXY_PID" 2>/dev/null; then + kill "$CDP_PROXY_PID" 2>/dev/null || true + wait "$CDP_PROXY_PID" 2>/dev/null || true fi - # Expose CDP on 0.0.0.0 so Docker port mapping can reach it from the host. - socat TCP-LISTEN:${CDP_PORT},fork,reuseaddr,bind=0.0.0.0 TCP:127.0.0.1:${CDP_INTERNAL_PORT} & - SOCAT_PID=$! - echo "Started socat CDP forwarder (PID $SOCAT_PID): 0.0.0.0:${CDP_PORT} -> 127.0.0.1:${CDP_INTERNAL_PORT}" + # Expose CDP on 0.0.0.0 via an authenticated proxy (replaces raw socat). + python3 /usr/local/bin/cdp-auth-proxy.py & + CDP_PROXY_PID=$! + echo "Started CDP auth proxy (PID $CDP_PROXY_PID): 0.0.0.0:${CDP_PORT} -> 127.0.0.1:${CDP_INTERNAL_PORT}" - # Verify the socat-forwarded CDP port is reachable + # Verify the proxied CDP port is reachable (with auth) local fwd_ready=false for i in 1 2 3 4 5; do - if curl -s "http://127.0.0.1:${CDP_PORT}/json/version" | grep -q "webSocketDebuggerUrl"; then - echo "✅ CDP forwarded and reachable on port ${CDP_PORT} (attempt $i)" + if curl -s -H "Authorization: Bearer ${TOOL_SERVER_TOKEN}" "http://127.0.0.1:${CDP_PORT}/json/version" | grep -q "webSocketDebuggerUrl"; then + echo "✅ CDP auth proxy reachable on port ${CDP_PORT} (attempt $i)" fwd_ready=true break fi sleep 1 done if [ "$fwd_ready" = false ]; then - echo "WARNING: CDP not reachable via socat on port ${CDP_PORT}" - echo " socat PID $SOCAT_PID alive: $(kill -0 $SOCAT_PID 2>&1 && echo yes || echo no)" + echo "WARNING: CDP not reachable via auth proxy on port ${CDP_PORT}" + echo " proxy PID $CDP_PROXY_PID alive: $(kill -0 $CDP_PROXY_PID 2>&1 && echo yes || echo no)" echo " ss output: $(ss -tlnp 2>/dev/null | grep ${CDP_PORT} || echo 'port not listening')" return 1 fi diff --git a/pyproject.toml b/pyproject.toml index 5edddf74..0029c1c3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -239,6 +239,9 @@ ignore = [ ] [tool.ruff.lint.per-file-ignores] +"containers/**/*.py" = [ + "T201", # print is the only logging mechanism in container scripts +] "tests/**/*.py" = [ "S106", # Possible hardcoded password "S108", # Possible insecure usage of temporary file/directory @@ -255,6 +258,11 @@ ignore = [ "PLR0912", # Too many branches (dispatch functions) "PLR0915", # Too many statements (complex browser task handling) ] +"strix/interface/tool_components/browser_renderer.py" = [ + "PLR0911", # Too many return statements (action dispatcher) + "PLR0912", # Too many branches + "PLR0915", # Too many statements +] "strix/telemetry/tracer.py" = [ "PLR0912", # Too many branches (save_run_data is legitimately complex) "PLR0915", # Too many statements diff --git a/strix/interface/tool_components/browser_renderer.py b/strix/interface/tool_components/browser_renderer.py index d09cca4b..ce1a550d 100644 --- a/strix/interface/tool_components/browser_renderer.py +++ b/strix/interface/tool_components/browser_renderer.py @@ -18,24 +18,22 @@ def _get_style_colors() -> dict[Any, str]: @register_tool_renderer class BrowserRenderer(BaseToolRenderer): - tool_name: ClassVar[str] = "browser_action" + tool_name: ClassVar[str] = "browser_actions" css_classes: ClassVar[list[str]] = ["tool-call", "browser-tool"] - SIMPLE_ACTIONS: ClassVar[dict[str, str]] = { - "back": "going back in browser history", - "forward": "going forward in browser history", - "scroll_down": "scrolling down", - "scroll_up": "scrolling up", - "refresh": "refreshing browser tab", - "close_tab": "closing browser tab", - "switch_tab": "switching browser tab", - "list_tabs": "listing browser tabs", - "view_source": "viewing page source", - "get_console_logs": "getting console logs", - "screenshot": "taking screenshot of browser tab", - "wait": "waiting...", - "close": "closing browser", - } + # -- palette (used only for highlights) ---------------------------- + NAV: ClassVar[str] = "#06b6d4" # cyan — links / URLs + INTERACT: ClassVar[str] = "#3b82f6" # blue — targets / values + OBSERVE: ClassVar[str] = "#a78bfa" # purple — data fields + EXEC: ClassVar[str] = "#f59e0b" # amber — task text + LIFE: ClassVar[str] = "#10b981" # teal — lifecycle values + OK: ClassVar[str] = "#22c55e" # green — success + ERR: ClassVar[str] = "#ef4444" # red — error + DIM: ClassVar[str] = "dim" # gray — prose / labels + + # ----------------------------------------------------------------- + # helpers + # ----------------------------------------------------------------- @classmethod def _get_token_color(cls, token_type: Any) -> str | None: @@ -50,87 +48,286 @@ class BrowserRenderer(BaseToolRenderer): def _highlight_js(cls, code: str) -> Text: lexer = get_lexer_by_name("javascript") text = Text() - for token_type, token_value in lexer.get_tokens(code): if not token_value: continue color = cls._get_token_color(token_type) text.append(token_value, style=color) - return text + @classmethod + def _head(cls) -> Text: + text = Text() + text.append("@ ", style=cls.DIM) + return text + + @classmethod + def _icon(cls, color: str) -> Text: + text = Text() + text.append("◈ ", style=color) + return text + + @classmethod + def _status_mark(cls, status: str) -> Text: + text = Text() + if status == "completed": + text.append(" ✓", style=f"dim {cls.OK}") + elif status in ("failed", "error"): + text.append(" ✗", style=f"dim {cls.ERR}") + return text + + @classmethod + def _append_fields(cls, text: Text, res: dict[str, Any]) -> None: + """Append a dim summary of returned history fields.""" + fields = res.get("fields") + if not fields or not isinstance(fields, dict): + return + names = sorted(fields) + text.append("\n ") + text.append("fields ", style=cls.OBSERVE) + text.append(" ".join(names), style=cls.DIM) + + # ----------------------------------------------------------------- + # public + # ----------------------------------------------------------------- + @classmethod def render(cls, tool_data: dict[str, Any]) -> Static: args = tool_data.get("args", {}) status = tool_data.get("status", "unknown") + result = tool_data.get("result") action = args.get("action", "") - content = cls._build_content(action, args) + content = cls._build_content(action, args, status, result) css_classes = cls.get_css_classes(status) return Static(content, classes=css_classes) - @classmethod - def _build_url_action(cls, text: Text, label: str, url: str | None, suffix: str = "") -> None: - text.append(label, style="#06b6d4") - if url: - text.append(url, style="#06b6d4") - if suffix: - text.append(suffix, style="#06b6d4") + # ----------------------------------------------------------------- + # content builder + # ----------------------------------------------------------------- @classmethod - def _build_content(cls, action: str, args: dict[str, Any]) -> Text: - text = Text() - text.append("🌐 ") + def _build_content( + cls, + action: str, + args: dict[str, Any], + status: str, + result: Any, + ) -> Text: + if action == "run": + return cls._build_run(args, status, result) - if action in cls.SIMPLE_ACTIONS: - text.append(cls.SIMPLE_ACTIONS[action], style="#06b6d4") - return text - - url = args.get("url") - - url_actions = { - "launch": ("launching ", " on browser" if url else "browser"), - "goto": ("navigating to ", ""), - "new_tab": ("opening tab ", ""), + # --- simple one-liners ---------------------------------------- + simple = { + "back": "going back", + "close": "closing browser", + "close_tab": "closing tab", + "screenshot": "taking screenshot", + "state": "reading page state", + "extract": "extracting content", } - if action in url_actions: - label, suffix = url_actions[action] - if action == "launch" and not url: - text.append("launching browser", style="#06b6d4") - else: - cls._build_url_action(text, label, url, suffix) + if action in simple: + text = cls._head() + text.append(simple[action], style=cls.DIM) + text.append_text(cls._status_mark(status)) return text - click_actions = { - "click": "clicking", - "double_click": "double clicking", - "hover": "hovering", - } - if action in click_actions: - text.append(click_actions[action], style="#06b6d4") + # --- launch ---------------------------------------------------- + if action == "launch": + mode = "local" if args.get("use_local") else "sandboxed" + text = cls._icon(cls.LIFE) + text.append("launching browser", style=f"bold {cls.LIFE}") + text.append(f" {mode}", style=cls.DIM) return text - handlers: dict[str, tuple[str, str | None]] = { - "type": ("typing ", args.get("text")), - "press_key": ("pressing key ", args.get("key")), - "save_pdf": ("saving PDF to ", args.get("file_path")), - } - if action in handlers: - label, value = handlers[action] - text.append(label, style="#06b6d4") - if value: - text.append(str(value), style="#06b6d4") + # --- navigate -------------------------------------------------- + if action == "open": + text = cls._head() + text.append("navigating to ", style=cls.DIM) + url = args.get("url", "") + if len(url) > 80: + url = url[:77] + "..." + text.append(url, style=f"{cls.NAV} underline") + text.append_text(cls._status_mark(status)) return text - if action == "execute_js": - text.append("executing javascript", style="#06b6d4") - js_code = args.get("js_code") - if js_code: + # --- pointer actions ------------------------------------------- + if action in ("click", "dblclick", "rightclick", "hover"): + labels = { + "click": "clicking", + "dblclick": "double clicking", + "rightclick": "right clicking", + "hover": "hovering over", + } + text = cls._head() + text.append(labels[action], style=cls.DIM) + index = args.get("index") + if index is not None: + text.append(f" #{index}", style=f"bold {cls.INTERACT}") + elif args.get("x") is not None and args.get("y") is not None: + text.append(f" ({args['x']}, {args['y']})", style=cls.INTERACT) + text.append_text(cls._status_mark(status)) + return text + + # --- text input ------------------------------------------------ + if action in ("type", "input"): + label = "typing" if action == "type" else "inputting" + text = cls._head() + text.append(label, style=cls.DIM) + t = args.get("text") + if t: + preview = t if len(t) <= 60 else t[:57] + "..." + text.append(f' "{preview}"', style=cls.INTERACT) + text.append_text(cls._status_mark(status)) + return text + + # --- scroll ---------------------------------------------------- + if action == "scroll": + d = args.get("direction", "down") + text = cls._head() + text.append("scrolling ", style=cls.DIM) + text.append(d, style=cls.INTERACT) + text.append_text(cls._status_mark(status)) + return text + + # --- tab switch ------------------------------------------------ + if action == "switch": + text = cls._head() + text.append("switching to tab ", style=cls.DIM) + text.append(str(args.get("tab", "?")), style=f"bold {cls.NAV}") + text.append_text(cls._status_mark(status)) + return text + + # --- keyboard -------------------------------------------------- + if action == "keys": + text = cls._head() + text.append("pressing ", style=cls.DIM) + text.append(args.get("keys", ""), style=f"bold {cls.INTERACT}") + text.append_text(cls._status_mark(status)) + return text + + # --- select ---------------------------------------------------- + if action == "select": + text = cls._head() + text.append("selecting ", style=cls.DIM) + text.append(args.get("value", ""), style=f"bold {cls.INTERACT}") + text.append_text(cls._status_mark(status)) + return text + + # --- eval js --------------------------------------------------- + if action == "eval": + text = cls._head() + text.append("executing javascript", style=cls.DIM) + text.append_text(cls._status_mark(status)) + js = args.get("js") + if js: text.append("\n") - text.append_text(cls._highlight_js(js_code)) + text.append_text(cls._highlight_js(js)) return text + # --- cookies --------------------------------------------------- + if action == "cookies": + sub = args.get("subcommand", "") + text = cls._head() + text.append("cookies ", style=cls.DIM) + text.append(sub, style=cls.OBSERVE) + text.append_text(cls._status_mark(status)) + return text + + # --- wait ------------------------------------------------------ + if action == "wait": + sub = args.get("subcommand", "") + target = args.get("selector") or args.get("text") or "" + text = cls._head() + if status == "completed": + text.append(f"waited for {sub}", style=cls.DIM) + text.append_text(cls._status_mark(status)) + else: + text.append(f"waiting for {sub}", style=cls.DIM) + if target: + text.append(" ", style=cls.DIM) + text.append(target, style=cls.OBSERVE) + return text + + # --- get ------------------------------------------------------- + if action == "get": + sub = args.get("subcommand", "") + text = cls._head() + text.append("getting ", style=cls.DIM) + text.append(sub, style=cls.OBSERVE) + text.append_text(cls._status_mark(status)) + return text + + # --- fallback -------------------------------------------------- + text = cls._head() if action: - text.append(action, style="#06b6d4") + text.append(action, style=cls.DIM) + text.append_text(cls._status_mark(status)) + return text + + # ----------------------------------------------------------------- + # run task + # ----------------------------------------------------------------- + + @classmethod + def _build_run( + cls, + args: dict[str, Any], + status: str, + result: Any, + ) -> Text: + task = args.get("task", "") + + if status == "running": + text = cls._head() + text.append("running task", style=f"bold {cls.EXEC}") + if task: + text.append("\n ") + text.append(task, style=cls.EXEC) + rf = args.get("return_fields") + if rf and isinstance(rf, list): + text.append("\n ") + text.append("returning ", style=cls.OBSERVE) + text.append(" ".join(rf), style=cls.DIM) + return text + + # Completed or failed — show result + if status in ("completed", "failed", "error"): + res = result if isinstance(result, dict) else {} + has_error = "error" in res + + if has_error: + text = cls._head() + text.append("task failed", style=f"bold {cls.ERR}") + if task: + text.append("\n ") + text.append(task, style="dim strike") + text.append("\n ") + error_msg = str(res["error"]) + if len(error_msg) > 200: + error_msg = error_msg[:197] + "..." + text.append(error_msg, style=cls.ERR) + else: + text = cls._head() + text.append("task completed", style=f"bold {cls.OK}") + if task: + text.append("\n ") + text.append(task, style=cls.DIM) + output = res.get("result", "") + if output: + text.append("\n ") + output_str = str(output) + if len(output_str) > 300: + output_str = output_str[:297] + "..." + text.append(output_str, style=cls.DIM) + cls._append_fields(text, res) + return text + + # Unknown status + text = cls._head() + text.append("running task", style=cls.DIM) + if task: + text.append("\n ") + text.append(task, style=cls.DIM) return text diff --git a/strix/runtime/tool_server.py b/strix/runtime/tool_server.py index 2b23046f..913ce483 100644 --- a/strix/runtime/tool_server.py +++ b/strix/runtime/tool_server.py @@ -18,7 +18,6 @@ if not SANDBOX_MODE: raise RuntimeError("Tool server should only run in sandbox mode (STRIX_SANDBOX_MODE=true)") parser = argparse.ArgumentParser(description="Start Strix tool server") -parser.add_argument("--token", required=True, help="Authentication token") parser.add_argument("--host", default="0.0.0.0", help="Host to bind to") # nosec parser.add_argument("--port", type=int, required=True, help="Port to bind to") parser.add_argument( @@ -29,7 +28,11 @@ parser.add_argument( ) args = parser.parse_args() -EXPECTED_TOKEN = args.token + +# Read token from environment to avoid leaking it in /proc//cmdline. +EXPECTED_TOKEN = os.environ.get("TOOL_SERVER_TOKEN", "") +if not EXPECTED_TOKEN: + raise RuntimeError("TOOL_SERVER_TOKEN environment variable must be set") REQUEST_TIMEOUT = args.timeout app = FastAPI() @@ -136,7 +139,12 @@ async def register_agent( async def _check_cdp_health() -> dict[str, Any]: - """Probe the container-local Chromium CDP endpoint.""" + """Probe the container-local Chromium CDP endpoint. + + Returns only a boolean status — never the ``webSocketDebuggerUrl`` + or browser version, because the ``/health`` endpoint is + unauthenticated and those fields would leak session secrets. + """ cdp_port = os.getenv("CDP_INTERNAL_PORT", "19222") cdp_url = f"http://127.0.0.1:{cdp_port}/json/version" try: @@ -145,15 +153,31 @@ async def _check_cdp_health() -> dict[str, Any]: async with httpx.AsyncClient(timeout=3, trust_env=False) as client: resp = await client.get(cdp_url) if resp.status_code == 200: - data = resp.json() - return { - "status": "healthy", - "browser": data.get("Browser", "unknown"), - "ws_url": data.get("webSocketDebuggerUrl", ""), - } - return {"status": "unhealthy", "detail": f"HTTP {resp.status_code}"} - except Exception as exc: # noqa: BLE001 - return {"status": "unhealthy", "detail": str(exc)} + return {"status": "healthy"} + return {"status": "unhealthy"} + except Exception: # noqa: BLE001 + return {"status": "unhealthy"} + + +@app.get("/cdp/version") +async def cdp_version( + credentials: HTTPAuthorizationCredentials = security_dependency, # noqa: ARG001 +) -> dict[str, Any]: + """Return Chromium's ``/json/version`` data via the authenticated tool server. + + This lets the host discover the WebSocket debugger URL (which contains + a random browser GUID) without exposing ``/json/version`` on the + unauthenticated CDP port. + """ + cdp_port = os.getenv("CDP_INTERNAL_PORT", "19222") + cdp_url = f"http://127.0.0.1:{cdp_port}/json/version" + import httpx + + async with httpx.AsyncClient(timeout=5, trust_env=False) as client: + resp = await client.get(cdp_url) + resp.raise_for_status() + data: dict[str, Any] = resp.json() + return data @app.get("/health") diff --git a/strix/tools/browser/browser_actions.py b/strix/tools/browser/browser_actions.py index 8063f42d..84065d12 100644 --- a/strix/tools/browser/browser_actions.py +++ b/strix/tools/browser/browser_actions.py @@ -168,7 +168,12 @@ async def _cleanup_agent(agent: Any) -> None: await asyncio.wait_for(coro, timeout=5) -async def _run_agent_task(task: str, session: _BrowserSession) -> dict[str, Any]: +async def _run_agent_task( + task: str, + session: _BrowserSession, + *, + return_fields: list[str] | None = None, +) -> dict[str, Any]: """Run a browser-use task on the session's browser with resilient reconnection. Lifecycle: @@ -253,7 +258,7 @@ async def _run_agent_task(task: str, session: _BrowserSession) -> dict[str, Any] return {"error": f"Browser task failed: {exc}", "is_running": False} # --- Success path --- - interpreted = _interpret_agent_result(result) + interpreted = _interpret_agent_result(result, return_fields=return_fields) if "error" in interpreted: error_msg = str(interpreted["error"]) @@ -291,12 +296,51 @@ async def _run_agent_task(task: str, session: _BrowserSession) -> dict[str, Any] } -def _interpret_agent_result(result: Any) -> dict[str, Any]: +_JSON_PRIMITIVES = str | int | float | bool | None + + +def _json_safe(value: Any) -> Any: + """Coerce a value into a JSON-serialisable form.""" + if isinstance(value, _JSON_PRIMITIVES): + return value + if isinstance(value, list): + return [v if isinstance(v, _JSON_PRIMITIVES) else str(v) for v in value] + return str(value) + + +_HISTORY_ACCESSORS: dict[str, str] = { + "urls": "urls", + "screenshot_paths": "screenshot_paths", + "screenshots": "screenshots", + "action_names": "action_names", + "extracted_content": "extracted_content", + "errors": "errors", + "model_actions": "model_actions", + "model_outputs": "model_outputs", + "last_action": "last_action", + "final_result": "final_result", + "is_done": "is_done", + "has_errors": "has_errors", + "model_thoughts": "model_thoughts", + "action_results": "action_results", + "action_history": "action_history", + "number_of_steps": "number_of_steps", + "total_duration_seconds": "total_duration_seconds", +} + + +def _interpret_agent_result( + result: Any, + return_fields: list[str] | None = None, +) -> dict[str, Any]: """Inspect an ``AgentHistoryList`` and return a success or error dict. If the browser-use agent encountered errors (CDP failures, navigation errors, etc.) this returns an ``{"error": ...}`` dict so the calling strix agent treats the task as failed rather than succeeded. + + When *return_fields* is provided, the requested history accessors are + included in the response under a ``"fields"`` key. """ # Collect errors reported by the browser-use agent. errors: list[str] = [] @@ -343,24 +387,47 @@ def _interpret_agent_result(result: Any) -> dict[str, Any]: "is_running": False, } - return { + out: dict[str, Any] = { "message": "Task completed", "result": final_result or str(result), "is_running": False, } + # Attach optional fields requested by the caller. + if return_fields: + fields: dict[str, Any] = {} + available = ", ".join(sorted(_HISTORY_ACCESSORS)) + for field in return_fields: + accessor = _HISTORY_ACCESSORS.get(field) + if accessor is None: + fields[field] = f"unknown field (available: {available})" + continue + try: + attr = getattr(result, accessor, None) + value = attr() if callable(attr) else attr + fields[field] = _json_safe(value) + except Exception as exc: # noqa: BLE001 + logger.debug("Error extracting field %s: %s", field, exc) + fields[field] = f"error: {exc}" + out["fields"] = fields + + return out + # --------------------------------------------------------------------------- # CDP URL resolution # --------------------------------------------------------------------------- -def _resolve_cdp_url(agent_state: Any) -> str: - """Extract the CDP URL from agent_state's sandbox_info. +def _resolve_cdp_url(agent_state: Any) -> tuple[str, str]: + """Extract the CDP URL and auth token from agent_state's sandbox_info. The sandbox container runs Chromium with ``--remote-debugging-port`` and the Docker runtime maps it to an ephemeral host port stored in - ``sandbox_info["browser_cdp_port"]``. + ``sandbox_info["browser_cdp_port"]``. The auth token (shared with + the tool server) is used to authenticate against the CDP auth proxy. + + Returns ``(cdp_url, auth_token)``. """ if not hasattr(agent_state, "sandbox_info") or not agent_state.sandbox_info: raise ValueError( @@ -379,6 +446,8 @@ def _resolve_cdp_url(agent_state: Any) -> str: f"Available keys: {list(sandbox_info.keys())}" ) + auth_token: str = sandbox_info.get("auth_token", "") + # Resolve the Docker host (same logic used by the tool server URL). docker_host = os.getenv("DOCKER_HOST", "") if docker_host: @@ -400,7 +469,7 @@ def _resolve_cdp_url(agent_state: Any) -> str: host, sandbox_info.get("workspace_id", "?")[:12], ) - return cdp_url + return cdp_url, auth_token # --------------------------------------------------------------------------- @@ -440,6 +509,7 @@ async def browser_actions( http_only: bool = False, same_site: str | None = None, expires: float | None = None, + return_fields: list[str] | None = None, *, agent_state: Any = None, ) -> dict[str, Any]: @@ -511,13 +581,18 @@ async def browser_actions( "profile_directory": session.profile_directory or "auto", "is_running": True, } - cdp_url = _resolve_cdp_url(agent_state) - session = await _launch_browser(cdp_url, agent_id) + cdp_url, auth_token = _resolve_cdp_url(agent_state) + session = await _launch_browser(cdp_url, agent_id, auth_token) + # Strip auth token from the ws_url before returning — the + # token is a secret and must not leak to the calling agent. + import re + + safe_ws = re.sub(r"[?&]token=[^&]+", "", session.ws_url) return { "message": "Browser launched and ready", "mode": "sandboxed", "cdp_url": session.cdp_url, - "ws_url": session.ws_url, + "ws_url": safe_ws, "is_running": True, } @@ -537,7 +612,7 @@ async def browser_actions( session.local, task, ) - result = await _run_agent_task(task, session) + result = await _run_agent_task(task, session, return_fields=return_fields) logger.info("Browser task completed for agent %s", agent_id) # Agent.run() calls browser_session.kill() which destroys the diff --git a/strix/tools/browser/browser_actions_schema.xml b/strix/tools/browser/browser_actions_schema.xml index 9ca6ca8e..26410a06 100644 --- a/strix/tools/browser/browser_actions_schema.xml +++ b/strix/tools/browser/browser_actions_schema.xml @@ -26,6 +26,7 @@ **Agent mode:** - run: Execute a natural-language browser task. Requires 'task' parameter. + Optionally pass 'return_fields' to select which history data to include in the response. **Navigation:** - open: Navigate to a URL. Requires 'url'. @@ -149,6 +150,30 @@ Cookie expiration timestamp. Used by cookies set. + + Only used with action='run'. A list of history fields to include in the + response under a "fields" key. By default only the final result is returned. Use this + to request additional data from the agent run. + + Available fields: + - urls: List of visited URLs + - screenshot_paths: List of screenshot file paths + - screenshots: List of screenshots as base64 strings + - action_names: Names of executed actions + - extracted_content: Extracted content from all actions + - errors: List of errors (None for steps without errors) + - model_actions: All actions with their parameters + - model_outputs: All model outputs + - last_action: Last action in history + - final_result: Final extracted content (last step) + - is_done: Whether the agent completed + - has_errors: Whether any errors occurred + - model_thoughts: Agent reasoning process + - action_results: All ActionResult objects + - action_history: Truncated action history with essential fields + - number_of_steps: Number of steps taken + - total_duration_seconds: Total duration of all steps + Returns a dict with action-specific results. On error, returns an error field. @@ -267,6 +292,13 @@ Go to https://example.com/login, fill in username "admin" and password "secret", then click the login button + # Run a task and request specific output fields + + run + Navigate through the site and collect all page titles + ["urls", "extracted_content", "number_of_steps"] + + # Close the browser when done close diff --git a/strix/tools/browser/browser_manager.py b/strix/tools/browser/browser_manager.py index 1c9ba8e3..cd342451 100644 --- a/strix/tools/browser/browser_manager.py +++ b/strix/tools/browser/browser_manager.py @@ -36,6 +36,7 @@ _background_tasks: set[asyncio.Task[None]] = set() class _BrowserSession: __slots__ = ( + "auth_token", "browser", "cdp_url", "consecutive_failures", @@ -53,9 +54,11 @@ class _BrowserSession: cdp_url: str, ws_url: str, *, + auth_token: str = "", # nosec B107 local: bool = False, profile_directory: str | None = None, ) -> None: + self.auth_token = auth_token self.browser = browser self.cdp_url = cdp_url self.ws_url = ws_url @@ -84,26 +87,55 @@ _sessions: dict[str, _BrowserSession] = {} _lock = threading.Lock() +def _squelch_connection_error( + loop: asyncio.AbstractEventLoop, + context: dict[str, Any], +) -> None: + """Swallow 'Client is stopping' futures left behind by the CDP client. + + During browser shutdown the CDP WebSocket layer may reject pending + futures with ``ConnectionError('Client is stopping')``. Because no + caller ever awaits them, asyncio logs noisy "Future exception was + never retrieved" tracebacks. This handler silences that specific + case and delegates everything else to the default handler. + """ + exc = context.get("exception") + if isinstance(exc, ConnectionError) and "Client is stopping" in str(exc): + logger.debug("Suppressed CDP teardown noise: %s", exc) + return + loop.default_exception_handler(context) + + async def _safe_close_browser(browser: Any, label: str = "") -> None: """Best-effort close/stop of a Browser object, swallowing all errors.""" if browser is None: return + + loop = asyncio.get_running_loop() + original_handler = loop.get_exception_handler() + loop.set_exception_handler(_squelch_connection_error) + tag = f" ({label})" if label else "" - for method_name in ("close", "stop"): - fn = getattr(browser, method_name, None) - if not callable(fn): - continue - try: - coro = fn() - if asyncio.iscoroutine(coro): - await asyncio.wait_for(coro, timeout=10) - except TimeoutError: - logger.warning("Browser%s %s() timed out after 10s", tag, method_name) - except Exception as exc: # noqa: BLE001 - logger.debug("Browser%s %s() failed: %s", tag, method_name, exc) - else: - logger.debug("Browser%s closed via %s()", tag, method_name) - return + try: + for method_name in ("close", "stop"): + fn = getattr(browser, method_name, None) + if not callable(fn): + continue + try: + coro = fn() + if asyncio.iscoroutine(coro): + await asyncio.wait_for(coro, timeout=10) + except TimeoutError: + logger.warning("Browser%s %s() timed out after 10s", tag, method_name) + except Exception as exc: # noqa: BLE001 + logger.debug("Browser%s %s() failed: %s", tag, method_name, exc) + else: + logger.debug("Browser%s closed via %s()", tag, method_name) + return + finally: + # Give straggler futures a moment to settle before restoring. + await asyncio.sleep(0.1) + loop.set_exception_handler(original_handler) async def _refresh_browser(session: _BrowserSession) -> None: @@ -139,6 +171,7 @@ async def _refresh_browser(session: _BrowserSession) -> None: ws_url, _ = await asyncio.to_thread( _wait_for_cdp, session.cdp_url, + session.auth_token, max_attempts=10, interval=2.0, ) @@ -154,7 +187,10 @@ async def _refresh_browser(session: _BrowserSession) -> None: def _wait_for_cdp( - cdp_url: str, max_attempts: int = 30, interval: float = 1.0 + cdp_url: str, + auth_token: str = "", # nosec B107 + max_attempts: int = 30, + interval: float = 1.0, ) -> tuple[str, dict[str, Any]]: """Block until the CDP endpoint at *cdp_url* responds to ``/json/version``. @@ -164,21 +200,24 @@ def _wait_for_cdp( reachable from the host — we replace the host:port with the values from *cdp_url* (the Docker-mapped endpoint). - The Chromium process inside the sandbox container may still be starting - when the tool server is already healthy. We poll the CDP ``/json/version`` - endpoint before handing the URL to browser-use. + When *auth_token* is provided it is sent as a Bearer header on the + HTTP probe and appended as a ``?token=`` query parameter on the + rewritten WebSocket URL (for the CDP auth proxy). """ from urllib.parse import urlparse import httpx version_url = cdp_url.rstrip("/") + "/json/version" + headers: dict[str, str] = {} + if auth_token: + headers["Authorization"] = f"Bearer {auth_token}" last_error: str = "" for attempt in range(1, max_attempts + 1): try: with httpx.Client(trust_env=False, timeout=5) as client: - resp = client.get(version_url) + resp = client.get(version_url, headers=headers) if resp.status_code == 200 and "webSocketDebuggerUrl" in resp.text: version_info = resp.json() raw_ws = version_info.get("webSocketDebuggerUrl", "") @@ -193,13 +232,24 @@ def _wait_for_cdp( f"{parsed_cdp.hostname}:{parsed_cdp.port}", ) + # Append auth token so browser-use's WebSocket upgrade + # request passes through the CDP auth proxy. + if auth_token: + sep = "&" if "?" in ws_url else "?" + ws_url = f"{ws_url}{sep}token={auth_token}" + + # Log the WS URL with the token redacted to avoid + # leaking credentials into log files. + import re + + safe_ws = re.sub(r"([?&])token=[^&]+", r"\1token=REDACTED", ws_url) logger.info( "CDP ready at %s (attempt %d): browser=%s, ws_raw=%s, ws_rewritten=%s", cdp_url, attempt, version_info.get("Browser", "unknown"), raw_ws, - ws_url, + safe_ws, ) return ws_url, version_info last_error = f"HTTP {resp.status_code}, body={resp.text[:200]}" @@ -222,7 +272,7 @@ def _wait_for_cdp( ) -async def _launch_browser(cdp_url: str, agent_id: str) -> _BrowserSession: +async def _launch_browser(cdp_url: str, agent_id: str, auth_token: str = "") -> _BrowserSession: # nosec B107 """Create a new browser session connected to the sandbox browser via CDP. The *cdp_url* points to the Chromium instance running inside the Docker @@ -243,7 +293,7 @@ async def _launch_browser(cdp_url: str, agent_id: str) -> _BrowserSession: # Wait for the container's Chromium to be ready and get the rewritten WS URL. # We pass the ws:// URL directly so browser-use skips its own /json/version # fetch (which would get the container-internal address). - ws_url, version_info = await asyncio.to_thread(_wait_for_cdp, cdp_url) + ws_url, version_info = await asyncio.to_thread(_wait_for_cdp, cdp_url, auth_token) logger.info( "Chromium version: %s, protocol: %s, user-agent: %s", version_info.get("Browser", "?"), @@ -260,7 +310,9 @@ async def _launch_browser(cdp_url: str, agent_id: str) -> _BrowserSession: ws_url, ) - session = _BrowserSession(browser=browser, cdp_url=cdp_url, ws_url=ws_url) + session = _BrowserSession( + browser=browser, cdp_url=cdp_url, ws_url=ws_url, auth_token=auth_token + ) with _lock: if agent_id in _sessions: @@ -380,15 +432,27 @@ def cleanup_agent(agent_id: str) -> None: def _close_all() -> None: with _lock: - agent_ids = list(_sessions.keys()) - for aid in agent_ids: - with contextlib.suppress(Exception): - with _lock: - session = _sessions.pop(aid, None) - if session is not None: - # At atexit there's no running loop, so create one. - with contextlib.suppress(Exception): - asyncio.run(_shutdown_session(session)) + sessions = [_sessions.pop(aid) for aid in list(_sessions) if aid in _sessions] + if not sessions: + return + for session in sessions: + session.invalidated = True + browser = session.browser + session.browser = None + if browser is None: + continue + # Attempt a synchronous best-effort close without spinning up a + # full asyncio.run(), which hangs when a signal handler fires + # inside the new loop's select(). + for method_name in ("close", "stop"): + fn = getattr(browser, method_name, None) + if not callable(fn): + continue + with contextlib.suppress(Exception): + coro = fn() + if asyncio.iscoroutine(coro): + coro.close() # discard; can't await at atexit + break atexit.register(_close_all) @@ -428,14 +492,17 @@ async def _reinitialize_after_agent(session: _BrowserSession) -> None: # --------------------------------------------------------------------------- -def _check_cdp_alive(cdp_url: str) -> bool: +def _check_cdp_alive(cdp_url: str, auth_token: str = "") -> bool: # nosec B107 """Quick check that the CDP endpoint is reachable (non-blocking from sync).""" import httpx version_url = cdp_url.rstrip("/") + "/json/version" + headers: dict[str, str] = {} + if auth_token: + headers["Authorization"] = f"Bearer {auth_token}" try: with httpx.Client(trust_env=False, timeout=5) as client: - resp = client.get(version_url) + resp = client.get(version_url, headers=headers) return resp.status_code == 200 and "webSocketDebuggerUrl" in resp.text except Exception: # noqa: BLE001 return False @@ -449,7 +516,7 @@ async def _wait_for_cdp_recovery(session: _BrowserSession, task_num: int) -> boo max_attempts = int(_CDP_RECOVERY_TIMEOUT / _CDP_RECOVERY_INTERVAL) for attempt in range(1, max_attempts + 1): await asyncio.sleep(_CDP_RECOVERY_INTERVAL) - if await asyncio.to_thread(_check_cdp_alive, session.cdp_url): + if await asyncio.to_thread(_check_cdp_alive, session.cdp_url, session.auth_token): logger.info( "Task #%d: CDP recovered after %ds", task_num, @@ -504,7 +571,7 @@ async def _ensure_healthy_session(session: _BrowserSession, task_num: int) -> st reason, ) # Make sure CDP is alive first (may need watchdog restart). - cdp_alive = await asyncio.to_thread(_check_cdp_alive, session.cdp_url) + cdp_alive = await asyncio.to_thread(_check_cdp_alive, session.cdp_url, session.auth_token) if not cdp_alive and not await _wait_for_cdp_recovery(session, task_num): return ( f"Chromium CDP at {session.cdp_url} is not responding after " @@ -519,7 +586,7 @@ async def _ensure_healthy_session(session: _BrowserSession, task_num: int) -> st return None # 2) Normal pre-flight: verify CDP is alive. - if await asyncio.to_thread(_check_cdp_alive, session.cdp_url): + if await asyncio.to_thread(_check_cdp_alive, session.cdp_url, session.auth_token): return None # all good # 3) CDP down — wait for watchdog, then refresh.