diff --git a/containers/Dockerfile b/containers/Dockerfile index 9943266a..cabb6b39 100644 --- a/containers/Dockerfile +++ b/containers/Dockerfile @@ -117,6 +117,21 @@ ENV AGENT_BROWSER_EXECUTABLE_PATH=/usr/bin/chromium ENV AGENT_BROWSER_USER_AGENT="Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36" ENV AGENT_BROWSER_ARGS="--disable-blink-features=AutomationControlled,--no-first-run,--no-default-browser-check,--lang=en-US" ENV AGENT_BROWSER_SCREENSHOT_DIR=/workspace/.agent-browser-screenshots +ENV AGENT_BROWSER_IDLE_TIMEOUT_MS=180000 +USER root +RUN set -eu; \ + { \ + for var in AGENT_BROWSER_EXECUTABLE_PATH AGENT_BROWSER_USER_AGENT \ + AGENT_BROWSER_ARGS AGENT_BROWSER_SCREENSHOT_DIR \ + AGENT_BROWSER_IDLE_TIMEOUT_MS; do \ + eval "value=\${$var}"; \ + printf 'export %s="${%s:-%s}"\n' "$var" "$var" "$value"; \ + done; \ + } > /tmp/agent-browser.sh; \ + install -m 0644 /tmp/agent-browser.sh /etc/profile.d/agent-browser.sh; \ + rm /tmp/agent-browser.sh; \ + env -i bash -lc 'test "${AGENT_BROWSER_IDLE_TIMEOUT_MS}" = "180000"' +USER pentester RUN /home/pentester/.npm-global/bin/agent-browser doctor --offline --quick RUN set -eux; \ diff --git a/pyproject.toml b/pyproject.toml index 269a0710..257beb81 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "strix-agent" -version = "1.5.1" +version = "1.5.3" description = "Open-source AI Hackers for your apps" readme = "README.md" license = "Apache-2.0" diff --git a/strix/agents/factory.py b/strix/agents/factory.py index e40e60e8..9d599a54 100644 --- a/strix/agents/factory.py +++ b/strix/agents/factory.py @@ -160,7 +160,9 @@ def _schema_types(spec: dict[str, Any]) -> set[str]: def _decode_structured(value: str, types: set[str]) -> Any: stripped = value.strip() if not stripped: - return value + # An empty string is the model's "no value" for a list/dict param; give it + # the empty container so it validates instead of failing the type check. + return [] if "array" in types else {} try: decoded = json.loads(stripped) except json.JSONDecodeError: diff --git a/strix/agents/prompts/system_prompt.jinja b/strix/agents/prompts/system_prompt.jinja index 5fc697d7..23493d2d 100644 --- a/strix/agents/prompts/system_prompt.jinja +++ b/strix/agents/prompts/system_prompt.jinja @@ -39,6 +39,8 @@ INTERACTIVE BEHAVIOR: - To end the whole engagement, call the lifecycle tool: finish_scan (root) or agent_finish (subagent). - A turn that ends with plain text and no tool call does NOT stop you: the system nudges you to continue and will re-run you. Do not rely on going silent to pause — it will not pause you. - Answering a user question: put the answer in respond_to_user's message. Do not write the answer as plain text and then fall silent — that does not reach a stopping point, it just triggers a continuation nudge. +- If all you want to do is reply and stop, that whole turn is ONE respond_to_user call carrying the answer. Do not write the answer as text and then call respond_to_user as well: the user reads it twice. +- If you do end a turn on plain text and the nudge arrives, your words already reached the user. Do not restate them: call respond_to_user with NO message to simply wait, or with only whatever you still need to add. - You may include brief explanatory text before a tool call, and you can narrate while you work — plain text is shown to the user as you go. Narrating is free; respond_to_user is specifically the act of WAITING for the user, so do not call it just to give a status update. - Respond naturally when the user asks questions or gives instructions. - While actively working on a task, every turn should carry exactly one tool call — use think to plan, the appropriate tool to act, and respond_to_user only when you genuinely need the user. @@ -261,7 +263,13 @@ Remember: A single well-validated high-impact vulnerability is worth more than d AGENT ISOLATION & SANDBOXING: - All agents run in the same shared Docker container for efficiency -- Each agent has its own: browser sessions, terminal sessions +- Each agent has its own terminal sessions +- Browsers are NOT per-agent by default: `agent-browser` with no `--session` is one + shared browser, so a concurrent agent's navigation invalidates your page and refs. + Pass `--session ` for any browser work of your own — then it is + yours alone. Each session is a full Chromium (~340 MB) on this shared box, so keep + one, not several, and `agent-browser --session close` when you're done with + the target; an idle browser is reclaimed automatically after 3 minutes - All agents share the same /workspace directory and proxy history - Agents can see each other's files and proxy traffic for better collaboration diff --git a/strix/config/models.py b/strix/config/models.py index 19c78eb4..671bd599 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -657,27 +657,31 @@ def _install_openrouter_stream_cost_capture() -> None: litellm.OpenrouterConfig = _StrixOpenrouterConfig # type: ignore[misc] -_OPENROUTER_ATTRIBUTION_HEADERS = { +OPENROUTER_ATTRIBUTION_HEADERS = { "HTTP-Referer": "https://strix.ai", "X-Title": "Strix", "X-OpenRouter-Categories": "cli-agent", } +def is_openrouter_model(model_name: str | None) -> bool: + return bool(model_name) and "openrouter/" in (model_name or "").strip().lower() + + def _configure_openrouter_attribution(model_name: str | None) -> None: import litellm current: object = litellm.headers existing: dict[str, str] = current if isinstance(current, dict) else {} - if not model_name or "openrouter/" not in model_name.strip().lower(): - if any(key in existing for key in _OPENROUTER_ATTRIBUTION_HEADERS): + if not is_openrouter_model(model_name): + if any(key in existing for key in OPENROUTER_ATTRIBUTION_HEADERS): remaining = { - k: v for k, v in existing.items() if k not in _OPENROUTER_ATTRIBUTION_HEADERS + k: v for k, v in existing.items() if k not in OPENROUTER_ATTRIBUTION_HEADERS } litellm.headers = remaining or None # type: ignore[assignment] return - litellm.headers = {**existing, **_OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment] + litellm.headers = {**existing, **OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment] def _configure_extra_headers(llm: LlmSettings) -> None: diff --git a/strix/core/execution.py b/strix/core/execution.py index 91ceb8af..bd99e7c3 100644 --- a/strix/core/execution.py +++ b/strix/core/execution.py @@ -830,7 +830,7 @@ async def _append_tool_required_message( "execution and never hands control to the user: it is shown to the user, and the " "run continues. Continue immediately and call exactly one tool. " "If you have something to tell the user and nothing to do until they reply, " - "call respond_to_user. " + "call respond_to_user — with no message if you have already said it. " "If you are blocked waiting for another agent, call wait_for_agents. " f"If the whole engagement is complete, call {finish_tool}. " "Otherwise use the appropriate execution or planning tool. " diff --git a/strix/core/inputs.py b/strix/core/inputs.py index a89248b2..a1106ae2 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -10,10 +10,12 @@ from openai.types.shared import Reasoning from strix.config.models import ( DEFAULT_MODEL_RETRY, + OPENROUTER_ATTRIBUTION_HEADERS, bedrock_route_supports_prompt_caching, is_bedrock_route, is_claude_model, is_known_openai_bare_model, + is_openrouter_model, model_supports_reasoning, request_timeout_extra_args, ) @@ -138,6 +140,15 @@ def build_root_task(scan_config: dict[str, Any]) -> str: "target to assess: the instructions below are the only source of " "truth for what to do." ) + elif not parts and user_instructions: + # Neither a target nor a directory, but there is an instruction: the user + # declined the mount, so the instruction is all there is. Say so, or the + # agent goes looking for a scope that was never given. + parts.append( + "\n\nNo scan target and no working directory were provided. The " + "instructions below are the only source of truth for what to do; " + "work from them and from what you can reach yourself." + ) parts.extend(_render_diff_scope(diff_scope)) @@ -192,13 +203,15 @@ def make_model_settings( request_timeout: float | None = None, prompt_cache: bool = True, extra_headers: dict[str, str] | None = None, + has_tools: bool = True, ) -> ModelSettings: + headers = _request_headers(model_name, extra_headers) model_settings = ModelSettings( - parallel_tool_calls=False, + parallel_tool_calls=False if has_tools else None, retry=DEFAULT_MODEL_RETRY, include_usage=True, extra_args=request_timeout_extra_args(request_timeout), - extra_headers=dict(extra_headers) if extra_headers else None, + extra_headers=headers, ) if ( reasoning_effort is not None @@ -221,6 +234,17 @@ def make_model_settings( return model_settings +def _request_headers( + model_name: str, extra_headers: dict[str, str] | None +) -> dict[str, str] | None: + headers: dict[str, str] = {} + if is_openrouter_model(model_name): + headers.update(OPENROUTER_ATTRIBUTION_HEADERS) + if extra_headers: + headers.update(extra_headers) + return headers or None + + def _reasoning_settings( effort: ReasoningEffort, extra_args: dict[str, Any] | None, diff --git a/strix/core/runner.py b/strix/core/runner.py index 6c197b54..8726f819 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import contextlib import io import json @@ -429,7 +430,6 @@ async def run_strix_scan( except BudgetExceededError as exc: logger.info("Scan %s stopped: %s", scan_id, exc) if root_id is not None: - await coordinator.cancel_descendants(root_id) with contextlib.suppress(Exception): await coordinator.set_status(root_id, "stopped") return None @@ -442,19 +442,28 @@ async def run_strix_scan( scan_id, ) if root_id is not None: - await coordinator.cancel_descendants(root_id) with contextlib.suppress(Exception): await coordinator.set_status(root_id, "stopped") return None + except (asyncio.CancelledError, KeyboardInterrupt): + logger.info("Scan %s interrupted by the user", scan_id) + if root_id is not None: + with contextlib.suppress(Exception): + await coordinator.set_status(root_id, "running") + raise except BaseException: logger.exception("Strix scan %s failed", scan_id) if root_id is not None: - await coordinator.cancel_descendants(root_id) with contextlib.suppress(Exception): await coordinator.set_status(root_id, "failed") raise finally: configure_spill_writer(None) + # Settle descendants before closing sessions: on a clean finish a child + # can still be mid-turn, and closing its session underneath it crashes it. + if root_id is not None: + with contextlib.suppress(Exception): + await coordinator.cancel_descendants(root_id) for s in sessions_to_close: with contextlib.suppress(Exception): s.close() diff --git a/strix/core/sessions.py b/strix/core/sessions.py index 8879d359..9286b662 100644 --- a/strix/core/sessions.py +++ b/strix/core/sessions.py @@ -4,6 +4,8 @@ from __future__ import annotations import asyncio import logging +import sqlite3 +from contextlib import contextmanager from typing import TYPE_CHECKING, Any, cast from weakref import WeakKeyDictionary @@ -12,7 +14,7 @@ from agents.memory import SQLiteSession if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Iterator from pathlib import Path from agents.items import TResponseInputItem @@ -22,9 +24,25 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) +class _PooledConnectionSession(SQLiteSession): + @contextmanager + def _locked_connection(self) -> Iterator[sqlite3.Connection]: + with self._lock: + if self._closed: + raise RuntimeError("SQLiteSession is closed") + if self._is_memory_db: + yield self._shared_connection + return + connection = sqlite3.connect(str(self.db_path), check_same_thread=False) + try: + yield connection + finally: + connection.close() + + def open_agent_session(agent_id: str, path: Path) -> SQLiteSession: path.parent.mkdir(parents=True, exist_ok=True) - return SQLiteSession(session_id=agent_id, db_path=path) + return _PooledConnectionSession(session_id=agent_id, db_path=path) async def seed_initial_input(session: Session, initial_input: Any) -> bool: diff --git a/strix/interface/cli_args.py b/strix/interface/cli_args.py index ec4082cc..6c672437 100644 --- a/strix/interface/cli_args.py +++ b/strix/interface/cli_args.py @@ -328,10 +328,11 @@ def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser parser.error(f"--resume {args.resume}: run.json unreadable: {exc}") args.targets_info = state.get("targets_info") or [] - # A target-less run has no targets_info at all: it works in a mounted - # directory, driven by its instruction. + # A target-less run has no targets_info at all. It is driven by its + # instruction, over a mounted working directory or over nothing when the + # mount was declined, so either of those is enough to resume it. workspace_mount = state.get("workspace_mount") or None - if not args.targets_info and not workspace_mount: + if not args.targets_info and not workspace_mount and not state.get("user_instruction"): parser.error(f"--resume {args.resume}: run.json has no targets_info") for target in args.targets_info: diff --git a/strix/interface/main.py b/strix/interface/main.py index 4da8f089..6eb26920 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -224,6 +224,7 @@ async def warm_up_llm(show_model_warning: bool = True) -> None: request_timeout=llm.timeout, prompt_cache=False, extra_headers=settings.dedupe.extra_headers, + has_tools=False, ) if deduper_extra: merged = {**(deduper_settings.extra_args or {}), **deduper_extra} diff --git a/strix/interface/scan_setup.py b/strix/interface/scan_setup.py index 4c9abe6c..599dc1e1 100644 --- a/strix/interface/scan_setup.py +++ b/strix/interface/scan_setup.py @@ -78,6 +78,7 @@ async def preflight_model_connection( request_timeout=resolved_settings.llm.timeout, prompt_cache=False, extra_headers=resolved_settings.llm.extra_headers, + has_tools=False, ) await asyncio.wait_for( model.get_response( diff --git a/strix/interface/tui/backend/controller.py b/strix/interface/tui/backend/controller.py index 3784604d..d74bbba6 100644 --- a/strix/interface/tui/backend/controller.py +++ b/strix/interface/tui/backend/controller.py @@ -138,13 +138,6 @@ class TuiController: self.error = detail self.notify_changed() - def enter_setup(self) -> None: - """Return a session to the start screen, e.g. on a declined mount.""" - self.setup_mode = True - self.scan_started = False - self.scan_state = "setup" - self.notify_changed() - def add_message(self, text: str, level: str = "info") -> None: self._append_message(text, level) self.notify_changed() @@ -356,14 +349,12 @@ class TuiController: if not isinstance(approved, bool): raise TypeError("approved must be a boolean") self.pending_workspace_mount = None - if not approved: - # Nothing was prepared, so return to the start screen untouched. - self.workspace_mount = None - self.enter_setup() - return {"approved": False} - self.workspace_mount = mount + # Declining skips the mount, it does not abandon the scan. The prompt is + # the whole of the input either way; the working directory is only an + # extra the agent may look at, so the run goes ahead without one. + self.workspace_mount = mount if approved else None await self._begin_scan(self._pending_verify) - return {"approved": True} + return {"approved": approved} async def _send_message(self, payload: dict[str, Any]) -> dict[str, Any]: agent_id = self._required_string(payload, "agent_id") diff --git a/strix/interface/tui/internal/app/setup.go b/strix/interface/tui/internal/app/setup.go index 4a6c4b18..a7525fbf 100644 --- a/strix/interface/tui/internal/app/setup.go +++ b/strix/interface/tui/internal/app/setup.go @@ -66,13 +66,10 @@ func (m *Model) submitSetupPrompt(value string) (tea.Model, tea.Cmd) { } // answerMountConfirmation replies to the working-directory mount the backend is -// waiting on. Declining returns to the start screen, so the prompt goes back in -// the composer to be edited or given a target instead. +// waiting on. Either answer starts the scan - declining only means it runs +// without the directory - so the prompt stays with the run rather than coming +// back to the composer. func (m *Model) answerMountConfirmation(approved bool) tea.Cmd { - if !approved && m.pendingPrompt != "" { - m.input.SetValue(m.pendingPrompt) - m.resizeViewport() - } m.pendingPrompt = "" return send(m.client, "setup.confirm_mount", map[string]any{"approved": approved}) } diff --git a/strix/interface/tui/internal/app/setup_prompt_test.go b/strix/interface/tui/internal/app/setup_prompt_test.go index a18f8234..63a0170f 100644 --- a/strix/interface/tui/internal/app/setup_prompt_test.go +++ b/strix/interface/tui/internal/app/setup_prompt_test.go @@ -247,13 +247,10 @@ func TestMountConfirmationAnswers(t *testing.T) { if payload.Approved != tc.approved { t.Fatalf("%s: approved=%v, want %v", tc.name, payload.Approved, tc.approved) } - // Declining returns to the start screen, so the prompt comes back. - want := "" - if !tc.approved { - want = "find auth bugs in the login flow" - } - if got := model.input.Value(); got != want { - t.Fatalf("%s: composer = %q, want %q", tc.name, got, want) + // Either answer launches, so the prompt stays with the run rather than + // coming back to the composer. + if got := model.input.Value(); got != "" { + t.Fatalf("%s: composer = %q, want it cleared", tc.name, got) } if model.pendingPrompt != "" { t.Fatalf("%s: held prompt was not cleared: %q", tc.name, model.pendingPrompt) @@ -290,3 +287,101 @@ func TestSetupPromptWithTargetLaunches(t *testing.T) { t.Fatalf("setup.start (%d) must come after setup.set_instruction (%d): %v", start, instr, types) } } + +// The prompt's buttons are buttons: clicking Cancel has to answer the backend, +// which it could not do while the mouse handler had no case for this modal. +func TestMountPromptButtonsAreClickable(t *testing.T) { + for _, testCase := range []struct { + label string + approved bool + }{ + {mountConfirmLabel, true}, + {mountCancelLabel, false}, + } { + connection := &recordingConn{} + model := New(&Client{conn: connection}) + model.width, model.height = 130, 40 + model.snapshot = protocol.Snapshot{SetupMode: true, WorkingDir: "/Users/me/code/api"} + updated, _ := model.submit("find auth bugs in the login flow") + model = updated.(Model) + connection.Reset() + model.snapshot = protocol.Snapshot{ + ScanStarted: true, ScanState: "preparing", PendingMount: "/Users/me/code/api", + } + model.syncMountPrompt() + + left, top, panel := model.mountPromptBounds() + clicked := false + for row, line := range strings.Split(panel, "\n") { + plain := ansi.Strip(line) + index := strings.Index(plain, testCase.label) + if index < 0 { + continue + } + updated, cmd := model.updateModalMouse(tea.MouseMsg{ + X: left + ansi.StringWidth(plain[:index]) + 1, Y: top + row, + Button: tea.MouseButtonLeft, Action: tea.MouseActionPress, + }) + model = updated.(Model) + envelopes := drainCommands(t, cmd, connection) + if len(envelopes) != 1 || envelopes[0].Type != "setup.confirm_mount" { + t.Fatalf("clicking %s sent %v", testCase.label, commandTypes(envelopes)) + } + var payload struct { + Approved bool `json:"approved"` + } + if err := json.Unmarshal(envelopes[0].Payload, &payload); err != nil { + t.Fatal(err) + } + if payload.Approved != testCase.approved { + t.Fatalf("clicking %s answered approved=%v", testCase.label, payload.Approved) + } + clicked = true + break + } + if !clicked { + t.Fatalf("%s was not found in the prompt", testCase.label) + } + } +} + +// Skipping the mount runs the scan without a directory. It must not throw the +// session back to the start screen, and it must not hand the prompt back: the +// run has it. +func TestSkippingTheMountKeepsTheScanRunning(t *testing.T) { + connection := &recordingConn{} + model := New(&Client{conn: connection}) + model.width, model.height = 130, 40 + model.snapshot = protocol.Snapshot{SetupMode: true, WorkingDir: "/Users/me/code/api"} + updated, _ := model.submit("find auth bugs in the login flow") + model = updated.(Model) + model.snapshot = protocol.Snapshot{ + ScanStarted: true, ScanState: "preparing", PendingMount: "/Users/me/code/api", + } + model.syncMountPrompt() + if model.modal != modalConfirmMount { + t.Fatal("the prompt did not open") + } + + model.modalChoice = 1 + updated, _ = model.updateModal(tea.KeyMsg{Type: tea.KeyEnter}) + model = updated.(Model) + + // The backend answers by starting the scan with no mount. + model.handleEnvelope(stateEnvelope(t, 2, protocol.Snapshot{ + ScanStarted: true, ScanState: "running", + })) + + if model.modal != modalNone { + t.Fatalf("the prompt is still open: %v", model.modal) + } + if model.snapshot.SetupMode { + t.Fatal("skipping the mount fell back to the start screen") + } + if got := model.input.Value(); got != "" { + t.Fatalf("the prompt came back to the composer: %q", got) + } + if model.pendingPrompt != "" { + t.Fatalf("the held prompt was not released: %q", model.pendingPrompt) + } +} diff --git a/strix/interface/tui/internal/app/update.go b/strix/interface/tui/internal/app/update.go index 2a22c4a9..692a83b5 100644 --- a/strix/interface/tui/internal/app/update.go +++ b/strix/interface/tui/internal/app/update.go @@ -474,6 +474,18 @@ func (m Model) updateModalMouse(msg tea.MouseMsg) (tea.Model, tea.Cmd) { m.modalChoice = 1 return m.updateModal(tea.KeyMsg{Type: tea.KeyEnter}) } + case modalConfirmMount: + left, top, panel := m.mountPromptBounds() + if labelHitAt(panel, mountConfirmLabel, left, top, msg.X, msg.Y) { + m.modalChoice = 0 + cmd := m.answerMountConfirmation(true) + return m, cmd + } + if labelHitAt(panel, mountCancelLabel, left, top, msg.X, msg.Y) { + m.modalChoice = 1 + cmd := m.answerMountConfirmation(false) + return m, cmd + } case modalVulnerability: for _, button := range m.reportButtons() { if button == reportCopy || button == reportDone { @@ -507,7 +519,14 @@ func (m Model) centeredViewBounds(view string) (left, top, width, height int) { func (m Model) centeredLabelHit(view, label string, x, y int) bool { left, top, _, _ := m.centeredViewBounds(view) - for row, line := range strings.Split(view, "\n") { + return labelHitAt(view, label, left, top, x, y) +} + +// labelHitAt reports whether a click landed on a label drawn in a panel whose +// top-left corner is at (left, top). The mount prompt is docked in a corner +// rather than centered, so it cannot use the centered bounds. +func labelHitAt(panel, label string, left, top, x, y int) bool { + for row, line := range strings.Split(panel, "\n") { plain := ansi.Strip(line) index := strings.Index(plain, label) if index < 0 || y != top+row { @@ -593,7 +612,8 @@ func (m Model) updateModal(key tea.KeyMsg) (tea.Model, tea.Cmd) { case "esc": if m.modal == modalConfirmMount { // The backend is waiting on an answer; escape declines it. - return m, m.answerMountConfirmation(false) + cmd := m.answerMountConfirmation(false) + return m, cmd } m.closeModal() return m, nil @@ -604,7 +624,10 @@ func (m Model) updateModal(key tea.KeyMsg) (tea.Model, tea.Cmd) { modal, choice := m.modal, m.modalChoice if modal == modalConfirmMount { // The snapshot closes this prompt once the backend has the answer. - return m, m.answerMountConfirmation(choice == 0) + // Bound to a variable first: the call restores the held prompt into + // the composer, and that has to be in the model being returned. + cmd := m.answerMountConfirmation(choice == 0) + return m, cmd } m.closeModal() if choice == 1 { diff --git a/strix/interface/tui/internal/app/view.go b/strix/interface/tui/internal/app/view.go index 9ac0a55e..78a3d9f5 100644 --- a/strix/interface/tui/internal/app/view.go +++ b/strix/interface/tui/internal/app/view.go @@ -271,6 +271,23 @@ func (m Model) viewInner() string { return m.toastOverlay(main) } +// mountPromptBounds is where the working-directory prompt is drawn. It is placed +// by cornerOverlay rather than centered, so a click has to be tested against +// these bounds and not the ones the other modals use. +func (m Model) mountPromptBounds() (left, top int, panel string) { + panel = m.modalView() + if panel == "" { + return 0, 0, "" + } + _, _, chatWidth, _ := m.layout() + left = max(0, min(chatWidth, m.width)-lipgloss.Width(panel)) + statusH := 0 + if m.statusVisible() { + statusH = 1 + } + return left, max(0, m.inputTop()-statusH-lipgloss.Height(panel)), panel +} + // cornerOverlay splices a panel in directly above the composer, right-aligned // with it, leaving the rest of the view visible behind it. func (m Model) cornerOverlay(view, panel string) string { diff --git a/strix/interface/tui/internal/app/vulnerabilities.go b/strix/interface/tui/internal/app/vulnerabilities.go index b8d988bf..bd5fa1e6 100644 --- a/strix/interface/tui/internal/app/vulnerabilities.go +++ b/strix/interface/tui/internal/app/vulnerabilities.go @@ -220,6 +220,13 @@ func (m Model) confirmView(title string, width int, border, titleColor lipgloss. return m.confirmDialog(title, "", width, border, titleColor, red, "Yes", "No") } +// The mount prompt's buttons, named so the renderer and the click test cannot +// drift apart. +const ( + mountConfirmLabel = "Mount" + mountCancelLabel = "Skip" +) + // mountConfirmView asks before a target-less scan mounts the working directory. // It is a compact prompt docked in the corner of the live view: nothing is // prepared until it is answered, and the directory is a workspace rather than a @@ -232,8 +239,8 @@ func (m Model) mountConfirmView() string { } title := render.Bold(amber).Render("△ Mount working directory?") body := render.Col(white).Render(truncatePath(dir, width-4)) + "\n" + - render.Dim().Render("writable in the sandbox") - return m.cornerPrompt(title, body, width, "Confirm", "Cancel") + render.Dim().Render("writable in the sandbox · skip to run without it") + return m.cornerPrompt(title, body, width, mountConfirmLabel, mountCancelLabel) } // truncatePath keeps the tail of a path visible, which is the part that diff --git a/strix/llm/compaction.py b/strix/llm/compaction.py index 09d36454..e40caf6a 100644 --- a/strix/llm/compaction.py +++ b/strix/llm/compaction.py @@ -294,6 +294,7 @@ async def _summarize(model: str, prompt: str, max_tokens: int) -> str | None: request_timeout=llm.timeout, prompt_cache=False, extra_headers=llm.extra_headers, + has_tools=False, ).resolve(ModelSettings(max_tokens=max_tokens)) try: response = ( diff --git a/strix/report/dedupe.py b/strix/report/dedupe.py index f848a6d6..1cc0a66a 100644 --- a/strix/report/dedupe.py +++ b/strix/report/dedupe.py @@ -62,6 +62,7 @@ def _dedupe_model_settings( # must never receive the main endpoint's credentials. A dedicated model # gets its own DEDUPE_LLM_EXTRA_HEADERS instead. extra_headers=dedupe.extra_headers if dedupe.model else llm.extra_headers, + has_tools=False, ) extra = _dedupe_extra_args(dedupe) if extra: diff --git a/strix/report/pricing.py b/strix/report/pricing.py new file mode 100644 index 00000000..57c89959 --- /dev/null +++ b/strix/report/pricing.py @@ -0,0 +1,54 @@ +"""LiteLLM model-name resolution for local cost estimates.""" + +from __future__ import annotations + +from functools import lru_cache +from typing import Any, cast + + +@lru_cache(maxsize=512) +def resolve_litellm_model(model: str) -> str | None: + """Return a provider-qualified model name that LiteLLM can price.""" + try: + import litellm + + normalized = model.strip() + for prefix in ("litellm/", "any-llm/", "openai/"): + if normalized.startswith(prefix): + normalized = normalized.removeprefix(prefix) + break + if not normalized: + return None + + model_cost = cast( + "dict[str, dict[str, Any]]", + getattr(litellm, "model_cost"), # noqa: B009 + ) + bare_entry = model_cost.get(normalized) + if "/" not in normalized and isinstance(bare_entry, dict): + provider = bare_entry.get("litellm_provider") + if isinstance(provider, str) and provider: + return f"{provider}/{normalized}" + if "/" in normalized and isinstance(bare_entry, dict): + return normalized + + names = [normalized] + if "/" in normalized: + names.append(normalized.rsplit("/", 1)[-1]) + for name in names: + matches = sorted(key for key in model_cost if key.endswith(f"/{name}")) + if not matches: + continue + prices = { + ( + model_cost[key].get("input_cost_per_token"), + model_cost[key].get("output_cost_per_token"), + ) + for key in matches + if isinstance(model_cost.get(key), dict) + } + if len(matches) == 1 or len(prices) == 1: + return matches[0] + return None # noqa: TRY300 + except Exception: # noqa: BLE001 + return None diff --git a/strix/report/state.py b/strix/report/state.py index d4d60ea6..f7ca890c 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -14,6 +14,7 @@ from agents.usage import Usage from strix.config import subscription from strix.config.loader import load_settings from strix.core.paths import run_dir_for +from strix.report.pricing import resolve_litellm_model from strix.report.sarif import write_sarif from strix.report.usage import LLMUsageLedger from strix.report.writer import ( @@ -698,10 +699,13 @@ def _estimate_response_cost(kwargs: Any, completion_response: Any) -> float | No candidates.append(model.rsplit("/", 1)[-1]) for candidate in candidates: + resolved = resolve_litellm_model(candidate) + if not resolved: + continue try: value = completion_cost( - completion_response={"model": candidate, "usage": usage_payload}, - model=candidate, + completion_response={"model": resolved, "usage": usage_payload}, + model=resolved, ) except Exception: # nosec B112 # noqa: BLE001, S112 continue diff --git a/strix/report/usage.py b/strix/report/usage.py index e3ddf494..3d6be050 100644 --- a/strix/report/usage.py +++ b/strix/report/usage.py @@ -7,6 +7,8 @@ from typing import Any from agents.usage import Usage, deserialize_usage, serialize_usage +from strix.report.pricing import resolve_litellm_model + logger = logging.getLogger(__name__) @@ -18,7 +20,9 @@ class LLMUsageLedger: self._total_usage = Usage() self._agent_usage: dict[str, Usage] = {} self._agent_metadata: dict[str, dict[str, str]] = {} - self._total_cost = 0.0 + self._observed_cost = 0.0 + self._estimated_cost = 0.0 + self._has_observed_cost = False # When True, tokens are still tracked but cost stays $0 — the run is on a # model subscription, so there is no metered per-token charge to report. self.zero_cost = False @@ -44,10 +48,10 @@ class LLMUsageLedger: if model: metadata["model"] = model - if not self.zero_cost and not _is_litellm_routed(model): + if not self.zero_cost: estimated = _estimate_litellm_cost(usage, model) if estimated: - self._total_cost += estimated + self._estimated_cost += estimated return True @@ -55,15 +59,18 @@ class LLMUsageLedger: if self.zero_cost: return if isinstance(cost, int | float) and cost > 0: - self._total_cost += float(cost) + self._observed_cost += float(cost) + self._has_observed_cost = True @property def total_cost(self) -> float: - return _round_cost(self._total_cost) + if self.zero_cost: + return 0.0 + return _round_cost(self._observed_cost if self._has_observed_cost else self._estimated_cost) def to_record(self) -> dict[str, Any]: record = serialize_usage(self._total_usage) - record["cost"] = _round_cost(self._total_cost) + record["cost"] = self.total_cost record["agents"] = [] agent_tokens = {aid: _resolve_total_tokens(u) for aid, u in self._agent_usage.items()} @@ -72,7 +79,7 @@ class LLMUsageLedger: usage = self._agent_usage[agent_id] metadata = self._agent_metadata.get(agent_id, {}) agent_cost = ( - self._total_cost * (agent_tokens[agent_id] / total_tokens) if total_tokens else 0.0 + self.total_cost * (agent_tokens[agent_id] / total_tokens) if total_tokens else 0.0 ) agent_record = serialize_usage(usage) @@ -92,7 +99,9 @@ class LLMUsageLedger: self._total_usage = Usage() self._agent_usage.clear() self._agent_metadata.clear() - self._total_cost = 0.0 + self._observed_cost = 0.0 + self._estimated_cost = 0.0 + self._has_observed_cost = False if not isinstance(raw_usage, dict): return @@ -103,7 +112,9 @@ class LLMUsageLedger: logger.exception("Failed to hydrate aggregate llm_usage from run.json") self._total_usage = Usage() - self._total_cost = _float_or_zero(raw_usage.get("cost")) + persisted_cost = _float_or_zero(raw_usage.get("cost")) + self._observed_cost = persisted_cost + self._estimated_cost = persisted_cost for raw_agent in raw_usage.get("agents") or []: if not isinstance(raw_agent, dict): @@ -136,15 +147,6 @@ def _resolve_total_tokens(usage: Usage) -> int: return prompt + completion -def _is_litellm_routed(model: str | None) -> bool: - if not model: - return False - name = model.strip().lower() - if "/" not in name: - return False - return not name.startswith("openai/") - - def _usage_has_activity(usage: Usage) -> bool: return bool( usage.requests @@ -201,24 +203,23 @@ def _estimate_litellm_entry_cost(entry: Any, model: str) -> float | None: candidates = [model] if "/" in model: - candidates.append(model.split("/", 1)[-1]) + candidates.append(model.rsplit("/", 1)[-1]) - cost: Any = None for candidate in candidates: + resolved = resolve_litellm_model(candidate) + if not resolved: + continue try: cost = completion_cost( - completion_response={"model": candidate, "usage": usage_payload}, - model=model, + completion_response={"model": resolved, "usage": usage_payload}, + model=resolved, ) - break except Exception: # nosec B112 # noqa: BLE001, S112 continue - - if cost is None: - logger.debug("LiteLLM cost estimate unavailable for model %s", model) - return None - - return cost if isinstance(cost, int | float) and cost >= 0 else None + if cost > 0: + return float(cost) + logger.debug("LiteLLM cost estimate unavailable for model %s", model) + return None def _litellm_model_name(model: str | None) -> str | None: diff --git a/strix/skills/tooling/agent_browser.md b/strix/skills/tooling/agent_browser.md index db074e86..a254bfaf 100644 --- a/strix/skills/tooling/agent_browser.md +++ b/strix/skills/tooling/agent_browser.md @@ -58,6 +58,26 @@ agent-browser screenshot The browser stays running across commands so these feel like a single session. Use `agent-browser close` (or `close --all`) when you're done. +The default session is **shared with every other agent in the sandbox** — if +another agent navigates it, your page and your refs are gone from under you. So +claim your own by passing `--session ` on **every** command: + +```bash +agent-browser --session recon-3 open https://example.com +agent-browser --session recon-3 snapshot -i +agent-browser --session recon-3 close # when done with the target +``` + +The examples in the rest of this skill omit `--session` to keep them readable; +keep passing yours. Each session is a separate Chromium (~340 MB) on a shared +box, so hold one rather than several, and close it when you're finished. + +A browser left idle for 3 minutes is reclaimed automatically to free memory for +the other agents; the next command relaunches it, but the page, tabs, refs and +cookies are gone. If you're authenticated and about to go do something else for a +while, save the state first (see +[Persist session across runs](#persist-session-across-runs)). + ## Reading a page ```bash @@ -307,6 +327,16 @@ agent-browser --session b fill @e1 "bob@test.com" `AGENT_BROWSER_SESSION=myapp` sets the default session for the current shell. +Use a session named after yourself for your own work — that's what keeps a +concurrent agent from navigating the page out from under you. Every session is a +separate Chromium though, so hold one at a time rather than a collection, and +close each one when its flow is finished: + +```bash +agent-browser --session a close +agent-browser --session b close +``` + ### Mock network requests ```bash @@ -368,8 +398,11 @@ 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; later commands reuse it. A daemon left idle for 3 minutes shuts itself +down to free memory for the other agents, so an `open` after a long gap is a +fresh browser rather than a resumed one — expect to re-navigate, and re-`state +load` if you were logged in. Distinguish the 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 diff --git a/strix/tools/respond/tool.py b/strix/tools/respond/tool.py index 796050c3..12538f0c 100644 --- a/strix/tools/respond/tool.py +++ b/strix/tools/respond/tool.py @@ -15,7 +15,7 @@ def _ctx(ctx: RunContextWrapper) -> dict[str, Any]: @function_tool -async def respond_to_user(ctx: RunContextWrapper, message: str) -> str: +async def respond_to_user(ctx: RunContextWrapper, message: str = "") -> str: """Answer the user and hand control back to them. This is the ONLY way to yield to the user. Delivering the message and @@ -45,6 +45,10 @@ async def respond_to_user(ctx: RunContextWrapper, message: str) -> str: have followed the tool calls that led here. Lead with the answer or the decision you need, and if you are blocked, say exactly what you need from them. + + Omit it when you have just said your piece as plain text and + only need to wait: that text has already reached them, and + repeating it makes them read the same answer twice. """ inner = _ctx(ctx) coordinator = coordinator_from_context(inner) diff --git a/strix/tools/todo/tools.py b/strix/tools/todo/tools.py index 07761f5f..135f3bc0 100644 --- a/strix/tools/todo/tools.py +++ b/strix/tools/todo/tools.py @@ -110,12 +110,19 @@ def _get_agent_todos(agent_id: str) -> dict[str, dict[str, Any]]: def _normalize_priority(priority: str | None, default: str = "normal") -> str: - candidate = (priority or default or "normal").lower() + candidate = str(priority or default or "normal").strip().lower() if candidate not in VALID_PRIORITIES: raise ValueError(f"Invalid priority. Must be one of: {', '.join(VALID_PRIORITIES)}") return candidate +def _coerce_priority(priority: str | None, default: str = "normal") -> str: + try: + return _normalize_priority(priority, default) + except ValueError: + return default + + def _sorted_todos(agent_id: str) -> list[dict[str, Any]]: todos_list = [ {**todo, "todo_id": todo_id} for todo_id, todo in _get_agent_todos(agent_id).items() @@ -285,11 +292,16 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str: - ``description`` (str, optional): extra context or acceptance criteria. - ``priority`` (str, optional): one of ``"low"`` / - ``"normal"`` / ``"high"`` / ``"critical"``. Defaults to - ``"normal"``. + ``"normal"`` / ``"high"`` / ``"critical"``. Anything else, + including omitting it, falls back to ``"normal"`` rather + than failing. Example: ``[{"title": "Probe /admin", "priority": "high"}, {"title": "Check JWT alg=none"}]``. + + A title already on the list, or repeated within this call, is + skipped rather than duplicated; skipped titles come back under + ``skipped``. """ agent_id = _agent_id_from(ctx) try: @@ -302,13 +314,21 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str: ) agent_todos = _get_agent_todos(agent_id) + seen = {todo["title"].strip().lower() for todo in agent_todos.values()} created: list[dict[str, Any]] = [] + skipped: list[dict[str, str]] = [] for task in tasks: - task_priority = _normalize_priority(task.get("priority")) + title = task["title"] + key = title.lower() + if key in seen: + skipped.append({"title": title, "reason": "duplicate title"}) + continue + seen.add(key) + task_priority = _coerce_priority(task.get("priority")) todo_id = str(uuid.uuid4())[:6] timestamp = datetime.now(UTC).isoformat() agent_todos[todo_id] = { - "title": task["title"], + "title": title, "description": task.get("description"), "priority": task_priority, "status": "pending", @@ -316,7 +336,7 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str: "updated_at": timestamp, "completed_at": None, } - created.append({"todo_id": todo_id, "title": task["title"], "priority": task_priority}) + created.append({"todo_id": todo_id, "title": title, "priority": task_priority}) except (ValueError, TypeError) as e: return json.dumps( {"success": False, "error": f"Failed to create todo: {e}"}, @@ -330,6 +350,7 @@ async def create_todo(ctx: RunContextWrapper, todos: str) -> str: "success": True, "created": created, "created_count": len(created), + "skipped": skipped, "todos": _sorted_todos(agent_id), "total_count": len(_get_agent_todos(agent_id)), }, diff --git a/tests/test_agent_factory_tool_arguments.py b/tests/test_agent_factory_tool_arguments.py index 49dabf5d..70908f26 100644 --- a/tests/test_agent_factory_tool_arguments.py +++ b/tests/test_agent_factory_tool_arguments.py @@ -70,7 +70,6 @@ async def test_encoded_list_is_decoded_for_an_array_parameter(schema: dict[str, "auth", "Endpoint /admin leaks user data, and session tokens never expire", '"auth"', - "", ], ) async def test_free_form_strings_are_never_split_into_an_array(value: str) -> None: @@ -79,6 +78,29 @@ async def test_free_form_strings_are_never_split_into_an_array(value: str) -> No assert parsed["tags"] == value +@pytest.mark.asyncio +@pytest.mark.parametrize("schema", [_ARRAY, _NULLABLE_ARRAY]) +@pytest.mark.parametrize("value", ["", " "]) +async def test_empty_string_becomes_an_empty_array(schema: dict[str, Any], value: str) -> None: + parsed = await _roundtrip(schema, {"tags": value}) + + assert parsed["tags"] == [] + + +@pytest.mark.asyncio +async def test_empty_string_becomes_an_empty_object() -> None: + parsed = await _roundtrip(_OBJECT, {"modifications": ""}) + + assert parsed["modifications"] == {} + + +@pytest.mark.asyncio +async def test_empty_string_for_a_string_parameter_is_untouched() -> None: + parsed = await _roundtrip(_STRING, {"todos": ""}) + + assert parsed["todos"] == "" + + @pytest.mark.asyncio async def test_encoded_mapping_is_decoded_for_an_object_parameter() -> None: parsed = await _roundtrip(_OBJECT, {"modifications": '{"method": "POST"}'}) diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index 30d4db44..6db31145 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -143,7 +143,7 @@ def test_cost_callback_estimates_cost_with_bare_model_fallback() -> None: } def fake_completion_cost(**kwargs: object) -> float: - if kwargs["model"] == "gpt-4o-mini": + if kwargs["model"] == "openai/gpt-4o-mini": return 0.025 raise ValueError(kwargs["model"]) diff --git a/tests/test_execution.py b/tests/test_execution.py index 6364f801..d389bde3 100644 --- a/tests/test_execution.py +++ b/tests/test_execution.py @@ -1228,3 +1228,40 @@ async def test_wait_kind_survives_a_snapshot_round_trip() -> None: assert restored.wait_kinds["root"] == "user" assert restored.idle_resume_counts["root"] == 1 assert await execution._plain_waiting_timeout(restored, "root") is None + + +@pytest.mark.asyncio +async def test_interactive_nudge_offers_waiting_without_repeating() -> None: + """The nudge is the instruction an agent reads when it is stranded here. + + It is where the option to wait on what was already said has to be, not only + in the system prompt: an agent that ended a turn on plain text reasons off + this text, and without the clause it restates its answer to reach a tool + call, so the user reads it twice. + + The clause holds whatever the turn did, because the agent is the one who + knows whether it spoke — this fires for a turn that produced no text at all. + """ + items = await execution._append_tool_required_message( + session=None, + context={"parent_id": None}, + attempt=1, + limit=3, + interactive=True, + ) + + assert "with no message if you have already said it" in items[0]["content"] + + +@pytest.mark.asyncio +async def test_autonomous_nudge_does_not_offer_the_user() -> None: + """There is nobody attached to an autonomous run to wait for.""" + items = await execution._append_tool_required_message( + session=None, + context={"parent_id": None}, + attempt=1, + limit=3, + interactive=False, + ) + + assert "respond_to_user" not in items[0]["content"] diff --git a/tests/test_inputs.py b/tests/test_inputs.py index ed233262..2ff9a603 100644 --- a/tests/test_inputs.py +++ b/tests/test_inputs.py @@ -299,6 +299,16 @@ def test_make_model_settings_forces_required_for_anyllm_routed_openai_model() -> assert settings.tool_choice == "required" +def test_make_model_settings_disables_parallel_tool_calls_by_default() -> None: + assert make_model_settings("none", model_name="gpt-4o").parallel_tool_calls is False + + +def test_make_model_settings_omits_parallel_tool_calls_without_tools() -> None: + settings = make_model_settings("none", model_name="gpt-4o", has_tools=False) + + assert settings.parallel_tool_calls is None + + def test_make_model_settings_sets_request_timeout() -> None: settings = make_model_settings( "none", @@ -351,3 +361,32 @@ def test_make_model_settings_timeout_survives_reasoning_resolve() -> None: assert settings.extra_args is not None assert settings.extra_args["timeout"] == 120.0 + + +def test_openrouter_attribution_rides_on_the_request_headers() -> None: + # litellm.headers is ignored once a request carries any header of its own, + # so the attribution must be part of the per-request headers. + headers = make_model_settings( + None, model_name="openrouter/anthropic/claude-sonnet-4-5" + ).extra_headers + assert headers == { + "HTTP-Referer": "https://strix.ai", + "X-Title": "Strix", + "X-OpenRouter-Categories": "cli-agent", + } + + +def test_openrouter_attribution_absent_for_other_providers() -> None: + assert make_model_settings(None, model_name="anthropic/claude-sonnet-4-5").extra_headers is None + + +def test_user_headers_override_openrouter_attribution() -> None: + headers = make_model_settings( + None, + model_name="openrouter/anthropic/claude-sonnet-4-5", + extra_headers={"X-Title": "Custom", "X-Tenant": "acme"}, + ).extra_headers + assert headers is not None + assert headers["X-Title"] == "Custom" + assert headers["X-Tenant"] == "acme" + assert headers["HTTP-Referer"] == "https://strix.ai" diff --git a/tests/test_pricing.py b/tests/test_pricing.py new file mode 100644 index 00000000..abff873d --- /dev/null +++ b/tests/test_pricing.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +from unittest.mock import patch + +import litellm +from agents.usage import Usage + +from strix.report.pricing import resolve_litellm_model +from strix.report.usage import LLMUsageLedger + + +def test_resolves_common_bare_model_names() -> None: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("deepseek-v4-flash") == "deepseek/deepseek-v4-flash" + assert resolve_litellm_model("openai/deepseek-v4-flash") == "deepseek/deepseek-v4-flash" + assert resolve_litellm_model("grok-4.5") == "xai/grok-4.5" + assert resolve_litellm_model("MiniMax-M3") == "minimax/MiniMax-M3" + + +def test_resolver_returns_none_for_unresolvable_model() -> None: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("provider/not-a-real-model") is None + + +def test_ledger_uses_estimate_when_routed_provider_reports_no_cost() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + + with patch("litellm.completion_cost", return_value=0.42): + ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash") + + assert ledger.total_cost == 0.42 + + +def test_ledger_prefers_observed_cost_over_estimate() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + + with patch("litellm.completion_cost", return_value=0.42): + ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash") + ledger.record_observed_cost(0.17) + + assert ledger.total_cost == 0.17 + + +def test_hydrated_estimate_continues_accumulating_new_estimates() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + ledger.hydrate({"cost": 0.42}) + + with patch("litellm.completion_cost", return_value=0.17): + ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash") + + assert ledger.total_cost == 0.59 + + +def test_zero_cost_disables_both_observed_and_estimated_costs() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + ledger.zero_cost = True + + with patch("litellm.completion_cost", return_value=0.42) as estimate: + ledger.record(agent_id="a", usage=usage, model="deepseek-v4-flash") + ledger.record_observed_cost(1.0) + + estimate.assert_not_called() + assert ledger.total_cost == 0.0 + + +def test_resolver_uses_provider_when_bare_entry_has_one() -> None: + original = litellm.model_cost + litellm.model_cost = { + "example": { + "litellm_provider": "example-provider", + "input_cost_per_token": 1.0, + "output_cost_per_token": 2.0, + } + } + try: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("example") == "example-provider/example" + finally: + litellm.model_cost = original + resolve_litellm_model.cache_clear() + + +def test_resolver_does_not_guess_between_differently_priced_providers() -> None: + original = litellm.model_cost + litellm.model_cost = { + "provider-a/example": { + "input_cost_per_token": 1.0, + "output_cost_per_token": 2.0, + }, + "provider-b/example": { + "input_cost_per_token": 3.0, + "output_cost_per_token": 4.0, + }, + } + try: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("example") is None + finally: + litellm.model_cost = original + resolve_litellm_model.cache_clear() diff --git a/tests/test_respond_to_user.py b/tests/test_respond_to_user.py index 254bd357..fc59c267 100644 --- a/tests/test_respond_to_user.py +++ b/tests/test_respond_to_user.py @@ -64,3 +64,31 @@ async def test_a_message_that_already_arrived_is_taken_instead_of_parking() -> N assert result["wait_outcome"] == "message_arrived" assert result["pending_messages"] == 1 assert coordinator.statuses["root"] == "running" + + +async def _call_without_message(context: dict[str, Any]) -> dict[str, Any]: + ctx = ToolContext( + context=context, + tool_name="respond_to_user", + tool_call_id="call-1", + tool_arguments="{}", + ) + raw = await respond_to_user.on_invoke_tool(ctx, "{}") + return json.loads(raw) # type: ignore[no-any-return] + + +@pytest.mark.asyncio +async def test_parks_without_a_message() -> None: + """An agent that has already said its piece as plain text can just wait. + + The nudge is what leaves it here, and while a message was required the only + way to stop was to send the same answer a second time. + """ + context = await _context(interactive=True) + + result = await _call_without_message(context) + + assert result["success"] is True + assert result["wait_outcome"] == "waiting" + assert result["message"] == "" + assert context["coordinator"].statuses["root"] == "waiting" diff --git a/tests/test_runner_interrupt.py b/tests/test_runner_interrupt.py new file mode 100644 index 00000000..054f5b7a --- /dev/null +++ b/tests/test_runner_interrupt.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +import asyncio +import types +from typing import Any + +import pytest +from agents import ModelSettings + +import strix.tools.notes.tools as notes_tools +import strix.tools.todo.tools as todo_tools +from strix.core import runner +from strix.core.agents import AgentCoordinator +from strix.runtime import session_manager + + +def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: + monkeypatch.setattr(runner, "run_dir_for", lambda _scan_id: tmp_path) + monkeypatch.setattr(runner, "runtime_state_dir", lambda _run_dir: tmp_path) + monkeypatch.setattr(runner, "setup_scan_logging", lambda _run_dir: lambda: None) + monkeypatch.setattr(runner, "set_scan_id", lambda _scan_id: None) + + settings = types.SimpleNamespace( + llm=types.SimpleNamespace( + model="openai/gpt-4o", + reasoning_effort="high", + force_required_tool_choice=False, + timeout=300, + prompt_cache=True, + extra_headers=None, + ), + runtime=types.SimpleNamespace(max_context_images=3), + ) + monkeypatch.setattr(runner, "load_settings", lambda: settings) + monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _settings: None) + monkeypatch.setattr( + runner, "uses_chat_completions_tool_schema", lambda _model, _settings: False + ) + monkeypatch.setattr(todo_tools, "hydrate_todos_from_disk", lambda _state_dir: None) + monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _state_dir: None) + + async def _create_or_reuse(*_args: Any, **_kwargs: Any) -> dict[str, Any]: + return {"client": object(), "session": object(), "caido_client": None} + + async def _cleanup(*_args: Any, **_kwargs: Any) -> None: + return None + + monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse) + monkeypatch.setattr(session_manager, "cleanup", _cleanup) + monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task") + monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: "") + monkeypatch.setattr(runner, "make_model_settings", lambda *_a, **_k: ModelSettings()) + monkeypatch.setattr(runner, "build_strix_agent", lambda **_kwargs: object()) + monkeypatch.setattr(runner, "make_child_factory", lambda **_kwargs: lambda **_k: object()) + monkeypatch.setattr(runner, "open_agent_session", lambda _root_id, _db: object()) + + +def _root_status(coordinator: AgentCoordinator) -> str: + roots = [aid for aid, parent in coordinator.parent_of.items() if parent is None] + assert len(roots) == 1 + return coordinator.statuses[roots[0]] + + +@pytest.mark.parametrize("interrupt", [KeyboardInterrupt, asyncio.CancelledError]) +@pytest.mark.asyncio +async def test_user_interrupt_leaves_the_root_running_for_resume( + monkeypatch: pytest.MonkeyPatch, tmp_path: Any, interrupt: type[BaseException] +) -> None: + _wire_runner(monkeypatch, tmp_path) + + async def _interrupt(*_args: Any, **_kwargs: Any) -> None: + raise interrupt() + + monkeypatch.setattr(runner, "run_agent_loop", _interrupt) + coordinator = AgentCoordinator() + + with pytest.raises(interrupt): + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-test", + image="img", + coordinator=coordinator, + ) + + assert _root_status(coordinator) == "running" + + +@pytest.mark.asyncio +async def test_a_real_crash_still_marks_root_failed( + monkeypatch: pytest.MonkeyPatch, tmp_path: Any +) -> None: + _wire_runner(monkeypatch, tmp_path) + + async def _boom(*_args: Any, **_kwargs: Any) -> None: + raise RuntimeError("boom") + + monkeypatch.setattr(runner, "run_agent_loop", _boom) + coordinator = AgentCoordinator() + + with pytest.raises(RuntimeError, match="boom"): + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-test", + image="img", + coordinator=coordinator, + ) + + assert _root_status(coordinator) == "failed" diff --git a/tests/test_runner_teardown.py b/tests/test_runner_teardown.py new file mode 100644 index 00000000..4433c17d --- /dev/null +++ b/tests/test_runner_teardown.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import asyncio +import types +from typing import Any + +import pytest +from agents import ModelSettings + +import strix.tools.notes.tools as notes_tools +import strix.tools.todo.tools as todo_tools +from strix.core import runner +from strix.core.agents import AgentCoordinator +from strix.runtime import session_manager + + +def _wire_runner(monkeypatch: pytest.MonkeyPatch, tmp_path: Any) -> None: + monkeypatch.setattr(runner, "run_dir_for", lambda _scan_id: tmp_path) + monkeypatch.setattr(runner, "runtime_state_dir", lambda _run_dir: tmp_path) + monkeypatch.setattr(runner, "setup_scan_logging", lambda _run_dir: lambda: None) + monkeypatch.setattr(runner, "set_scan_id", lambda _scan_id: None) + + settings = _settings() + monkeypatch.setattr(runner, "load_settings", lambda: settings) + monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _s: None) + monkeypatch.setattr(runner, "uses_chat_completions_tool_schema", lambda _m, _s: False) + monkeypatch.setattr(todo_tools, "hydrate_todos_from_disk", lambda _d: None) + monkeypatch.setattr(notes_tools, "hydrate_notes_from_disk", lambda _d: None) + + async def _create_or_reuse(*_a: Any, **_k: Any) -> dict[str, Any]: + return {"client": object(), "session": object(), "caido_client": None} + + async def _cleanup(*_a: Any, **_k: Any) -> None: + return None + + monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse) + monkeypatch.setattr(session_manager, "cleanup", _cleanup) + monkeypatch.setattr(runner, "build_root_task", lambda _c: "task") + monkeypatch.setattr(runner, "build_scope_context", lambda _c: "") + monkeypatch.setattr(runner, "make_model_settings", lambda *_a, **_k: ModelSettings()) + monkeypatch.setattr(runner, "build_strix_agent", lambda **_k: object()) + monkeypatch.setattr(runner, "make_child_factory", lambda **_k: lambda **_kk: object()) + monkeypatch.setattr(runner, "open_agent_session", lambda _root_id, _db: object()) + + +def _settings() -> Any: + return types.SimpleNamespace( + llm=types.SimpleNamespace( + model="openai/gpt-4o", + reasoning_effort="high", + force_required_tool_choice=False, + timeout=300, + prompt_cache=True, + extra_headers=None, + ), + runtime=types.SimpleNamespace(max_context_images=3), + ) + + +@pytest.mark.asyncio +async def test_a_live_child_is_settled_before_sessions_close( + monkeypatch: pytest.MonkeyPatch, tmp_path: Any +) -> None: + _wire_runner(monkeypatch, tmp_path) + coordinator = AgentCoordinator() + child_started = asyncio.Event() + child_task: dict[str, asyncio.Task[None]] = {} + + async def _root_finishes(**kwargs: Any) -> None: + root_id = kwargs["agent_id"] + + async def _child_mid_turn() -> None: + child_started.set() + await asyncio.sleep(3600) + + await coordinator.register("child", "Child", parent_id=root_id) + task = asyncio.create_task(_child_mid_turn()) + child_task["t"] = task + await coordinator.attach_runtime("child", task=task) + await child_started.wait() + + monkeypatch.setattr(runner, "run_agent_loop", _root_finishes) + + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-test", + image="img", + coordinator=coordinator, + ) + + task = child_task["t"] + assert task.done(), "the child task was left running past scan teardown" + assert task.cancelled(), "the child was not cancelled cleanly on a finish" diff --git a/tests/test_session_fd.py b/tests/test_session_fd.py new file mode 100644 index 00000000..976758be --- /dev/null +++ b/tests/test_session_fd.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import asyncio +from pathlib import Path +from typing import Any, cast + +import pytest + +from strix.core.sessions import open_agent_session + + +def _count_open_fds() -> int | None: + for path in (Path("/proc/self/fd"), Path("/dev/fd")): + if path.is_dir(): + return len(list(path.iterdir())) + return None + + +@pytest.mark.asyncio +async def test_sessions_hold_no_descriptors_while_parked(tmp_path: Path) -> None: + """Descriptor use must track live operations, not the number of sessions. + + The SDK keeps a connection per (session, pool thread) open for the session's + whole life. An agent parks rather than exits, so its session lives for the + scan, and fan-out multiplies those handles until the process runs out of file + descriptors (#1018). A session that is not mid-operation should hold none. + """ + baseline = _count_open_fds() + if baseline is None: + pytest.skip("no /proc/self/fd or /dev/fd on this platform") + + sessions = [open_agent_session(f"a{i}", tmp_path / f"s{i}.db") for i in range(60)] + try: + for _ in range(4): + await asyncio.gather( + *(s.add_items([{"role": "user", "content": "x"}]) for s in sessions) + ) + await asyncio.gather(*(s.get_items() for s in sessions)) + parked = _count_open_fds() + assert parked is not None + # 60 parked sessions, yet descriptors are back at the baseline. + assert parked - baseline <= 5, f"parked fds grew by {parked - baseline}" + finally: + for s in sessions: + s.close() + + +@pytest.mark.asyncio +async def test_in_flight_descriptors_track_concurrency_not_session_count( + tmp_path: Path, +) -> None: + baseline = _count_open_fds() + if baseline is None: + pytest.skip("no /proc/self/fd or /dev/fd on this platform") + + sessions = [open_agent_session(f"a{i}", tmp_path / f"s{i}.db") for i in range(200)] + peak = baseline + try: + + async def sample() -> None: + nonlocal peak + for _ in range(500): + current = _count_open_fds() + if current is not None: + peak = max(peak, current) + await asyncio.sleep(0) + + async def load() -> None: + for _ in range(4): + await asyncio.gather( + *(s.add_items([{"role": "user", "content": "x"}]) for s in sessions) + ) + + await asyncio.gather(load(), sample()) + # 200 sessions, but peak is bounded by the thread pool, well under 200. + assert peak - baseline < 100, f"in-flight fds peaked at +{peak - baseline}" + finally: + for s in sessions: + s.close() + + +@pytest.mark.asyncio +async def test_history_survives_the_per_operation_connection(tmp_path: Path) -> None: + session = open_agent_session("agent-1", tmp_path / "agents.db") + try: + for i in range(30): + await session.add_items([{"role": "user", "content": f"m{i}"}]) + items = [cast("dict[str, Any]", i) for i in await session.get_items()] + assert [i["content"] for i in items] == [f"m{i}" for i in range(30)] + finally: + session.close() + + +@pytest.mark.asyncio +async def test_concurrent_sessions_sharing_one_file_stay_consistent(tmp_path: Path) -> None: + db = tmp_path / "shared.db" + sessions = [open_agent_session(f"a{i}", db) for i in range(10)] + try: + await asyncio.gather( + *(s.add_items([{"role": "user", "content": s.session_id}]) for s in sessions) + ) + # Each session sees only its own row despite sharing the file. + for s in sessions: + items = [cast("dict[str, Any]", i) for i in await s.get_items()] + assert [i["content"] for i in items] == [s.session_id] + finally: + for s in sessions: + s.close() + + +@pytest.mark.asyncio +async def test_a_closed_session_refuses_operations(tmp_path: Path) -> None: + session = open_agent_session("agent-1", tmp_path / "agents.db") + await session.add_items([{"role": "user", "content": "x"}]) + session.close() + with pytest.raises(RuntimeError, match="closed"): + await session.add_items([{"role": "user", "content": "y"}]) diff --git a/tests/test_todo.py b/tests/test_todo.py new file mode 100644 index 00000000..036bcacc --- /dev/null +++ b/tests/test_todo.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +import json +from typing import Any + +import pytest +from agents.tool_context import ToolContext + +from strix.tools.todo import tools +from strix.tools.todo.tools import _coerce_priority, create_todo + + +@pytest.fixture(autouse=True) +def _isolate_store() -> Any: + tools._todos_storage.clear() + yield + tools._todos_storage.clear() + + +async def _create(todos: list[Any], agent_id: str = "root") -> dict[str, Any]: + ctx = ToolContext( + context={"agent_id": agent_id}, + tool_name="create_todo", + tool_call_id="call-1", + tool_arguments="{}", + ) + raw = await create_todo.on_invoke_tool(ctx, json.dumps({"todos": json.dumps(todos)})) + return json.loads(raw) # type: ignore[no-any-return] + + +def test_unknown_priority_falls_back_to_normal() -> None: + assert _coerce_priority("medium") == "normal" + assert _coerce_priority("urgent") == "normal" + assert _coerce_priority("high") == "high" + + +@pytest.mark.asyncio +async def test_one_bad_priority_no_longer_discards_the_batch() -> None: + result = await _create( + [ + {"title": "Recon", "priority": "medium"}, + {"title": "Probe /admin", "priority": "sky-high"}, + {"title": "Report"}, + ] + ) + + assert result["success"] is True + assert result["created_count"] == 3 + by_title = {c["title"]: c["priority"] for c in result["created"]} + assert by_title["Recon"] == "normal" + assert by_title["Probe /admin"] == "normal" + assert by_title["Report"] == "normal" + + +@pytest.mark.asyncio +async def test_duplicate_titles_within_a_batch_are_skipped() -> None: + result = await _create( + [ + {"title": "Subdomain enumeration"}, + {"title": "Content discovery"}, + {"title": "Subdomain enumeration"}, + {"title": "content discovery"}, + ] + ) + + assert result["created_count"] == 2 + assert {c["title"] for c in result["created"]} == { + "Subdomain enumeration", + "Content discovery", + } + assert len(result["skipped"]) == 2 + assert all(s["reason"] == "duplicate title" for s in result["skipped"]) + + +@pytest.mark.asyncio +async def test_a_title_already_on_the_list_is_not_created_again() -> None: + await _create([{"title": "Crawl with katana"}]) + result = await _create([{"title": "crawl with katana"}, {"title": "JS analysis"}]) + + assert [c["title"] for c in result["created"]] == ["JS analysis"] + assert [s["title"] for s in result["skipped"]] == ["crawl with katana"] + assert result["total_count"] == 2 + + +def test_coerce_never_raises() -> None: + assert _coerce_priority("nonsense") == "normal" + assert _coerce_priority(None) == "normal" + assert _coerce_priority("high") == "high" + for value in (2, ["high"], {"p": 1}, True): + assert _coerce_priority(value) == "normal" # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_non_string_priority_does_not_fail_the_batch() -> None: + result = await _create( + [ + {"title": "Recon", "priority": 2}, + {"title": "Probe", "priority": ["high"]}, + {"title": "Report"}, + ] + ) + + assert result["success"] is True + assert result["created_count"] == 3 + assert {c["priority"] for c in result["created"]} == {"normal"} diff --git a/tests/test_tui_backend_controller.py b/tests/test_tui_backend_controller.py index f6cebe13..0a4ff15c 100644 --- a/tests/test_tui_backend_controller.py +++ b/tests/test_tui_backend_controller.py @@ -234,12 +234,11 @@ async def test_confirming_the_mount_starts_the_scan_without_a_target() -> None: @pytest.mark.asyncio -async def test_declining_the_mount_returns_to_the_start_screen() -> None: - started = False +async def test_declining_the_mount_runs_without_one() -> None: + started: list[bool] = [] - async def start(_verify: bool = True) -> None: - nonlocal started - started = True + async def start(verify: bool = True) -> None: + started.append(verify) os.environ["STRIX_LLM"] = "anthropic/claude-sonnet-4" os.environ["ANTHROPIC_API_KEY"] = "test-key" @@ -250,14 +249,34 @@ async def test_declining_the_mount_returns_to_the_start_screen() -> None: result = await controller.handle("setup.confirm_mount", {"approved": False}) assert result == {"approved": False} - # Nothing was prepared, so the session goes back to the start screen and can - # be launched again. - assert started is False + # Declining skips the directory; it does not abandon the scan. + assert started == [False] assert controller.workspace_mount is None assert controller.pending_workspace_mount is None - assert controller.setup_mode is True - assert controller.scan_started is False - assert controller.scan_state == "setup" + assert controller.setup_mode is False + assert controller.scan_started is True + assert controller.scan_state == "running" + + +@pytest.mark.asyncio +async def test_approving_the_mount_runs_with_it() -> None: + started: list[bool] = [] + + async def start(verify: bool = True) -> None: + started.append(verify) + + os.environ["STRIX_LLM"] = "anthropic/claude-sonnet-4" + os.environ["ANTHROPIC_API_KEY"] = "test-key" + loader._cached = None + controller = TuiController(args(), on_start=start) + await controller.handle("setup.start", {"verify": False, "mount_working_dir": True}) + + result = await controller.handle("setup.confirm_mount", {"approved": True}) + + assert result == {"approved": True} + assert started == [False] + assert controller.workspace_mount == str(Path.cwd()) + assert controller.scan_state == "running" @pytest.mark.asyncio diff --git a/tests/test_tui_resume_history.py b/tests/test_tui_resume_history.py index f1edd8e8..20c010e5 100644 --- a/tests/test_tui_resume_history.py +++ b/tests/test_tui_resume_history.py @@ -7,19 +7,26 @@ shows what the user actually typed; resuming has to match that. from __future__ import annotations +import ast import json import sqlite3 +from pathlib import Path from typing import TYPE_CHECKING, Any import pytest +from strix.core import execution from strix.core.paths import runtime_state_dir from strix.interface.tui.backend.live_view import TuiLiveView as GoTuiLiveView -from strix.interface.tui.live_view import TuiLiveView, _is_internal_agent_turn +from strix.interface.tui.live_view import ( + _INTERNAL_TURN_PREFIXES, + TuiLiveView, + _is_internal_agent_turn, +) if TYPE_CHECKING: - from pathlib import Path + from types import ModuleType def _write_run(run_dir: Path, items: list[dict[str, Any]], agent_id: str = "root") -> None: @@ -176,11 +183,60 @@ def test_internal_turn_classifier_matches_every_injected_form() -> None: "[CRITICAL] Turn budget: 480/500 used (96%).", "== Inherited context from parent (background only) ==", "Your previous message ended a turn without a tool call.", - "Your previous response ended the autonomous Strix run without a lifecycle tool call.", + "Your previous response ended the autonomous run without a lifecycle tool call.", ): assert _is_internal_agent_turn(content), content +def _injected_strings(module: ModuleType) -> list[str]: + """Every string a module can inject, and nothing it merely mentions. + + Parsing rather than searching the text keeps comments out of it, so a stale + copy of a message left in a comment cannot pass for the message itself. It + also joins adjacent literals for free, which the line wrapping needs, and + docstrings are dropped because they describe the code rather than run in it. + """ + tree = ast.parse(Path(module.__file__ or "").read_text(encoding="utf-8")) + docstrings = set() + for node in ast.walk(tree): + if not isinstance(node, ast.Module | ast.ClassDef | ast.FunctionDef | ast.AsyncFunctionDef): + continue + first = node.body[0] if node.body else None + if isinstance(first, ast.Expr) and isinstance(first.value, ast.Constant): + docstrings.add(id(first.value)) + + literals: list[str] = [] + for node in ast.walk(tree): + if isinstance(node, ast.Constant): + if isinstance(node.value, str) and id(node) not in docstrings: + literals.append(node.value) + elif isinstance(node, ast.JoinedStr): + literals.append( + "".join( + part.value + for part in node.values + if isinstance(part, ast.Constant) and isinstance(part.value, str) + ) + ) + return literals + + +def test_internal_turn_prefixes_still_match_what_is_injected() -> None: + """The classifier copies sentences out of another module, so they can drift. + + Both nudges are written inline in strix.core.execution, so there is nothing to + import and compare against. Read them back out of what that module can inject. + """ + injected = _injected_strings(execution) + nudges = [prefix for prefix in _INTERNAL_TURN_PREFIXES if prefix.startswith("Your previous")] + assert nudges, "the no-tool-call nudges are no longer in the classifier" + for nudge in nudges: + assert any(nudge in literal for literal in injected), ( + f"the classifier expects {nudge!r}, which strix.core.execution no longer " + f"injects. A resumed scan would show that nudge as the user's own message." + ) + + def test_internal_turn_classifier_keeps_bracketed_user_text() -> None: """A leading bracket is not enough: typed text often starts with one.""" for content in ( diff --git a/uv.lock b/uv.lock index a4561526..523e7c64 100644 --- a/uv.lock +++ b/uv.lock @@ -2378,7 +2378,7 @@ wheels = [ [[package]] name = "strix-agent" -version = "1.5.1" +version = "1.5.3" source = { editable = "." } dependencies = [ { name = "caido-sdk-client" },