Compare commits

..
Author SHA1 Message Date
Ahmed Allam df95d3ad2b build(container): pin gitleaks instead of resolving the latest release 2026-08-02 15:56:20 +00:00
Ahmed Allam d582e142f4 ci: publish the sandbox image with buildx for amd64 and arm64 2026-08-02 14:49:22 +00:00
dbc427d816 feat(runtime): mount local targets instead of copying them in (#958)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-02 07:45:10 -07:00
Ahmed AllamandAhmed Allam b6cf156e95 fix(tools): tell a waiting parent when stop_agent stops its child 2026-08-02 15:43:43 +03:00
Ahmed AllamandAhmed Allam 002712284a fix(core): wake the parent when a child ends without a completion report 2026-08-02 15:43:43 +03:00
797b37467e perf(cli): ~10x faster startup via lazy imports (#920)
* perf(cli): fast startup — lazy heavy imports + onedir standalone build

* perf(cli): drop legacy single-file compat from install/self-update

* perf(cli): simplify — drop constants module and extra lazy-import refactors

* refactor(update): strix --update just re-runs the install script

* perf(cli): drop packaging/install/update changes; deepen lazy imports instead

Reverts the onedir build, install.sh, and self-update changes so release
mechanics stay untouched. Startup cost is addressed purely by deferring
heavy imports (agents/openai, config.models, report state/writer, docker)
until a scan actually runs; DEFAULT_MAX_TURNS moves to strix.config.settings
so argparse no longer pulls the agents SDK.

---------

Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-02 06:24:31 +03:00
c240068c2c fix(tools): accept both the string and structured form of every tool argument (#957)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-01 18:53:24 -07:00
2e7040240d feat(config): accept STRIX_REASONING_EFFORT=max for providers that support it (#956)
Co-authored-by: Ahmed Allam <allam@usestrix.com>
2026-08-01 17:10:01 -07:00
22d668d538 docs(llm-providers): explain the structured tool_calls requirement for local endpoints (#520) (#901)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-01 16:29:18 -07:00
Ahmed AllamandAhmed Allam f77805e5bc docs(prompts): text-only turns no longer end an autonomous run 2026-08-02 02:15:51 +03:00
Ahmed AllamandAhmed Allam 1c1fa49961 refactor(tools): split wait_for_message into respond_to_user + wait_for_agents
One tool was doing three jobs (wait on the user, wait on other agents, and
- wrongly - wait for a long-running command), so the driver had to guess which
one an agent meant and used parent_id as the proxy: the root waits for a human,
everyone else waits for agents. That proxy is wrong, since the user can message
any agent from the TUI's agent tree.

Tool identity now carries the intent, and the coordinator records it as a
wait_kind that survives snapshot/restore:

  respond_to_user  -> wait_kind="user",   never auto-resumed (root or not)
  wait_for_agents  -> wait_kind="agents", auto-resumed on a 300s timer
  recovery exhaust -> wait_kind="stalled"

respond_to_user fuses the message and the yield into one call, so there is no
way to answer and then forget to stop - the two-step that gpt-4o-mini skipped
2/2 in live testing. Plain text still renders as before.

Auto-resume is also bounded now: an agent that re-parks after every timeout
burned a model turn every 300s for the rest of the scan (and, since parked
children notify their parent, spammed the parent's inbox on the same cycle).
After _MAX_IDLE_AUTO_RESUMES it stays parked until a real message arrives.
2026-08-02 02:15:51 +03:00
Ahmed AllamandAhmed Allam 742f382836 docs(core): correct the rationale for notifying a stalled child's parent
The user can message any agent from the TUI, not only the root, so the
justification is that the parent is an agent with no other way to learn
the child parked - not that the child has no human resumer.
2026-08-02 02:15:51 +03:00
Ahmed AllamandAhmed Allam 8f1bb64d16 fix(core): tell the parent when an interactive subagent parks
Parking is self-service only for the root, which the user is watching.
A parked child owes its parent a report it can no longer send, so the
parent would wait out its full timeout for nothing.
2026-08-02 02:15:51 +03:00
Ahmed AllamandAhmed Allam 49057f267f fix(tools): halve the wait_for_message ceiling to 300s
A mutual wait between two agents resolves only when both hit their cap,
so the ceiling is the worst-case idle burn. Name the constants instead of
repeating the literal, and align the interactive auto-resume timeout.
2026-08-02 02:15:51 +03:00
Ahmed AllamandAhmed Allam 6eec34df24 fix(core): persist the tool-call recovery counter across resumes
An exhausted agent parked in 'waiting' got a fresh nudge budget on every
600s auto-resume, so a wedged agent could nudge-park-nudge indefinitely.
Track the count on the coordinator, snapshot it, and reset it only on
real input or an explicit lifecycle tool.
2026-08-02 02:15:51 +03:00
Ahmed AllamandAhmed Allam f6f9469e00 fix(core): stop interactive runs stalling on a missing tool call
Interactive turns ended by plain text left the agent parked in 'waiting'
forever. Require an explicit lifecycle tool in both modes and nudge a
text-only turn back into a tool call, bounded by a recovery limit.
2026-08-02 02:15:51 +03:00
dc7cc50f80 docs(prompt): teach agents to recognize Caido proxy error pages instead of chasing them (#955)
* docs(prompt): teach agents to recognize Caido proxy error pages

* docs(prompt): tighten Caido proxy error page section

---------

Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-02 02:01:34 +03:00
devin-ai-integration[bot]andGitHub 5602bc23ca fix: pre-v1-style lifecycle resilience — mailbox delivery, uniform revival, unexitable runner, waiting timeout, broader retries, crash-safe identity (#923) 2026-08-01 11:17:08 -07:00
alex sandGitHub a9deb84260 fix(llm): surface structured provider refusals (#944)
* fix(llm): surface structured provider refusals

* fix(llm): settle refused autonomous agents
2026-07-31 10:57:42 -04:00
chunguscodesandAhmed Allam 76e97e6a59 fix(llm): avoid auth during ChatGPT lookup
LiteLLM treats provider-qualified metadata lookups as an auth path.
Use the underlying model slug so context sizing cannot block the scan
loop in a device-code poll.
2026-07-31 03:45:41 +03:00
Ahmed AllamandAhmed Allam 885b2ca5c5 test(llm): cover the full run loop against a non-streaming gateway; drop README note
Adds an integration test that drives Runner.run_streamed against a
non-streaming gateway through _NonStreamingModel: the synthetic terminal
event feeds the runner, which executes the tool call and continues to a
final answer over two non-streaming turns. Removes the README env-var note.
2026-07-30 08:30:06 +03:00
Ahmed AllamandAhmed Allam 980216860e feat(llm): opt-in LLM_DISABLE_STREAMING for non-streaming OpenAI-compatible endpoints
Some OpenAI-compatible gateways don't support Server-Sent Events (or
deliver them unreliably), but the SDK run loop Strix uses only issues
streamed requests, so such a gateway fails every turn. Add an opt-in
LLM_DISABLE_STREAMING setting that wraps the resolved model in
_NonStreamingModel: each turn makes one non-streaming get_response and
replays the completed result as a single terminal stream event, so tool
calls, usage, and the rest of the agent loop are unchanged. Subscription
(ChatGPT) models are always streamed and are not wrapped.
2026-07-30 08:30:06 +03:00
devin-ai-integration[bot]andGitHub d4e58b2cd0 fix(llm): pass LLM_EXTRA_HEADERS through ModelSettings so they reach the agent loop (#937) 2026-07-29 19:38:06 -07:00
Ahmed AllamandAhmed Allam e9ebdc502f fix(llm): apply LLM_EXTRA_HEADERS on native OpenAI route even without a custom base 2026-07-30 04:13:25 +03:00
Ahmed AllamandAhmed Allam ebb3a62a99 feat(llm): custom request headers for OpenAI-compatible endpoints via LLM_EXTRA_HEADERS 2026-07-30 04:13:25 +03:00
1a2fa89972 fix(runtime): label docker sandbox containers with the run id for teardown (#933)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-29 08:05:42 -07:00
alex sandGitHub 9de747d135 fix(cost): capture OpenRouter streamed usage.cost (fixes $0 kimi-k3 c… (#929)
* fix(cost): capture OpenRouter streamed usage.cost (fixes $0 kimi-k3 cost)

* refactor(cost): encapsulate streamed OpenRouter cost cache, clear per run

* test(cost): resolve OpenRouter handler via LiteLLM provider pipeline
2026-07-28 23:28:34 -04:00
b313d78f60 Scope viewer session cookie to the bound port (#922)
Co-authored-by: Jonathan Singer <jonathansinger@Mac-4078.lan>
2026-07-27 20:37:54 -04:00
e037d8d727 fix: recoverable guardrail blocks and decoupled crash-notify (#919)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-27 16:28:48 -07:00
alex sandGitHub fade37025d fix viewer tool call collisions across agents (#917) 2026-07-27 18:51:02 -04:00
Ahmed AllamandAhmed Allam f968f8e5a7 fix(cli): align View label spacing in final panel 2026-07-27 15:41:22 -07:00
Ahmed AllamandAhmed Allam ac0014fe65 chore: release v1.4.1 2026-07-27 12:57:39 -07:00
86282e83a8 fix(tls): replace raw urllib with requests for external HTTPS calls (frozen-build cert failures) (#903)
Co-authored-by: Jonathan Singer <jonathansinger@Mac-4051.lan>
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-27 12:34:26 -07:00
72 changed files with 4052 additions and 1160 deletions
+112
View File
@@ -0,0 +1,112 @@
name: Sandbox Image
on:
workflow_dispatch:
inputs:
tag:
description: 'Image tag to publish (e.g. 1.2.0)'
required: true
latest:
description: 'Also tag as latest'
type: boolean
default: true
permissions:
contents: read
env:
IMAGE: ghcr.io/${{ github.repository_owner }}/strix-sandbox
jobs:
build:
strategy:
fail-fast: false
matrix:
include:
- os: ubuntu-22.04
platform: linux/amd64
- os: ubuntu-22.04-arm
platform: linux/arm64
runs-on: ${{ matrix.os }}
permissions:
contents: read
packages: write
steps:
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
with:
persist-credentials: false
- uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0
- uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
with:
registry: ghcr.io
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Build and push by digest
id: build
uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0
with:
context: .
file: containers/Dockerfile
platforms: ${{ matrix.platform }}
provenance: mode=max
sbom: true
outputs: type=image,name=${{ env.IMAGE }},push-by-digest=true,name-canonical=true,push=true
- name: Export digest
env:
DIGEST: ${{ steps.build.outputs.digest }}
run: |
set -euo pipefail
mkdir -p /tmp/digests
touch "/tmp/digests/${DIGEST#sha256:}"
- uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1
with:
name: digest-${{ runner.arch }}
path: /tmp/digests/*
if-no-files-found: error
retention-days: 1
publish:
needs: build
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- uses: actions/download-artifact@3e5f45b2cfb9172054b4087a40e8e0b5a5461e7c # v8.0.1
with:
path: /tmp/digests
pattern: digest-*
merge-multiple: true
- uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0
- uses: docker/login-action@dbcb813823bdd20940b903addbd779551569679f # v4.6.0
with:
registry: ghcr.io
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Create manifest list
env:
TAG: ${{ inputs.tag }}
ALSO_LATEST: ${{ inputs.latest }}
run: |
set -euo pipefail
tags=(-t "${IMAGE}:${TAG}")
if [ "${ALSO_LATEST}" = "true" ]; then
tags+=(-t "${IMAGE}:latest")
fi
digests=()
for file in /tmp/digests/*; do
digests+=("${IMAGE}@sha256:$(basename "$file")")
done
docker buildx imagetools create "${tags[@]}" "${digests[@]}"
docker buildx imagetools inspect "${IMAGE}:${TAG}"
+2 -2
View File
@@ -162,6 +162,7 @@ USER root
ARG TRUFFLEHOG_VERSION=3.95.9
RUN curl -sSfL https://raw.githubusercontent.com/trufflesecurity/trufflehog/main/scripts/install.sh | sh -s -- -b /home/pentester/.local/bin "v${TRUFFLEHOG_VERSION}" && \
chown -R pentester:pentester /home/pentester/.local
ARG GITLEAKS_VERSION=8.30.1
RUN set -eux; \
ARCH="$(uname -m)"; \
case "$ARCH" in \
@@ -169,8 +170,7 @@ RUN set -eux; \
aarch64|arm64) GITLEAKS_ARCH="arm64" ;; \
*) echo "Unsupported architecture: $ARCH" >&2; exit 1 ;; \
esac; \
TAG="$(curl -fsSL https://api.github.com/repos/gitleaks/gitleaks/releases/latest | jq -r .tag_name)"; \
curl -fsSL "https://github.com/gitleaks/gitleaks/releases/download/${TAG}/gitleaks_${TAG#v}_linux_${GITLEAKS_ARCH}.tar.gz" -o /tmp/gitleaks.tgz; \
curl -fsSL "https://github.com/gitleaks/gitleaks/releases/download/v${GITLEAKS_VERSION}/gitleaks_${GITLEAKS_VERSION}_linux_${GITLEAKS_ARCH}.tar.gz" -o /tmp/gitleaks.tgz; \
tar -xzf /tmp/gitleaks.tgz -C /tmp; \
install -m 0755 /tmp/gitleaks /usr/local/bin/gitleaks; \
rm -f /tmp/gitleaks /tmp/gitleaks.tgz
+16
View File
@@ -1,6 +1,22 @@
#!/bin/bash
set -e
if [ -n "${STRIX_HOST_UID:-}" ] && [ "${STRIX_HOST_UID}" != "0" ] && [ "${STRIX_HOST_UID}" != "$(id -u)" ]; then
exec sudo -E -- bash -c '
set -e
gid="${STRIX_HOST_GID:-$STRIX_HOST_UID}"
old_uid="$1"
old_gid="$2"
export PATH="$3"
shift 3
sed -i "s|^pentester:x:${old_uid}:${old_gid}:|pentester:x:${STRIX_HOST_UID}:${gid}:|" /etc/passwd
sed -i "s|^pentester:x:${old_gid}:|pentester:x:${gid}:|" /etc/group
chown -R "${STRIX_HOST_UID}:${gid}" /home/pentester /app/certs
chown "${STRIX_HOST_UID}:${gid}" /workspace
exec setpriv --reuid "${STRIX_HOST_UID}" --regid "${gid}" --init-groups "$0" "$@"
' "$0" "$(id -u)" "$(id -g)" "$PATH" "$@"
fi
CAIDO_PORT=48080
CAIDO_LOG="/tmp/caido_startup.log"
+16 -6
View File
@@ -19,6 +19,14 @@ Configure Strix using environment variables or a config file.
Custom API base URL. Also accepts `OPENAI_API_BASE`, `LITELLM_BASE_URL`, or `OLLAMA_API_BASE`.
</ParamField>
<ParamField path="LLM_EXTRA_HEADERS" type="string">
Extra HTTP headers sent on every LLM request, as a JSON object (e.g.
`{"X-Feature-Key":"value","X-Tenant":"acme"}`). Useful for OpenAI-compatible
gateways that require attribution or routing headers in addition to the bearer
token. The bearer token itself still comes from `LLM_API_KEY`. Applies to both
the LiteLLM and native OpenAI routing paths.
</ParamField>
<ParamField path="LLM_TIMEOUT" default="300" type="integer">
Request timeout in seconds for LLM calls.
</ParamField>
@@ -28,7 +36,7 @@ Configure Strix using environment variables or a config file.
</ParamField>
<ParamField path="STRIX_REASONING_EFFORT" default="high" type="string">
Control thinking effort for reasoning models. Valid values: `none`, `minimal`, `low`, `medium`, `high`, `xhigh`. Defaults to `medium` for quick scan mode.
Control thinking effort for reasoning models. Valid values: `none`, `minimal`, `low`, `medium`, `high`, `xhigh`, `max`. Defaults to `medium` for quick scan mode.
</ParamField>
<ParamField path="STRIX_MEMORY_COMPRESSOR_TIMEOUT" default="30" type="integer">
@@ -55,6 +63,12 @@ affecting the agents that do the actual testing.
model runs on a different endpoint than the main model.
</ParamField>
<ParamField path="DEDUPE_LLM_EXTRA_HEADERS" type="string">
Optional JSON object of extra HTTP headers sent on every deduplication-model
request, e.g. `{"X-Feature-Key":"value"}`. A dedicated dedupe model never
inherits `LLM_EXTRA_HEADERS`; set this when its endpoint needs custom headers.
</ParamField>
<ParamField path="STRIX_DEDUPE_REASONING_EFFORT" type="string">
Reasoning effort for the deduplication model. Defaults to the model's own
baseline when unset.
@@ -92,7 +106,7 @@ When remote vars are set, Strix dual-writes telemetry to both local JSONL and th
## Docker Configuration
<ParamField path="STRIX_IMAGE" default="ghcr.io/usestrix/strix-sandbox:1.0.0" type="string">
<ParamField path="STRIX_IMAGE" default="ghcr.io/usestrix/strix-sandbox:1.2.0" type="string">
Docker image to use for the sandbox container.
</ParamField>
@@ -104,10 +118,6 @@ When remote vars are set, Strix dual-writes telemetry to both local JSONL and th
Runtime backend for the sandbox environment.
</ParamField>
<ParamField path="STRIX_MAX_LOCAL_COPY_MB" default="1024" type="integer">
Maximum size (in MB) of a local directory target that Strix will copy into the sandbox file-by-file. Larger targets exit early with a suggestion to use `--mount` instead. Set to `0` to disable the check.
</ParamField>
## Sandbox Configuration
<ParamField path="STRIX_SANDBOX_EXECUTION_TIMEOUT" default="120" type="integer">
+52
View File
@@ -54,3 +54,55 @@ If you use LM Studio, vLLM, or other runners:
export STRIX_LLM="openai/local-model"
export LLM_API_BASE="http://localhost:1234/v1" # Adjust port as needed
```
### Gateways that require custom headers
Some OpenAI-compatible gateways require extra HTTP headers (for attribution or
tenant routing) alongside the bearer token. Set them with `LLM_EXTRA_HEADERS` as
a JSON object — they are sent on every request:
```bash
export STRIX_LLM="openai/your-model"
export LLM_API_BASE="https://your-gateway.example/v1"
export LLM_API_KEY="your-bearer-token" # sent as Authorization: Bearer ...
export LLM_EXTRA_HEADERS='{"X-Feature-Key":"value","X-Tenant":"acme"}'
```
For endpoints behind a private CA, point Strix at your certificate bundle with
the standard `SSL_CERT_FILE=/path/to/ca-bundle.pem` — never disable TLS
verification against a real endpoint.
## Tool calling must return structured `tool_calls`
Strix is entirely tool-driven: every working turn must be a **native** function/tool call. If your inference server returns the tool call as plain assistant text instead of a structured `tool_calls` field, Strix never sees a call it can execute, so the agent makes no real progress — it re-prompts the model for a tool call and gives up once its recovery attempts are exhausted.
This is almost always an **inference-server configuration** problem, not a model or Strix problem. Common symptoms are the model printing a call as text such as:
```text
<tool_call>{"name": "exec_command", "arguments": {"cmd": "nmap ..."}}</tool_call>
exec_command(cmd="nmap ...", timeout=180)
{"action": "exec_command", "params": {"cmd": "nmap ..."}}
```
The fix belongs on the inference server: it must be configured to parse the model's tool tokens into structured `tool_calls`. A correctly configured endpoint either returns a structured call or rejects the request outright — it never leaks the call as text.
### Fixes by server
**llama.cpp (`llama-server`)**
- Run with `--jinja` and a correct tool-use chat template (`--chat-template` / `--chat-template-file` matching the model). Recent builds enable `--jinja` by default — **upgrade** if yours doesn't.
- For thinking models, align or disable reasoning (`--reasoning-format`, `-rea off`) so it doesn't break tool-call parsing.
- A low temperature (e.g. `--temp 0.2`) improves tool-call reliability.
**Ollama**
- Use a recent Ollama and a model whose template wires tools. Modern Ollama refuses tools (`tools param requires --jinja flag`) if the template lacks tool support.
- For reasoning models (e.g. qwen3), disable the model's **thinking** mode — thinking left on frequently pushes the tool call into the text `content` instead of the structured `tool_calls` field. Turn it off on the Ollama side (a non-thinking model variant, or `think: false` in the model's parameters / `Modelfile`).
- Raise **`num_ctx`** to at least 16k32k. Strix sends a large system prompt plus many tool schemas; at Ollama's small default context the tool definitions are truncated out of the prompt and the model stops emitting valid calls. A short test prompt can look fine while a real scan fails, so set this explicitly rather than inferring it from a quick check.
**vLLM**
- Start with `--enable-auto-tool-choice`, a matching `--tool-call-parser` (`hermes`, `qwen3_xml`, or `llama3_json`), and a matching `--reasoning-parser` for reasoning models.
A low sampling temperature (roughly 0.20.6, depending on the family) also measurably reduces malformed tool calls on open-weight models. Set it on the server or in your model's parameters.
<Warning>
Even correctly configured, small models (< ~30B) emit malformed or text-form tool calls far more often than frontier models. Prefer a capable model for reliable agentic behavior.
</Warning>
+6 -19
View File
@@ -6,33 +6,23 @@ description: "Command-line options for Strix"
## Basic Usage
```bash
strix (--target <target> | --target-list <path> | --mount <path>) [options]
strix (--target <target> | --target-list <path>) [options]
```
## Options
<ParamField path="--target, -t" type="string">
Target to test. Accepts URLs, repositories, local directories, domains, or IP addresses. Can be specified multiple times. Fresh runs require at least one target source: `--target`, `--target-list`, or `--mount`.
Target to test. Accepts URLs, repositories, local directories, domains, or IP addresses. Can be specified multiple times. Fresh runs require at least one target source: `--target` or `--target-list`.
<Note>
A local directory is mounted into the sandbox live and **writable**, so the agent edits your real files (`.git` excepted). Commit or stash first.
</Note>
</ParamField>
<ParamField path="--target-list" type="string">
Path to a file containing targets, one per non-empty, non-comment line. Lines starting with `#` are ignored. Can be specified multiple times and combined with `--target`.
</ParamField>
<ParamField path="--mount" type="string">
Bind-mount a local directory into the sandbox (read-only) instead of copying it in file-by-file. Use this for large repositories that are too big to stream into the container. Can be specified multiple times.
Strix copies local `--target` directories into the sandbox one file at a time, which stalls on very large trees. When a local target exceeds the copy limit (see `STRIX_MAX_LOCAL_COPY_MB`, default 1024 MB) Strix exits early and asks you to re-run with `--mount`.
<Note>
The mount is read-only to protect your source from accidental modification. This is not a hard security boundary: a root process inside the container can remount it writable, so treat `--mount` as "scan my own code", not as isolation from untrusted code.
</Note>
<Note>
The size pre-flight only covers local directory targets. Remote repositories (cloned at scan time) are not size-checked.
</Note>
</ParamField>
<ParamField path="--instruction" type="string">
Custom instructions for the scan. Use for credentials, focus areas, or specific testing approaches.
</ParamField>
@@ -140,9 +130,6 @@ strix -t https://github.com/org/app -t https://staging.example.com
# Targets from a file
strix --target-list ./targets.txt
# Large local repository — bind-mount instead of copying it in
strix --mount ./huge-monorepo
```
## Exit Codes
+2 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "strix-agent"
version = "1.4.0"
version = "1.4.1"
description = "Open-source AI Hackers for your apps"
readme = "README.md"
license = "Apache-2.0"
@@ -220,6 +220,7 @@ ignore = [
# Stdlib HTTP handler overrides (do_GET/do_POST).
"strix/interface/auth_cli.py" = ["N802"]
"tests/test_codex_streaming.py" = ["N802"]
"tests/test_disable_streaming.py" = ["N802"]
"tests/test_report_pdf.py" = ["S105", "S106"]
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
# circular dependency with strix.telemetry / strix.interface.viewer.report_pdf.
+1 -1
View File
@@ -4,7 +4,7 @@ set -euo pipefail
APP=strix
REPO="usestrix/strix"
STRIX_IMAGE="ghcr.io/usestrix/strix-sandbox:1.1.0"
STRIX_IMAGE="ghcr.io/usestrix/strix-sandbox:1.2.0"
MUTED='\033[0;2m'
RED='\033[0;31m'
+97 -7
View File
@@ -23,7 +23,7 @@ from strix.tools.agents_graph.tools import (
send_message_to_agent,
stop_agent,
view_agent_graph,
wait_for_message,
wait_for_agents,
)
from strix.tools.finish.tool import finish_scan
from strix.tools.load_skill.tool import load_skill
@@ -49,6 +49,7 @@ from strix.tools.reporting.tool import (
get_report,
list_reports,
)
from strix.tools.respond.tool import respond_to_user
from strix.tools.thinking.tool import think
from strix.tools.todo.tools import (
create_todo,
@@ -142,6 +143,83 @@ def _with_bounded_result(tool: FunctionTool) -> FunctionTool:
return tool
def _schema_types(spec: dict[str, Any]) -> set[str]:
types: set[str] = set()
raw = spec.get("type")
if isinstance(raw, str):
types.add(raw)
elif isinstance(raw, list):
types.update(t for t in raw if isinstance(t, str))
for variant in spec.get("anyOf") or ():
if isinstance(variant, dict):
types |= _schema_types(variant)
types.discard("null")
return types
def _decode_structured(value: str, types: set[str]) -> Any:
stripped = value.strip()
if not stripped:
return value
try:
decoded = json.loads(stripped)
except json.JSONDecodeError:
return value
wanted = list if "array" in types else dict
return decoded if isinstance(decoded, wanted) else value
def _coerce_argument(value: Any, spec: dict[str, Any]) -> Any:
types = _schema_types(spec)
if not types or value is None:
return value
if isinstance(value, list | dict) and "string" in types and not types & {"array", "object"}:
return json.dumps(value, ensure_ascii=False)
if isinstance(value, str) and types & {"array", "object"} and "string" not in types:
return _decode_structured(value, types)
return value
def _coerce_arguments(raw_input: str, schema: dict[str, Any]) -> str:
properties = schema.get("properties")
if not isinstance(properties, dict) or not properties:
return raw_input
try:
payload = json.loads(raw_input) if raw_input else None
except json.JSONDecodeError:
return raw_input
if not isinstance(payload, dict):
return raw_input
changed = False
for key, value in payload.items():
spec = properties.get(key)
if not isinstance(spec, dict):
continue
coerced = _coerce_argument(value, spec)
if coerced is not value:
payload[key] = coerced
changed = True
if not changed:
return raw_input
return json.dumps(payload, ensure_ascii=False)
def _with_coerced_arguments(tool: FunctionTool) -> FunctionTool:
if getattr(tool, "_strix_coerced", False):
return tool
invoke_tool = tool.on_invoke_tool
schema = tool.params_json_schema
async def invoke(ctx: Any, raw_input: str) -> Any:
return await invoke_tool(ctx, _coerce_arguments(raw_input, schema))
tool.on_invoke_tool = invoke
tool._strix_coerced = True # type: ignore[attr-defined]
return tool
def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
invoke_tool = tool.on_invoke_tool
@@ -211,11 +289,13 @@ def _configure_filesystem_tools(toolset: Any, *, chat_completions: bool) -> None
if isinstance(tool, CustomTool):
setattr(toolset, name, _custom_tool_as_function_tool(tool))
elif isinstance(tool, FunctionTool):
setattr(toolset, name, _function_tool_with_error_result(tool))
setattr(
toolset, name, _function_tool_with_error_result(_with_coerced_arguments(tool))
)
elif isinstance(tool, CustomTool):
setattr(toolset, name, _bound_custom_tool(tool))
elif isinstance(tool, FunctionTool):
setattr(toolset, name, _with_bounded_result(tool))
setattr(toolset, name, _with_bounded_result(_with_coerced_arguments(tool)))
def _make_filesystem_configurator(*, chat_completions: bool) -> Any:
@@ -328,7 +408,7 @@ def _configure_shell_tools(toolset: Any, *, chat_completions: bool) -> None:
for name, tool in vars(toolset).items():
if not isinstance(tool, FunctionTool):
continue
wrapped = tool
wrapped = _with_coerced_arguments(tool)
if tool.name == "exec_command":
wrapped = _wrap_exec_command(wrapped)
elif tool.name == "write_stdin":
@@ -345,6 +425,10 @@ def _make_shell_configurator(*, chat_completions: bool) -> Any:
return configure
# Tools that hand control away by parking the agent rather than ending the scan.
_PARKING_TOOLS: frozenset[str] = frozenset({"respond_to_user", "wait_for_agents"})
def _lifecycle_tool_completed(tool_name: str, output: Any) -> bool:
if tool_name == "agent_finish":
completion_key = "agent_completed"
@@ -363,7 +447,7 @@ def _lifecycle_tool_completed(tool_name: str, output: Any) -> bool:
def _wait_tool_parked(tool_name: str, output: Any) -> bool:
if tool_name != "wait_for_message" or not isinstance(output, str):
if tool_name not in _PARKING_TOOLS or not isinstance(output, str):
return False
try:
parsed = json.loads(output)
@@ -425,7 +509,7 @@ _BASE_TOOLS: tuple[Tool, ...] = (
scope_rules,
view_agent_graph,
send_message_to_agent,
wait_for_message,
wait_for_agents,
create_agent,
stop_agent,
)
@@ -509,13 +593,19 @@ def build_strix_agent(
)
agent_tools = [*_EXTRA_TOOLS, *(extra_tools or [])]
if interactive:
# Yielding to the user is only meaningful when one is attached.
agent_tools.append(respond_to_user)
if is_root:
tools: list[Tool] = [*_BASE_TOOLS, *agent_tools, finish_scan]
else:
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
_ensure_unique_tool_names(tools)
tools = [
_with_bounded_result(tool) if isinstance(tool, FunctionTool) else tool for tool in tools
_with_bounded_result(_with_coerced_arguments(tool))
if isinstance(tool, FunctionTool)
else tool
for tool in tools
]
logger.info(
+28 -19
View File
@@ -31,28 +31,26 @@ INTER-AGENT MESSAGES:
{% if interactive %}
INTERACTIVE BEHAVIOR:
- You are in an interactive conversation with a user
- CRITICAL: A message WITHOUT a tool call IMMEDIATELY STOPS your entire execution and waits for user input. This is a HARD SYSTEM CONSTRAINT, not a suggestion.
- Statements like "Planning the assessment..." or "I'll now scan..." or "Starting with..." WITHOUT a tool call will HALT YOUR WORK COMPLETELY. The system interprets no-tool-call as "I'm done, waiting for the user."
- If you want to plan, call the think tool. If you want to act, call the appropriate tool. There is NO valid reason to output text without a tool call while working on a task.
- The ONLY time you may send a message without a tool call is when you are genuinely DONE and presenting final results, or when you NEED the user to answer a question before continuing.
- EVERY message while working MUST contain exactly one tool call — this is what keeps execution moving. No tool call = execution stops.
- You may include brief explanatory text BEFORE the tool call
- Respond naturally when the user asks questions or gives instructions
- For simple conversation, acknowledgements, or direct questions that you can answer from current context, reply in plain text and stop. Do NOT call think just to prepare wording.
- If you use a tool to answer a user question (for example list_todos, view_agent_graph, or a file read), then after the tool result arrives, provide the answer in plain text and stop unless the user explicitly asked you to continue working.
- Never loop through think or other tools just to prepare, polish, confirm, or announce a final answer. Once you know the answer, say it.
- NEVER send empty messages — if you have nothing to do or say, call the wait_for_message tool
- If you catch yourself about to describe multiple steps without a tool call, STOP and call the think tool instead
- You are in an interactive conversation with a user.
- HOW EXECUTION ENDS: your turn ends ONLY when you make an explicit lifecycle tool call. Plain text NEVER ends your turn and NEVER hands control to the user — text is shown to the user, and then execution continues.
- To answer the user and hand control back, call respond_to_user. It delivers your message AND parks you for their reply in one call, so there is no way to answer and then forget to stop. This is the ONLY way to yield to the user.
- To wait on another AGENT (a child's report, a peer's reply), call wait_for_agents. That is not a way to reach the user.
- 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.
- 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.
- Never loop through think or other tools just to prepare, polish, confirm, or announce an answer. Once you know the answer, send it with respond_to_user.
{% else %}
AUTONOMOUS BEHAVIOR:
- Work autonomously by default
- You should NOT ask for user input or confirmation - you should always proceed with your task autonomously.
- Minimize user messaging: avoid redundancy and repetition; consolidate updates into a single concise message
- NEVER send an empty or blank message. If you have no content to output or need to wait (for user input, subagent results, or any other reason), you MUST call the wait_for_message tool (or another appropriate tool) instead of emitting an empty response.
- If there is nothing to execute and no user query to answer any more: do NOT send filler/repetitive text — either call wait_for_message or finish your work (subagents: agent_finish; root: finish_scan)
- While the agent loop is running, almost every output MUST be a tool call. Do NOT send plain text messages; act via tools. If idle, use wait_for_message; when done, use agent_finish (subagents) or finish_scan (root)
- A text-only turn — even one — IMMEDIATELY ends the scan/run with no report written. The lifecycle tools (``finish_scan`` for root, ``agent_finish`` for subagents) are the ONLY valid way to terminate. If you find yourself wanting to say "Done!" or "Scan complete" without a tool call, call the lifecycle tool instead — the report and termination signal both flow through it.
- NEVER send an empty or blank message. If you have no content to output or need to wait for subagent results, you MUST call the wait_for_agents tool (or another appropriate tool) instead of emitting an empty response.
- There is no user attached to this run, so there is nobody to ask and nothing to yield to. If there is nothing left to execute: do NOT send filler/repetitive text — either call wait_for_agents (only if you are genuinely expecting another agent to message you) or finish your work (subagents: agent_finish; root: finish_scan)
- While the agent loop is running, almost every output MUST be a tool call. Do NOT send plain text messages; act via tools. If waiting on another agent, use wait_for_agents; when done, use agent_finish (subagents) or finish_scan (root)
- A text-only turn does nothing: it neither ends the run nor yields — it just wastes a turn and forces a retry. The lifecycle tools (``finish_scan`` for root, ``agent_finish`` for subagents) are the ONLY way to terminate, and the report flows through them. If you find yourself wanting to say "Done!" or "Scan complete" without a tool call, call the lifecycle tool instead.
{% endif %}
</communication_rules>
@@ -446,8 +444,19 @@ PROXY & INTERCEPTION:
- Caido CLI - Modern web proxy (already running). Use the proxy tools
directly, or import `caido_api` from sandbox Python scripts.
- HTTPQL filters (for `list_requests`): quote string values, leave integers unquoted (`resp.code.eq:200`, not `"200"`); combine terms with `AND`/`OR` (there is no `NOT` — use the negated operator `ne`/`ncont`/`nregex`). Numeric fields (`resp.code`, `req.port`) use `eq`/`ne`/`gt`/`gte`/`lt`/`lte`; text fields (`req.host`, `req.path`, `req.method`, `req.raw`) use `cont`/`ncont`/`eq`/`regex`. Example: `resp.code.gte:200 AND resp.code.lt:300 AND req.host.cont:"api"`.
- NOTE: If you are seeing proxy errors when sending requests, it usually means you are not sending requests to a correct url/host/port.
- Ignore Caido proxy-generated 50x HTML error pages; these are proxy issues (might happen when requesting a wrong host or SSL/TLS issues, etc).
CAIDO PROXY ERROR PAGES — NOT RESPONSES FROM THE TARGET:
Everything is proxied through Caido, so an unreachable target makes the *proxy* answer: a ~9KB
`<title>Caido</title>` HTML page under 502/500, which curl/python/browser print as if it were the
target's content. The request never reached a server. It also appears in `list_requests` with no
response at all (`resp` null), unlike a real 502.
- Don't dump it; extract the cause with `curl -s ... | grep -A8 'c-title"'`.
- The `c-details` cause says what to fix: "Failed to query DNS" — host doesn't resolve, check
`dig +short <host>`, then correct or drop it; "Connection refused" — nothing on that port, check
`nc -z -v <host> <port>`; "TLS handshake"/"wrong version number" — scheme/port mismatch, flip
http/https; timeout — filtered or unreachable from the sandbox.
- NEVER treat these as target behavior: not a finding, not evidence, not a WAF, not a server
error. Fix the url/host/port/scheme and retry, or move on — do not keep re-requesting a dead host.
PROGRAMMING:
- Python 3, uv, Node.js/npm
+13 -20
View File
@@ -18,12 +18,12 @@ import logging
import secrets
import threading
import time
import urllib.error
import urllib.parse
import urllib.request
from pathlib import Path
from typing import TYPE_CHECKING, Any
import requests
if TYPE_CHECKING:
from collections.abc import Iterator
@@ -221,26 +221,19 @@ def _first(query: dict[str, list[str]], key: str) -> str | None:
def _post_form(payload: dict[str, str]) -> dict[str, Any]:
body = urllib.parse.urlencode(payload).encode("ascii")
request = urllib.request.Request( # noqa: S310 - fixed https OAuth endpoint
TOKEN_URL,
data=body,
headers={
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
},
method="POST",
)
try:
with urllib.request.urlopen( # noqa: S310 # nosec B310 - fixed https endpoint
request, timeout=_TOKEN_TIMEOUT
) as response:
data = json.loads(response.read() or b"{}")
except urllib.error.HTTPError as exc:
detail = exc.read().decode("utf-8", "replace")[:300]
raise CodexAuthError("token_http_error", f"HTTP {exc.code}: {detail}") from exc
except (urllib.error.URLError, TimeoutError, OSError) as exc:
response = requests.post(
TOKEN_URL,
data=payload,
headers={"Accept": "application/json"},
timeout=_TOKEN_TIMEOUT,
)
except requests.RequestException as exc:
raise CodexAuthError("unavailable", str(exc)) from exc
if response.status_code >= 400:
detail = response.text[:300]
raise CodexAuthError("token_http_error", f"HTTP {response.status_code}: {detail}")
data = json.loads(response.content or b"{}")
if not isinstance(data, dict):
raise CodexAuthError("bad_response", "token endpoint returned non-object")
return data
+231 -8
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import contextlib
import inspect
import os
import time
from typing import TYPE_CHECKING, Any
from agents import (
@@ -13,6 +14,8 @@ from agents import (
set_tracing_disabled,
)
from agents.model_settings import ModelSettings
from agents.models.fake_id import FAKE_RESPONSES_ID
from agents.models.interface import Model
from agents.models.multi_provider import MultiProvider
from agents.models.openai_responses import OpenAIResponsesModel
from agents.retry import (
@@ -21,6 +24,8 @@ from agents.retry import (
RetryPolicyContext,
retry_policies,
)
from openai.types.responses import Response, ResponseCompletedEvent
from openai.types.responses.response_usage import ResponseUsage
from openai.types.shared import Reasoning
from strix.config import codex
@@ -30,10 +35,17 @@ from strix.config.loader import load_settings
if TYPE_CHECKING:
from collections.abc import AsyncIterator
from agents.models.interface import Model, ModelProvider
from agents.agent_output import AgentOutputSchemaBase
from agents.handoffs import Handoff
from agents.items import ModelResponse, TResponseInputItem, TResponseStreamEvent
from agents.models.interface import ModelProvider, ModelTracing
from agents.retry import ModelRetryAdvice, ModelRetryAdviceRequest
from agents.tool import Tool
from agents.usage import Usage
from openai import AsyncOpenAI
from openai.types.responses.response_prompt_param import ResponsePromptParam
from strix.config.settings import ReasoningEffort, Settings
from strix.config.settings import LlmSettings, ReasoningEffort, Settings
def request_timeout_extra_args(timeout_s: float | None) -> dict[str, float] | None:
@@ -71,10 +83,13 @@ class _CodexResponsesModel(OpenAIResponsesModel):
effort = self._reasoning_effort
if effort and effort != "none":
# Clamp to efforts the backend accepts.
if effort == "minimal":
effort = "low"
elif effort == "xhigh":
effort = "high"
match effort:
case "minimal":
effort = "low"
case "xhigh" | "max":
effort = "high"
case _:
pass
overrides = overrides.resolve(ModelSettings(reasoning=Reasoning(effort=effort)))
return model_settings.resolve(overrides)
@@ -135,6 +150,124 @@ class _CodexResponsesModel(OpenAIResponsesModel):
await result
class _NonStreamingModel(Model):
"""Serve the SDK's streamed run loop from a single non-streaming request.
Some OpenAI-compatible gateways do not support Server-Sent Events, or
deliver them unreliably (dropping structured tool-call deltas, or stalling
mid-stream so the whole turn waits out the read timeout). The SDK run loop
Strix uses only issues streamed requests, so such a gateway fails every
turn. Opt in with ``LLM_DISABLE_STREAMING=true`` to wrap the resolved model
so each turn makes one non-streaming ``get_response`` (``stream:false`` on
the wire) and the completed result is replayed as a single terminal stream
event. The run loop then executes tools and emits run items from that final
response exactly as it would for a real stream, so nothing else changes.
"""
def __init__(self, inner: Model) -> None:
self._inner = inner
async def close(self) -> None:
await self._inner.close()
def get_retry_advice(self, request: ModelRetryAdviceRequest) -> ModelRetryAdvice | None:
return self._inner.get_retry_advice(request)
async def get_response(
self,
system_instructions: str | None,
input: str | list[TResponseInputItem], # noqa: A002
model_settings: ModelSettings,
tools: list[Tool],
output_schema: AgentOutputSchemaBase | None,
handoffs: list[Handoff],
tracing: ModelTracing,
*,
previous_response_id: str | None,
conversation_id: str | None,
prompt: ResponsePromptParam | None,
) -> ModelResponse:
return await self._inner.get_response(
system_instructions,
input,
model_settings,
tools,
output_schema,
handoffs,
tracing,
previous_response_id=previous_response_id,
conversation_id=conversation_id,
prompt=prompt,
)
async def stream_response(
self,
system_instructions: str | None,
input: str | list[TResponseInputItem], # noqa: A002
model_settings: ModelSettings,
tools: list[Tool],
output_schema: AgentOutputSchemaBase | None,
handoffs: list[Handoff],
tracing: ModelTracing,
*,
previous_response_id: str | None,
conversation_id: str | None,
prompt: ResponsePromptParam | None,
) -> AsyncIterator[TResponseStreamEvent]:
response = await self._inner.get_response(
system_instructions,
input,
model_settings,
tools,
output_schema,
handoffs,
tracing,
previous_response_id=previous_response_id,
conversation_id=conversation_id,
prompt=prompt,
)
yield _completed_stream_event(response, getattr(self._inner, "model", None))
def _completed_stream_event(
model_response: ModelResponse, model_name: object | None
) -> TResponseStreamEvent:
"""Wrap a non-streamed ``ModelResponse`` as the terminal event of a stream.
The run loop builds its authoritative per-turn response solely from the
``response.completed`` event, so a single event carrying the full output
and usage is all it needs.
"""
response = Response(
id=model_response.response_id or FAKE_RESPONSES_ID,
created_at=time.time(),
model=str(model_name) if model_name else "",
object="response",
output=list(model_response.output),
tool_choice="auto",
tools=[],
parallel_tool_calls=False,
usage=_response_usage(model_response.usage),
)
return ResponseCompletedEvent(
response=response,
sequence_number=0,
type="response.completed",
)
def _response_usage(usage: Usage | None) -> ResponseUsage | None:
if usage is None:
return None
return ResponseUsage(
input_tokens=usage.input_tokens,
output_tokens=usage.output_tokens,
total_tokens=usage.total_tokens,
input_tokens_details=usage.input_tokens_details,
output_tokens_details=usage.output_tokens_details,
)
class StrixProvider(MultiProvider):
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
so users type ``deepseek/deepseek-chat`` rather than
@@ -159,14 +292,21 @@ class StrixProvider(MultiProvider):
return self._get_fallback_provider("litellm"), original_model_name
def get_model(self, model_name: str | None) -> Model:
llm = load_settings().llm
slug = codex.subscription_model(model_name)
if slug:
# The ChatGPT subscription backend is always streamed; it has no
# non-streaming mode to fall back to, so LLM_DISABLE_STREAMING
# does not apply here.
return _CodexResponsesModel(
slug,
codex.get_subscription_client(),
reasoning_effort=load_settings().llm.reasoning_effort,
reasoning_effort=llm.reasoning_effort,
)
return super().get_model(model_name)
model = super().get_model(model_name)
if llm.disable_streaming:
return _NonStreamingModel(model)
return model
DEFAULT_MODEL_RETRY = ModelRetrySettings(
@@ -243,6 +383,7 @@ def configure_sdk_model_defaults(settings: Settings) -> None:
set_default_openai_api("chat_completions")
else:
set_default_openai_api("responses")
_configure_extra_headers(llm)
def _mirror_api_key_to_provider_env(model_name: str | None, api_key: str) -> None:
@@ -277,6 +418,51 @@ def _configure_litellm_compatibility() -> None:
litellm.suppress_debug_info = True
_register_litellm_cost_callback()
_install_openrouter_stream_cost_capture()
def _install_openrouter_stream_cost_capture() -> None:
"""Preserve OpenRouter's per-stream cost, which LiteLLM drops when streaming.
OpenRouter reports the real charge in ``usage.cost`` of the final stream
chunk, but LiteLLM rebuilds streamed responses from token-only fields and
discards it (its non-streamed path stashes the cost in hidden params; the
streaming path does not). Every scan streams, so without this the cost is
lost and Strix falls back to a cost-map estimate that is missing entirely
for new models (e.g. kimi-k3), reporting $0. Subclass the OpenRouter
streaming handler to record the cost keyed by response id so the cost
callback can recover the exact charge for the matching rebuilt response.
"""
import litellm
from litellm.llms.openrouter.chat.transformation import (
OpenRouterChatCompletionStreamingHandler,
OpenrouterConfig,
)
from strix.report.state import streamed_openrouter_costs
class _StrixOpenRouterStreamingHandler(OpenRouterChatCompletionStreamingHandler):
def chunk_parser(self, chunk: dict[str, Any]) -> Any:
stream = super().chunk_parser(chunk)
streamed_openrouter_costs.remember(
chunk.get("id") or getattr(stream, "id", None), chunk.get("usage")
)
return stream
class _StrixOpenrouterConfig(OpenrouterConfig):
def get_model_response_iterator(
self, streaming_response: Any, sync_stream: bool, json_mode: bool | None = False
) -> Any:
return _StrixOpenRouterStreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
# LiteLLM's provider-config factory reads litellm.OpenrouterConfig at call
# time, so overriding the attribute is enough for the subclass to take
# effect. (type: ignore — mypy rejects reassigning a class attribute.)
litellm.OpenrouterConfig = _StrixOpenrouterConfig # type: ignore[misc]
_OPENROUTER_ATTRIBUTION_HEADERS = {
@@ -302,6 +488,43 @@ def _configure_openrouter_attribution(model_name: str | None) -> None:
litellm.headers = {**existing, **_OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment]
def _configure_extra_headers(llm: LlmSettings) -> None:
"""Send user-provided default headers on every LLM request.
Some OpenAI-compatible endpoints require extra HTTP headers (e.g. request
attribution or tenant routing) alongside the bearer token. Users supply
them via ``LLM_EXTRA_HEADERS``; they are applied to both routing paths:
the LiteLLM route (``litellm.headers``) and the SDK-native OpenAI route
(a default client carrying ``default_headers``), so they take effect
regardless of the ``STRIX_LLM`` prefix.
"""
headers = llm.extra_headers
if not headers:
return
_merge_litellm_headers(headers)
_register_openai_client_with_headers(llm, headers)
def _merge_litellm_headers(headers: dict[str, str]) -> None:
import litellm
current: object = litellm.headers
existing: dict[str, str] = current if isinstance(current, dict) else {}
litellm.headers = {**existing, **headers} # type: ignore[assignment]
def _register_openai_client_with_headers(llm: LlmSettings, headers: dict[str, str]) -> None:
from agents import set_default_openai_client
from openai import AsyncOpenAI
client = AsyncOpenAI(
api_key=llm.api_key or "not-needed",
base_url=llm.api_base,
default_headers=dict(headers),
)
set_default_openai_client(client, use_for_tracing=False)
def _register_litellm_cost_callback() -> None:
import litellm
+16 -7
View File
@@ -8,7 +8,9 @@ from pydantic import AliasChoices, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
ReasoningEffort = Literal["none", "minimal", "low", "medium", "high", "xhigh"]
ReasoningEffort = Literal["none", "minimal", "low", "medium", "high", "xhigh", "max"]
DEFAULT_MAX_TURNS = 500
_BASE_CONFIG = SettingsConfigDict(
case_sensitive=False,
@@ -35,6 +37,10 @@ class LlmSettings(BaseSettings):
"OLLAMA_API_BASE",
),
)
extra_headers: dict[str, str] | None = Field(
default=None,
alias="LLM_EXTRA_HEADERS",
)
reasoning_effort: ReasoningEffort = Field(default="high", alias="STRIX_REASONING_EFFORT")
force_required_tool_choice: bool = Field(
default=False,
@@ -44,6 +50,10 @@ class LlmSettings(BaseSettings):
default=True,
alias="STRIX_PROMPT_CACHE",
)
disable_streaming: bool = Field(
default=False,
alias="LLM_DISABLE_STREAMING",
)
timeout: int = Field(default=300, alias="LLM_TIMEOUT")
@@ -57,6 +67,10 @@ class DedupeSettings(BaseSettings):
)
api_key: str | None = Field(default=None, alias="DEDUPE_LLM_API_KEY")
api_base: str | None = Field(default=None, alias="DEDUPE_LLM_API_BASE")
extra_headers: dict[str, str] | None = Field(
default=None,
alias="DEDUPE_LLM_EXTRA_HEADERS",
)
class ContextSettings(BaseSettings):
@@ -83,15 +97,10 @@ class RuntimeSettings(BaseSettings):
model_config = _BASE_CONFIG
image: str = Field(
default="ghcr.io/usestrix/strix-sandbox:1.1.0",
default="ghcr.io/usestrix/strix-sandbox:1.2.0",
alias="STRIX_IMAGE",
)
backend: str = Field(default="docker", alias="STRIX_RUNTIME_BACKEND")
# Hard cap on a local target's size before we refuse to stream it into the
# sandbox file-by-file (the SDK copies every file individually, which stalls
# on large repos). Above this, the user must bind-mount via ``--mount``.
# Set to 0 (or less) to disable the pre-flight check entirely.
max_local_copy_mb: int = Field(default=1024, alias="STRIX_MAX_LOCAL_COPY_MB")
# Max screenshot/image tool outputs kept live per agent context (0 = none).
max_context_images: int = Field(default=3, ge=0, alias="STRIX_MAX_CONTEXT_IMAGES")
+153 -37
View File
@@ -24,6 +24,11 @@ logger = logging.getLogger(__name__)
Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed", "budget_paused"]
# Why an agent parked. The user can message any agent, so this - not the agent's
# position in the tree - decides whether waiting is bounded: only an agent waiting
# on other agents is re-checked on a timer.
WaitKind = Literal["user", "agents", "stalled"]
@dataclass(slots=True)
class AgentRuntime:
@@ -32,6 +37,8 @@ class AgentRuntime:
stream: Any | None = None
interrupt_on_message: bool = False
wake: asyncio.Event = field(default_factory=asyncio.Event)
mailbox: list[dict[str, Any]] = field(default_factory=list)
user_wake_required: bool = False
class AgentCoordinator:
@@ -44,7 +51,11 @@ class AgentCoordinator:
self.metadata: dict[str, dict[str, Any]] = {}
self.pending_counts: dict[str, int] = {}
self.errors: dict[str, str] = {}
self.recovery_counts: dict[str, int] = {}
self.idle_resume_counts: dict[str, int] = {}
self.wait_kinds: dict[str, WaitKind] = {}
self.runtimes: dict[str, AgentRuntime] = {}
self._parent_notified: set[str] = set()
self._lock = asyncio.Lock()
self._snapshot_path: Path | None = None
self.is_shutting_down = False
@@ -179,11 +190,59 @@ class AgentCoordinator:
if agent_id in self.statuses:
self.statuses[agent_id] = "running"
self.errors.pop(agent_id, None)
self.wait_kinds.pop(agent_id, None)
self.runtimes.setdefault(agent_id, AgentRuntime()).user_wake_required = False
self._parent_notified.discard(agent_id)
await self._maybe_snapshot()
async def park_waiting(self, agent_id: str) -> None:
async def park_waiting(self, agent_id: str, *, wait_kind: WaitKind) -> None:
"""Park an agent, recording what it is waiting on so the driver can time it."""
async with self._lock:
if agent_id in self.statuses:
self.wait_kinds[agent_id] = wait_kind
await self.set_status(agent_id, "waiting")
async def wait_kind_of(self, agent_id: str) -> WaitKind | None:
async with self._lock:
return self.wait_kinds.get(agent_id)
async def record_recovery(self, agent_id: str) -> int:
"""Count a turn that ended without a lifecycle tool call; return the new total.
Persisted so a resumed agent cannot earn a fresh nudge budget on every
auto-resume and loop forever.
"""
async with self._lock:
count = self.recovery_counts.get(agent_id, 0) + 1
self.recovery_counts[agent_id] = count
await self._maybe_snapshot()
return count
async def reset_recovery(self, agent_id: str) -> None:
"""Clear the nudge budget after real progress (new message or a lifecycle tool)."""
async with self._lock:
if self.recovery_counts.pop(agent_id, None) is None:
return
await self._maybe_snapshot()
async def record_idle_resume(self, agent_id: str) -> int:
"""Count an auto-resume that no message triggered; return the new total.
An agent that parks again after every auto-resume would otherwise burn a
model turn per timeout for the rest of the scan.
"""
async with self._lock:
count = self.idle_resume_counts.get(agent_id, 0) + 1
self.idle_resume_counts[agent_id] = count
await self._maybe_snapshot()
return count
async def reset_idle_resumes(self, agent_id: str) -> None:
async with self._lock:
if self.idle_resume_counts.pop(agent_id, None) is None:
return
await self._maybe_snapshot()
async def set_status(
self, agent_id: str, status: Status | str, *, error: str | None = None
) -> None:
@@ -195,55 +254,71 @@ class AgentCoordinator:
self.errors[agent_id] = error
elif status == "running":
self.errors.pop(agent_id, None)
if status == "running":
# Running again means a fresh stint that owes its parent its own notice.
self._parent_notified.discard(agent_id)
runtime = self.runtimes.setdefault(agent_id, AgentRuntime())
runtime.user_wake_required = status in {"failed", "crashed"}
runtime.wake.set()
logger.info("agent.status %s=%s", agent_id, status)
await self._maybe_snapshot()
async def send(self, target_agent_id: str, message: dict[str, Any]) -> bool:
"""Deliver a user/peer message by appending it to the target SDK session."""
if message.get("from") == "user" and self._budget_paused:
async def claim_parent_notice(self, agent_id: str) -> bool:
"""Reserve the one notice a child owes its parent when it stops running.
A completion report and a terminal notice carry the same information, so
whichever comes first claims the slot and the other is skipped.
"""
async with self._lock:
if agent_id in self._parent_notified:
return False
self._parent_notified.add(agent_id)
return True
async def send(
self, target_agent_id: str, message: dict[str, Any], *, interrupt: bool = True
) -> bool:
"""Queue a user/peer message in the target's mailbox and wake it."""
from_user = message.get("from") == "user"
if from_user and self._budget_paused:
await self.resume_from_budget_pause(exclude=target_agent_id)
async with self._lock:
if target_agent_id not in self.statuses:
logger.debug("agent.send dropped unknown target=%s", target_agent_id)
return False
runtime = self.runtimes.setdefault(target_agent_id, AgentRuntime())
session = runtime.session
stream = runtime.stream
interrupt = runtime.interrupt_on_message
if session is None:
logger.warning(
"agent.send dropped target=%s because its SDK session is not attached",
target_agent_id,
)
return False
try:
async with session_write_lock(session):
await session.add_items([self._message_to_session_item(message)])
except Exception:
logger.exception(
"agent.send failed to append to SDK session target=%s",
target_agent_id,
)
return False
async with self._lock:
runtime.mailbox.append(dict(message))
self.pending_counts[target_agent_id] = self.pending_counts.get(target_agent_id, 0) + 1
self.runtimes.setdefault(target_agent_id, AgentRuntime()).wake.set()
if stream is not None and interrupt:
if from_user:
runtime.user_wake_required = False
runtime.wake.set()
stream = runtime.stream
interrupt_on_message = runtime.interrupt_on_message
if stream is not None and interrupt and interrupt_on_message:
stream.cancel(mode="immediate")
await self._maybe_snapshot()
return True
async def wait_for_message(self, agent_id: str) -> None:
async def wait_for_message(self, agent_id: str, *, timeout: float | None = None) -> bool:
"""Wait until a message is ready for ``agent_id``; False on ``timeout``."""
while True:
async with self._lock:
runtime = self.runtimes.setdefault(agent_id, AgentRuntime())
reserve_exit = self._reserve_stopped and self.parent_of.get(agent_id) is not None
if self._budget_stopped or reserve_exit or self.pending_counts.get(agent_id, 0) > 0:
return
wake = self.runtimes.setdefault(agent_id, AgentRuntime()).wake
pending_ready = (
self.pending_counts.get(agent_id, 0) > 0 and not runtime.user_wake_required
)
if self._budget_stopped or reserve_exit or pending_ready:
return True
wake = runtime.wake
wake.clear()
await wake.wait()
if timeout is None:
await wake.wait()
else:
try:
await asyncio.wait_for(wake.wait(), timeout)
except TimeoutError:
return False
async def consume_pending(
self,
@@ -251,17 +326,38 @@ class AgentCoordinator:
*,
include_items: bool = False,
) -> tuple[int, list[Any]]:
"""Drain the agent's mailbox into its own SDK session."""
async with self._lock:
count = self.pending_counts.get(agent_id, 0)
runtime = self.runtimes.setdefault(agent_id, AgentRuntime())
queued = list(runtime.mailbox)
runtime.mailbox.clear()
count = max(self.pending_counts.get(agent_id, 0), len(queued))
self.pending_counts[agent_id] = 0
session = self.runtimes.get(agent_id, AgentRuntime()).session
session = runtime.session
if count <= 0:
return 0, []
items = [self._message_to_session_item(m) for m in queued]
if items:
if session is None:
logger.warning(
"agent %s has no SDK session attached; %d queued messages were not persisted",
agent_id,
len(items),
)
else:
try:
async with session_write_lock(session):
await session.add_items(items)
except Exception:
logger.exception(
"failed to append %d queued messages to the session of %s",
len(items),
agent_id,
)
await self._maybe_snapshot()
if not include_items or session is None:
if not include_items:
return count, []
items = await session.get_items()
return count, list(items[-count:])
return count, items
async def request_stop(self, agent_id: str) -> None:
async with self._lock:
@@ -287,12 +383,15 @@ class AgentCoordinator:
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
async def cancel_descendants_graceful(self, agent_id: str) -> None:
async def cancel_descendants_graceful(self, agent_id: str) -> list[str]:
"""Stop a subtree leaves-first and report which agents were stopped."""
async with self._lock:
order = self._subtree_order_locked(agent_id)
for aid in reversed(order):
stopped = list(reversed(order))
for aid in stopped:
await self.request_stop(aid)
await self._maybe_snapshot()
return stopped
async def attach_stream(
self,
@@ -372,6 +471,14 @@ class AgentCoordinator:
"names": dict(self.names),
"metadata": {aid: dict(md) for aid, md in self.metadata.items()},
"pending_counts": dict(self.pending_counts),
"recovery_counts": dict(self.recovery_counts),
"idle_resume_counts": dict(self.idle_resume_counts),
"wait_kinds": dict(self.wait_kinds),
"mailboxes": {
aid: [dict(m) for m in runtime.mailbox]
for aid, runtime in self.runtimes.items()
if runtime.mailbox
},
"errors": dict(self.errors),
"budget_stopped": self._budget_stopped,
"reserve_stopped": self._reserve_stopped,
@@ -386,6 +493,15 @@ class AgentCoordinator:
self.metadata = {aid: dict(md) for aid, md in snap.get("metadata", {}).items()}
self.pending_counts = dict(snap.get("pending_counts", {}))
self.errors = dict(snap.get("errors", {}))
self.recovery_counts = dict(snap.get("recovery_counts", {}))
self.idle_resume_counts = dict(snap.get("idle_resume_counts", {}))
self.wait_kinds = dict(snap.get("wait_kinds", {}))
mailboxes = snap.get("mailboxes", {})
if isinstance(mailboxes, dict):
for aid, msgs in mailboxes.items():
if isinstance(msgs, list):
runtime = self.runtimes.setdefault(aid, AgentRuntime())
runtime.mailbox = [dict(m) for m in msgs if isinstance(m, dict)]
self._budget_stopped = bool(snap.get("budget_stopped", False))
self._reserve_stopped = bool(snap.get("reserve_stopped", False))
self._budget_paused = bool(snap.get("budget_paused", False))
+363 -103
View File
@@ -9,6 +9,7 @@ import uuid
from collections.abc import Callable
from typing import TYPE_CHECKING, Any, cast
import litellm
from agents import RunConfig, Runner
from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError
from agents.sandbox.errors import ExecTransportError
@@ -16,11 +17,10 @@ from docker import errors as docker_errors # type: ignore[import-untyped, unuse
from openai import (
APIConnectionError,
APIError,
APIStatusError,
APITimeoutError,
RateLimitError,
)
from strix.config import codex
from strix.core.hooks import (
BudgetExceededError,
BudgetPausedError,
@@ -30,6 +30,8 @@ from strix.core.inputs import child_initial_input
from strix.core.sessions import (
enforce_image_budget,
open_agent_session,
replace_session_items,
seed_initial_input,
strip_all_images_from_session,
)
from strix.llm.compaction import is_context_overflow, maybe_compact
@@ -54,6 +56,23 @@ _INPUT_REJECTION_CODES = frozenset({400, 404, 422})
_MAX_COMPACTIONS_PER_CYCLE = 2
class ProviderRefusalError(AgentsException):
"""Raised when a provider returns a structured refusal instead of an exception."""
def _structured_provider_refusal(result: Any) -> str | None:
for item in getattr(result, "new_items", ()) or ():
raw_item = getattr(item, "raw_item", None)
for content in getattr(raw_item, "content", ()) or ():
if getattr(content, "type", None) != "refusal":
continue
refusal = getattr(content, "refusal", None)
if isinstance(refusal, str) and refusal.strip():
return refusal.strip()
return "The model provider refused this request."
return None
def _run_config_model(run_config: RunConfig) -> str | None:
return run_config.model if isinstance(run_config.model, str) else None
@@ -88,10 +107,9 @@ async def _compact_session(
)
_TRANSIENT_MODEL_STATUS_CODES = frozenset({408, 500, 502, 503, 504})
_MAX_TRANSIENT_MODEL_RETRIES = 4
_MAX_TRANSIENT_MODEL_RETRIES = 5
_TRANSIENT_MODEL_RETRY_BASE_DELAY_S = 2.0
_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 30.0
_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 90.0
def _model_error_status_code(exc: BaseException) -> int | None:
@@ -100,15 +118,16 @@ def _model_error_status_code(exc: BaseException) -> int | None:
def _is_transient_model_error(exc: BaseException) -> bool:
if isinstance(exc, RateLimitError):
if codex.is_content_guardrail_error(exc):
return False
if isinstance(exc, APITimeoutError | APIConnectionError):
if isinstance(
exc, APITimeoutError | APIConnectionError | TimeoutError | ConnectionError | OSError
):
return True
if isinstance(exc, APIStatusError):
return exc.status_code in _TRANSIENT_MODEL_STATUS_CODES
if isinstance(exc, APIError):
return _model_error_status_code(exc) is None
return False
code = _model_error_status_code(exc)
if code is not None:
return bool(litellm._should_retry(code))
return isinstance(exc, APIError)
def _transient_model_retry_delay(attempt: int) -> float:
@@ -116,6 +135,40 @@ def _transient_model_retry_delay(attempt: int) -> float:
return min(delay, _TRANSIENT_MODEL_RETRY_MAX_DELAY_S)
async def _salvage_stream_to_session(
session: Session,
pre_run_items: list[Any],
stream: Any,
agent_id: str,
) -> None:
"""Persist a crashed run's full history so a revived agent loses no context."""
if stream is None:
return
try:
replay = list(stream.to_input_list())
except Exception:
logger.exception("could not build salvage history for %s", agent_id)
return
desired = list(pre_run_items) + replay
if len(desired) <= len(pre_run_items):
return
try:
await replace_session_items(session, desired)
except Exception:
logger.exception("salvaging crashed run history failed for %s", agent_id)
async def _seed_and_prepare_first_input(
session: Session | None, initial_input: Any, *, start_parked: bool
) -> Any:
"""Persist the opening input up front so it survives a first-turn crash."""
if initial_input and session is not None and not start_parked:
with contextlib.suppress(Exception):
if await seed_initial_input(session, initial_input):
return []
return initial_input
async def run_agent_loop(
*,
agent: Any,
@@ -138,6 +191,10 @@ async def run_agent_loop(
)
result: RunResultBase | None = None
first_cycle_input = await _seed_and_prepare_first_input(
session, initial_input, start_parked=start_parked
)
budget_stopped = coordinator.budget_stopped
reserve_stopped = coordinator.reserve_stopped
if budget_stopped:
@@ -151,31 +208,17 @@ async def run_agent_loop(
await coordinator.send(agent_id, _reserve_notice())
if not (start_parked and interactive):
if interactive:
with contextlib.suppress(BudgetPausedError):
result = await _run_cycle(
agent,
coordinator,
agent_id,
input_data=initial_input,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
interactive=interactive,
event_sink=event_sink,
hooks=hooks,
)
else:
result = await _run_noninteractive_until_lifecycle(
with contextlib.suppress(BudgetPausedError):
result = await _run_until_lifecycle(
agent,
coordinator,
agent_id,
initial_input=initial_input,
initial_input=first_cycle_input,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
interactive=interactive,
event_sink=event_sink,
hooks=hooks,
)
@@ -184,8 +227,9 @@ async def run_agent_loop(
return result
while True:
timeout = await _plain_waiting_timeout(coordinator, agent_id)
try:
await coordinator.wait_for_message(agent_id)
woke = await coordinator.wait_for_message(agent_id, timeout=timeout)
except asyncio.CancelledError:
return result
@@ -197,18 +241,46 @@ async def run_agent_loop(
await coordinator.set_status(agent_id, "stopped")
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
if woke:
# Real input is real progress, so the nudge budget starts over. A bare
# auto-resume is not: it must not hand a wedged agent a fresh budget.
await coordinator.reset_recovery(agent_id)
await coordinator.reset_idle_resumes(agent_id)
else:
idle_resumes = await coordinator.record_idle_resume(agent_id)
if idle_resumes >= _MAX_IDLE_AUTO_RESUMES:
logger.warning(
"agent %s auto-resumed %d times without hearing from anyone; "
"leaving it parked until a real message arrives",
agent_id,
idle_resumes,
)
await coordinator.park_waiting(agent_id, wait_kind="stalled")
await _notify_parent_on_stall(coordinator, agent_id)
continue
logger.info("agent %s reached its waiting timeout; auto-resuming", agent_id)
await coordinator.send(
agent_id,
{
"from": "system",
"type": "auto_resume",
"content": "Waiting timeout reached. Resuming execution.",
},
interrupt=False,
)
await coordinator.consume_pending(agent_id)
with contextlib.suppress(BudgetPausedError):
result = await _run_cycle(
result = await _run_until_lifecycle(
agent,
coordinator,
agent_id,
input_data=[],
initial_input=[],
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
interactive=interactive,
interactive=True,
event_sink=event_sink,
hooks=hooks,
)
@@ -359,7 +431,10 @@ async def respawn_subagents(
await coordinator.set_status(child_id, "crashed")
async def _run_noninteractive_until_lifecycle(
_INTERACTIVE_TOOL_RECOVERY_LIMIT = 3
async def _run_until_lifecycle(
agent: Any,
coordinator: AgentCoordinator,
agent_id: str,
@@ -369,14 +444,20 @@ async def _run_noninteractive_until_lifecycle(
context: dict[str, Any],
max_turns: int,
session: Session | None,
interactive: bool,
event_sink: StreamEventSink | None,
hooks: RunHooks[dict[str, Any]] | None,
) -> RunResultBase | None:
"""Non-chat mode keeps running until finish_scan / agent_finish settles status."""
"""Drive an agent until an explicit lifecycle tool settles its status.
A turn that ends without ``finish_scan``, ``agent_finish``,
``respond_to_user``, or ``wait_for_agents`` leaves the agent ``running``:
plain text never terminates a run and never yields to the user. Such a turn
is nudged back into a tool call, bounded by a recovery limit.
"""
result: RunResultBase | None = None
input_data: Any = initial_input
invalid_final_outputs = 0
invalid_final_output_limit = max(1, max_turns)
recovery_limit = _INTERACTIVE_TOOL_RECOVERY_LIMIT if interactive else max(1, max_turns)
while True:
if coordinator.budget_stopped:
@@ -387,7 +468,143 @@ async def _run_noninteractive_until_lifecycle(
await coordinator.set_status(agent_id, "stopped")
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
result = await _run_cycle(
if interactive:
result = await _run_cycle_parked(
agent,
coordinator,
agent_id,
input_data=input_data,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
event_sink=event_sink,
hooks=hooks,
)
else:
result = await _run_cycle(
agent,
coordinator,
agent_id,
input_data=input_data,
run_config=run_config,
context=context,
max_turns=max_turns,
session=session,
interactive=False,
event_sink=event_sink,
hooks=hooks,
)
status = await _agent_status(coordinator, agent_id)
if status != "running":
await coordinator.reset_recovery(agent_id)
return result
recoveries = await coordinator.record_recovery(agent_id)
logger.warning(
"agent %s ended a turn without a lifecycle tool call (interactive=%s); "
"forcing tool continuation (%d/%d): %s",
agent_id,
interactive,
recoveries,
recovery_limit,
_final_output_preview(result),
)
if recoveries >= recovery_limit:
return await _exhausted_recovery(coordinator, agent_id, result, interactive=interactive)
input_data = await _append_tool_required_message(
session=session,
context=context,
attempt=recoveries,
limit=recovery_limit,
interactive=interactive,
)
async def _exhausted_recovery(
coordinator: AgentCoordinator,
agent_id: str,
result: RunResultBase | None,
*,
interactive: bool,
) -> RunResultBase | None:
"""Settle an agent that never recovered into a tool call.
Interactive runs park instead of dying: a human is attached and can message
any agent, so the scan stays resumable. Autonomous runs have nobody to
resume them, so they fail loudly.
"""
if not interactive:
await coordinator.set_status(agent_id, "crashed")
await notify_parent_on_terminal(coordinator, agent_id, "crashed")
raise MaxTurnsExceeded(
"Agent exhausted recovery attempts without calling finish_scan or agent_finish."
)
logger.warning(
"agent %s exhausted tool-call recovery attempts; parking until a message arrives",
agent_id,
)
await coordinator.park_waiting(agent_id, wait_kind="stalled")
# A parked child owes its parent a completion report it can no longer send. The
# parent is an agent, not a watching human, so nothing else tells it to stop
# waiting and it burns its full timeout on a message that is never coming.
await _notify_parent_on_stall(coordinator, agent_id)
return result
_WAITING_AUTO_RESUME_TIMEOUT_S = 300.0
# An agent that parks again after every auto-resume makes no progress, so stop
# spending a model turn per timeout and leave it parked for a real message.
_MAX_IDLE_AUTO_RESUMES = 3
async def _plain_waiting_timeout(
coordinator: AgentCoordinator,
agent_id: str,
) -> float | None:
"""Auto-resume timeout for a parked agent; None waits until a message arrives.
Driven by what the agent is waiting on, not by where it sits in the graph:
the user can message any agent, so an agent awaiting a human parks
indefinitely whether or not it is the root. Only an agent awaiting other
agents is re-checked on a timer, and only until it has spent its idle
budget re-parking without hearing anything.
"""
async with coordinator._lock:
status = coordinator.statuses.get(agent_id)
has_error = agent_id in coordinator.errors
runtime = coordinator.runtimes.get(agent_id)
gated = runtime.user_wake_required if runtime is not None else False
wait_kind = coordinator.wait_kinds.get(agent_id)
idle_resumes = coordinator.idle_resume_counts.get(agent_id, 0)
if status != "waiting" or has_error or gated:
return None
if wait_kind != "agents" or idle_resumes >= _MAX_IDLE_AUTO_RESUMES:
return None
return _WAITING_AUTO_RESUME_TIMEOUT_S
async def _run_cycle_parked(
agent: Any,
coordinator: AgentCoordinator,
agent_id: str,
*,
input_data: Any,
run_config: RunConfig,
context: dict[str, Any],
max_turns: int,
session: Session | None,
event_sink: StreamEventSink | None,
hooks: RunHooks[dict[str, Any]] | None,
) -> RunResultBase | None:
"""Interactive run cycle that parks on any error instead of killing the runner."""
try:
return await _run_cycle(
agent,
coordinator,
agent_id,
@@ -396,39 +613,17 @@ async def _run_noninteractive_until_lifecycle(
context=context,
max_turns=max_turns,
session=session,
interactive=False,
interactive=True,
event_sink=event_sink,
hooks=hooks,
)
status = await _agent_status(coordinator, agent_id)
if status != "running":
return result
invalid_final_outputs += 1
logger.warning(
"agent %s produced non-lifecycle final output in non-interactive mode; "
"forcing tool continuation (%d/%d): %s",
agent_id,
invalid_final_outputs,
invalid_final_output_limit,
_final_output_preview(result),
)
if invalid_final_outputs >= invalid_final_output_limit:
await coordinator.set_status(agent_id, "crashed")
await _notify_parent_on_terminal(coordinator, agent_id, "crashed")
raise MaxTurnsExceeded(
"Agent exhausted non-interactive recovery attempts without calling "
"finish_scan or agent_finish."
)
input_data = await _append_noninteractive_tool_required_message(
session=session,
context=context,
attempt=invalid_final_outputs,
limit=invalid_final_output_limit,
)
except (BudgetExceededError, BudgetPausedError, SubagentBudgetReservedError):
raise
except Exception as exc:
logger.exception("error escaped the run cycle for %s; parking as failed", agent_id)
await coordinator.set_status(agent_id, "failed", error=str(exc) or type(exc).__name__)
await notify_parent_on_terminal(coordinator, agent_id, "failed")
return None
async def _run_cycle( # noqa: PLR0912, PLR0915
@@ -449,6 +644,8 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
compactions = 0
model_retries = 0
while True:
stream: Any = None
pre_run_items: list[Any] = []
try:
await coordinator.mark_running(agent_id)
if session is not None:
@@ -462,6 +659,8 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
await _compact_session(agent, session, run_config, force=False)
except Exception:
logger.exception("proactive compaction failed for %s", agent_id)
with contextlib.suppress(Exception):
pre_run_items = list(await session.get_items())
stream = Runner.run_streamed(
agent,
input=input_data,
@@ -482,6 +681,8 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
logger.exception("stream event sink failed for %s", agent_id)
if stream.run_loop_exception is not None:
raise stream.run_loop_exception
if refusal := _structured_provider_refusal(stream):
raise ProviderRefusalError(refusal)
except (BudgetExceededError, BudgetPausedError, SubagentBudgetReservedError):
raise
except RuntimeError as stream_exc:
@@ -572,6 +773,13 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
if session is not None:
input_data = []
continue
if session is not None:
await _salvage_stream_to_session(session, pre_run_items, stream, agent_id)
if isinstance(exc, ProviderRefusalError):
logger.warning("agent %s refused by the model provider: %s", agent_id, exc)
await coordinator.set_status(agent_id, "failed", error=str(exc))
await notify_parent_on_terminal(coordinator, agent_id, "failed")
return None
if not interactive:
raise
if isinstance(exc, MaxTurnsExceeded):
@@ -582,28 +790,10 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
status = "crashed"
logger.exception("agent run failed for %s; parking as %s", agent_id, status)
await coordinator.set_status(agent_id, status, error=str(exc) or type(exc).__name__)
await _notify_parent_on_terminal(coordinator, agent_id, status)
await notify_parent_on_terminal(coordinator, agent_id, status)
return None
else:
await _settle_run_result(coordinator, agent_id, interactive)
return stream
async def _settle_run_result(
coordinator: AgentCoordinator,
agent_id: str,
interactive: bool,
) -> None:
async with coordinator._lock:
current_status = coordinator.statuses.get(agent_id)
if current_status != "running":
return
if not interactive:
return
await coordinator.set_status(agent_id, "waiting")
return cast("RunResultBase | None", stream)
async def _agent_status(coordinator: AgentCoordinator, agent_id: str) -> Status | None:
@@ -621,23 +811,37 @@ def _final_output_preview(result: RunResultBase | None) -> str:
return text[:300]
async def _append_noninteractive_tool_required_message(
async def _append_tool_required_message(
*,
session: Session | None,
context: dict[str, Any],
attempt: int,
limit: int,
interactive: bool,
) -> list[dict[str, str]]:
finish_tool = "finish_scan" if context.get("parent_id") is None else "agent_finish"
message = (
"Your previous response ended the autonomous Strix run without a lifecycle tool call. "
"That is invalid in non-interactive mode; plain text final answers are ignored. "
"Continue immediately and call exactly one tool. "
f"If your work is complete, call {finish_tool}. "
"If you are blocked waiting for another agent, call wait_for_message. "
"Otherwise use the appropriate execution or planning tool. "
f"This is recovery attempt {attempt}/{limit}."
)
if interactive:
message = (
"Your previous message ended a turn without a tool call. Plain text never ends "
"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. "
"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. "
f"This is recovery attempt {attempt}/{limit}."
)
else:
message = (
"Your previous response ended the autonomous Strix run without a lifecycle tool "
"call. That is invalid in non-interactive mode; plain text final answers are "
"ignored. Continue immediately and call exactly one tool. "
f"If your work is complete, call {finish_tool}. "
"If you are blocked waiting for another agent, call wait_for_agents. "
"Otherwise use the appropriate execution or planning tool. "
f"This is recovery attempt {attempt}/{limit}."
)
item = {"role": "user", "content": message}
if session is None:
return [item]
@@ -647,6 +851,11 @@ async def _append_noninteractive_tool_required_message(
_TERMINAL_NOTICE = {
"completed": (
"[Agent completed] {name} ({agent_id}) finished and is no longer running, but it "
"sent no completion report. Stop waiting on this child; ask it directly if you "
"need its results."
),
"crashed": (
"[Agent crash] {name} ({agent_id}) terminated unexpectedly. "
"Stop waiting on this child unless you want to message it again."
@@ -657,14 +866,44 @@ _TERMINAL_NOTICE = {
"message it again."
),
"stopped": (
"[Agent capped] {name} ({agent_id}) hit its turn limit and was stopped "
"before finishing. It will not send a completion report, so stop waiting "
"on this child; account for its capped subtask and continue."
"[Agent stopped] {name} ({agent_id}) was stopped before finishing (turn limit "
"or an explicit stop). It will not send a completion report, so stop waiting "
"on this child; account for its unfinished subtask and continue."
),
}
async def _notify_parent_on_terminal(
_STALL_NOTICE = (
"[Agent stalled] {name} ({agent_id}) kept ending turns without a tool call and is "
"parked until it receives a message. It will not send a completion report on its "
"own: either message it with a concrete next step to unblock it, or stop waiting on "
"it and account for its unfinished subtask."
)
async def _notify_parent_on_stall(
coordinator: AgentCoordinator,
agent_id: str,
) -> None:
"""Tell the parent that a child parked mid-task, so it stops waiting blindly."""
async with coordinator._lock:
parent = coordinator.parent_of.get(agent_id)
name = coordinator.names.get(agent_id, agent_id)
if parent is None:
return
await coordinator.send(
parent,
{
"from": agent_id,
"type": "stalled",
"priority": "high",
"content": _STALL_NOTICE.format(name=name, agent_id=agent_id),
},
interrupt=False,
)
async def notify_parent_on_terminal(
coordinator: AgentCoordinator,
agent_id: str,
status: str,
@@ -677,6 +916,8 @@ async def _notify_parent_on_terminal(
name = coordinator.names.get(agent_id, agent_id)
if parent is None:
return
if not await coordinator.claim_parent_notice(agent_id):
return
await coordinator.send(
parent,
{
@@ -685,6 +926,7 @@ async def _notify_parent_on_terminal(
"priority": "high",
"content": template.format(name=name, agent_id=agent_id),
},
interrupt=False,
)
@@ -710,6 +952,21 @@ async def _notify_root_on_budget_reserve(coordinator: AgentCoordinator) -> None:
await coordinator.send(root, _reserve_notice())
async def _notify_parent_on_exit(
coordinator: AgentCoordinator,
agent_id: str,
) -> None:
"""Backstop for a child whose loop ended without telling its parent.
Every terminal state counts, including ``completed``: a child that skips its
completion report leaves the parent waiting on a message nobody will send.
"""
status = await _agent_status(coordinator, agent_id)
if status is None:
return
await notify_parent_on_terminal(coordinator, agent_id, status)
async def _start_child_runner(
*,
parent_ctx: dict[str, Any],
@@ -763,6 +1020,9 @@ async def _start_child_runner(
logger.info("child %s stopped after reaching the scan budget limit", child_id)
except SubagentBudgetReservedError:
logger.info("child %s stopped at the sub-agent budget reserve", child_id)
finally:
if not coordinator.is_shutting_down:
await _notify_parent_on_exit(coordinator, child_id)
task_handle = asyncio.create_task(_child_loop(), name=f"agent-{name}-{child_id}")
await coordinator.attach_runtime(child_id, task=task_handle)
+24 -6
View File
@@ -24,9 +24,6 @@ if TYPE_CHECKING:
from strix.config.settings import ReasoningEffort
DEFAULT_MAX_TURNS = 500
def _accepts_required_tool_choice(model_name: str | None) -> bool:
name = (model_name or "").strip().lower()
for prefix in ("litellm/", "any-llm/"):
@@ -62,8 +59,11 @@ def build_root_task(scan_config: dict[str, Any]) -> str:
)
elif ttype == "local_code":
path = details.get("target_path", "unknown")
suffix = ", read-only mount" if details.get("mount") else ""
sections["Local Codebases"].append(f"- {path} (available at: {workspace_path}{suffix})")
sections["Local Codebases"].append(
f"- {path} (available at: {workspace_path}; "
"this is the user's real directory, mounted live and writable — "
".git/.agents/.codex are read-only)"
)
elif ttype == "web_application":
sections["URLs"].append(f"- {details.get('target_url', '')}")
elif ttype == "ip_address":
@@ -132,12 +132,14 @@ def make_model_settings(
force_required_tool_choice: bool = False,
request_timeout: float | None = None,
prompt_cache: bool = True,
extra_headers: dict[str, str] | None = None,
) -> ModelSettings:
model_settings = ModelSettings(
parallel_tool_calls=False,
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,
)
if (
reasoning_effort is not None
@@ -145,7 +147,7 @@ def make_model_settings(
and model_supports_reasoning(model_name)
):
model_settings = model_settings.resolve(
ModelSettings(reasoning=Reasoning(effort=reasoning_effort)),
_reasoning_settings(reasoning_effort, model_settings.extra_args),
)
if force_required_tool_choice and _accepts_required_tool_choice(model_name):
model_settings = model_settings.resolve(ModelSettings(tool_choice="required"))
@@ -160,6 +162,22 @@ def make_model_settings(
return model_settings
def _reasoning_settings(
effort: ReasoningEffort,
extra_args: dict[str, Any] | None,
) -> ModelSettings:
"""``max`` is not in the OpenAI SDK's ``Reasoning.effort`` enum, so send it as
a raw body field instead — also keeping it clear of LiteLLM's DeepSeek mapping,
which collapses every ``reasoning_effort`` level to plain thinking-enabled.
Providers that don't support ``max`` reject the request.
"""
if effort != "max":
return ModelSettings(reasoning=Reasoning(effort=effort))
return ModelSettings(
extra_args={**(extra_args or {}), "extra_body": {"reasoning_effort": "max"}},
)
def _prompt_cache_extra_args(model_name: str) -> dict[str, Any] | None:
"""LiteLLM ``cache_control_injection_points`` for Claude prompt caching.
+12 -1
View File
@@ -23,6 +23,7 @@ from strix.config.models import (
configure_sdk_model_defaults,
uses_chat_completions_tool_schema,
)
from strix.config.settings import DEFAULT_MAX_TURNS
from strix.core.agents import AgentCoordinator
from strix.core.execution import (
respawn_subagents,
@@ -33,7 +34,6 @@ from strix.core.execution import (
)
from strix.core.hooks import BudgetExceededError, ReportUsageHooks, recomputed_budget_flags
from strix.core.inputs import (
DEFAULT_MAX_TURNS,
build_root_task,
build_scope_context,
make_model_settings,
@@ -53,6 +53,8 @@ if TYPE_CHECKING:
from agents.memory import SQLiteSession
from agents.result import RunResultBase
from strix.runtime.status import StatusSink
logger = logging.getLogger(__name__)
@@ -120,6 +122,7 @@ async def run_strix_scan(
event_sink: StreamEventSink | None = None,
root_instructions_override: str | None = None,
extra_system_prompt_context: dict[str, Any] | None = None,
status_sink: StatusSink | None = None,
) -> RunResultBase | None:
"""Run or resume one Strix scan against a sandbox.
@@ -129,6 +132,11 @@ async def run_strix_scan(
context before prompt rendering. Child agents keep the standard scan prompt
and context.
"""
def report(phase: str) -> None:
if status_sink is not None:
status_sink(phase)
if scan_id is None:
scan_id = f"scan-{uuid.uuid4().hex[:8]}"
@@ -219,7 +227,9 @@ async def run_strix_scan(
scan_id,
image=image,
local_sources=local_sources or [],
status_sink=status_sink,
)
report("Waiting for the first model response")
logger.info("Sandbox ready for scan %s", scan_id)
sandbox_session = bundle["session"]
@@ -250,6 +260,7 @@ async def run_strix_scan(
force_required_tool_choice=settings.llm.force_required_tool_choice,
request_timeout=settings.llm.timeout,
prompt_cache=settings.llm.prompt_cache,
extra_headers=settings.llm.extra_headers,
)
run_config = RunConfig(
model=resolved_model,
+13
View File
@@ -7,6 +7,7 @@ import logging
from typing import TYPE_CHECKING, Any, cast
from weakref import WeakKeyDictionary
from agents.items import ItemHelpers
from agents.memory import SQLiteSession
@@ -26,6 +27,18 @@ def open_agent_session(agent_id: str, path: Path) -> SQLiteSession:
return SQLiteSession(session_id=agent_id, db_path=path)
async def seed_initial_input(session: Session, initial_input: Any) -> bool:
"""Commit an agent's opening identity/task input before its first run cycle."""
items = ItemHelpers.input_to_new_input_list(initial_input)
if not items:
return False
async with session_write_lock(session):
if await session.get_items():
return False
await session.add_items(items)
return True
_IMAGE_REJECTED_TEXT = "[image rejected by the model]"
_IMAGE_ELIDED_TEXT = "[older screenshot elided to bound context memory]"
_INHERITED_IMAGE_TEXT = "[screenshot omitted from inherited context]"
+12 -1
View File
@@ -13,7 +13,7 @@ from rich.panel import Panel
from rich.text import Text
from strix.config import load_settings
from strix.core.inputs import DEFAULT_MAX_TURNS
from strix.config.settings import DEFAULT_MAX_TURNS
from strix.core.runner import run_strix_scan
from strix.report.state import ReportState, set_global_report_state
from strix.runtime import session_manager
@@ -21,6 +21,7 @@ from strix.runtime import session_manager
from .utils import (
build_live_stats_text,
format_vulnerability_report,
has_model_response,
)
@@ -135,11 +136,17 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915
set_global_report_state(report_state)
startup_phase: list[str] = ["Starting up"]
def create_live_status() -> Panel:
status_text = Text()
status_text.append("Penetration test in progress", style="bold #22c55e")
status_text.append("\n\n")
if not has_model_response(report_state):
status_text.append(f"{startup_phase[0]}...", style="dim")
status_text.append("\n\n")
stats_text = build_live_stats_text(report_state)
if stats_text:
status_text.append(stats_text)
@@ -152,6 +159,9 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915
padding=(1, 2),
)
def _note_startup_phase(phase: str) -> None:
startup_phase[:] = [phase]
try:
console.print()
@@ -186,6 +196,7 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915
interactive=bool(getattr(args, "interactive", False)),
max_budget_usd=getattr(args, "max_budget_usd", None),
max_turns=getattr(args, "max_turns", DEFAULT_MAX_TURNS),
status_sink=_note_startup_phase,
)
finally:
stop_updates.set()
+66 -62
View File
@@ -11,9 +11,6 @@ import sys
from datetime import UTC, datetime
from pathlib import Path
from agents.model_settings import ModelSettings
from agents.models.interface import ModelTracing
from docker.errors import DockerException
from rich.console import Console
from rich.panel import Panel
from rich.text import Text
@@ -24,17 +21,8 @@ from strix.config import (
load_settings,
persist_current,
)
from strix.config.models import (
RECOMMENDED_MODEL_NAMES,
StrixProvider,
configure_sdk_model_defaults,
is_known_openai_bare_model,
is_recommended_or_frontier_model,
)
from strix.core.inputs import DEFAULT_MAX_TURNS
from strix.config.settings import DEFAULT_MAX_TURNS
from strix.core.paths import run_dir_for, runtime_state_dir
from strix.interface.cli import run_cli
from strix.interface.tui import run_tui
from strix.interface.update_check import (
is_binary_install,
notify_update,
@@ -45,12 +33,11 @@ from strix.interface.update_check import (
from strix.interface.utils import (
assign_workspace_subdirs,
build_final_stats_text,
build_mount_targets_info,
check_docker_connection,
check_mountable_dir,
clone_repository,
collect_local_sources,
dedupe_local_targets,
find_oversized_local_targets,
generate_run_name,
image_exists,
infer_target_type,
@@ -61,8 +48,6 @@ from strix.interface.utils import (
rewrite_localhost_targets,
validate_config_file,
)
from strix.report.state import get_global_report_state
from strix.report.writer import read_run_record, write_run_record
from strix.telemetry import posthog, scarf
from strix.telemetry.logging import configure_dependency_logging
@@ -171,8 +156,8 @@ def validate_environment() -> None:
error_text.append("", style="white")
error_text.append("STRIX_REASONING_EFFORT", style="bold cyan")
error_text.append(
" - Reasoning effort level: none, minimal, low, medium, high, xhigh "
"(default: high)\n",
" - Reasoning effort level: none, minimal, low, medium, high, xhigh, "
"max (default: high)\n",
style="white",
)
@@ -310,6 +295,18 @@ def _subscription_error_hint(exc: BaseException) -> str | None:
async def warm_up_llm(show_model_warning: bool = True) -> None:
from agents.model_settings import ModelSettings
from agents.models.interface import ModelTracing
from strix.config.models import (
RECOMMENDED_MODEL_NAMES,
StrixProvider,
configure_sdk_model_defaults,
is_known_openai_bare_model,
is_recommended_or_frontier_model,
)
from strix.core.inputs import make_model_settings
console = Console()
logger.info("Warming up LLM connection")
@@ -382,7 +379,13 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
model.get_response(
system_instructions="You are a helpful assistant.",
input="Reply with just 'OK'.",
model_settings=ModelSettings(),
model_settings=make_model_settings(
None,
model_name=raw_model,
request_timeout=llm.timeout,
prompt_cache=False,
extra_headers=llm.extra_headers,
),
tools=[],
output_schema=None,
handoffs=[],
@@ -404,7 +407,19 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
# Match the runtime path: send the dedupe key/endpoint per call so a
# separate-provider dedupe model authenticates during warm-up too.
deduper_extra = _dedupe_extra_args(settings.dedupe)
deduper_settings = ModelSettings(extra_args=deduper_extra or None)
# A dedicated dedupe model may route to another provider, which must
# never receive the main endpoint's headers; it has its own
# DEDUPE_LLM_EXTRA_HEADERS.
deduper_settings = make_model_settings(
None,
model_name=dedupe_model,
request_timeout=llm.timeout,
prompt_cache=False,
extra_headers=settings.dedupe.extra_headers,
)
if deduper_extra:
merged = {**(deduper_settings.extra_args or {}), **deduper_extra}
deduper_settings = deduper_settings.resolve(ModelSettings(extra_args=merged))
await asyncio.wait_for(
deduper.get_response(
system_instructions="You are a helpful assistant.",
@@ -508,9 +523,6 @@ Examples:
# Local code analysis
strix --target ./my-project
# Large local repository (bind-mounted read-only instead of copied)
strix --mount ./huge-monorepo
# Domain penetration test
strix --target example.com
@@ -554,8 +566,9 @@ Examples:
type=str,
action="append",
help="Target to test (URL, repository, local directory path, domain name, or IP address). "
"Local directories are mounted into the sandbox writable. "
"Can be specified multiple times for multi-target scans. "
"Fresh runs require at least one of --target, --target-list, or --mount.",
"Fresh runs require --target or --target-list.",
)
parser.add_argument(
"--target-list",
@@ -565,15 +578,6 @@ Examples:
help="Path to a file containing targets, one per non-empty, non-comment line. "
"Can be specified multiple times and combined with --target.",
)
parser.add_argument(
"--mount",
type=str,
action="append",
metavar="PATH",
help="Bind-mount a local directory into the sandbox (read-only) instead of "
"copying it file-by-file. Use this for large repositories that are too big to "
"stream into the container. Can be specified multiple times.",
)
parser.add_argument(
"--instruction",
type=str,
@@ -704,9 +708,9 @@ Examples:
args.user_explicit_instruction = args.instruction if args.resume else None
if args.resume:
if args.target or args.target_list or args.mount:
if args.target or args.target_list:
parser.error(
"Cannot combine --resume with --target/--target-list/--mount. "
"Cannot combine --resume with --target/--target-list. "
"--resume picks up where the prior run left off, including the "
"original target list."
)
@@ -720,9 +724,9 @@ Examples:
f"or remove --resume to start over with the same targets."
)
else:
if not args.target and not args.target_list and not args.mount:
if not args.target and not args.target_list:
parser.error(
"the following arguments are required: -t/--target, --target-list, or --mount "
"the following arguments are required: -t/--target or --target-list "
"(or use --resume <run_name> to continue a prior scan)"
)
args.targets_info = []
@@ -745,37 +749,20 @@ Examples:
args.targets_info.append(
{"type": target_type, "details": target_dict, "original": display_target}
)
except ValueError:
parser.error(f"Invalid target '{target}'")
try:
args.targets_info.extend(build_mount_targets_info(args.mount or []))
except ValueError as e:
parser.error(str(e))
except ValueError as e:
parser.error(f"Invalid target '{target}': {e}")
args.targets_info = dedupe_local_targets(args.targets_info)
assign_workspace_subdirs(args.targets_info)
rewrite_localhost_targets(args.targets_info, HOST_GATEWAY_HOSTNAME)
max_local_copy_mb = load_settings().runtime.max_local_copy_mb
max_copy_bytes = max_local_copy_mb * 1024 * 1024
oversized = find_oversized_local_targets(args.targets_info, max_copy_bytes)
if oversized:
details = "; ".join(
f"{path} ({size / (1024 * 1024):.0f} MB)" for path, size in oversized
)
parser.error(
f"Local target too large to stream into the sandbox: {details}. "
f"The limit is {max_local_copy_mb} MB "
"(set STRIX_MAX_LOCAL_COPY_MB to change it). Re-run with "
"--mount <path> to bind-mount the directory instead of copying it."
)
return args
def _persist_run_record(args: argparse.Namespace) -> None:
from strix.report.writer import write_run_record
run_dir = run_dir_for(args.run_name)
run_dir.mkdir(parents=True, exist_ok=True)
run_record = {
@@ -799,6 +786,8 @@ def _persist_run_record(args: argparse.Namespace) -> None:
def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser) -> None:
"""Populate ``args.targets_info`` and friends from a prior run's run.json."""
from strix.report.writer import read_run_record
run_dir = run_dir_for(args.resume)
state_path = run_dir / "run.json"
if not state_path.exists():
@@ -819,6 +808,12 @@ def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser
if not isinstance(target, dict):
continue
details = target.get("details") or {}
if target.get("type") == "local_code" and details.get("target_path"):
try:
check_mountable_dir(Path(details["target_path"]).expanduser())
except ValueError as exc:
parser.error(f"--resume {args.resume}: {exc}")
continue
if target.get("type") != "repository":
continue
cloned = details.get("cloned_repo_path")
@@ -833,8 +828,7 @@ def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser
if args.instruction is None:
args.instruction = state.get("instruction")
if state.get("local_sources"):
args.local_sources = state.get("local_sources")
args.local_sources = collect_local_sources(args.targets_info)
if state.get("diff_scope"):
args.diff_scope = state.get("diff_scope")
persisted_scan_mode = state.get("scan_mode")
@@ -843,6 +837,8 @@ def _load_resume_state(args: argparse.Namespace, parser: argparse.ArgumentParser
def display_completion_message(args: argparse.Namespace, results_path: Path) -> None:
from strix.report.state import get_global_report_state
console = Console()
report_state = get_global_report_state()
@@ -884,7 +880,7 @@ def display_completion_message(args: argparse.Namespace, results_path: Path) ->
view_text = Text()
view_text.append("\n")
view_text.append("View", style="dim")
view_text.append(" ")
view_text.append(" ")
view_text.append(f"strix view {args.run_name}", style="#22c55e")
panel_parts.extend(["\n", view_text])
@@ -922,6 +918,8 @@ def display_completion_message(args: argparse.Namespace, results_path: Path) ->
def pull_docker_image() -> None:
from docker.errors import DockerException
console = Console()
client = check_docker_connection()
@@ -1068,11 +1066,17 @@ def main() -> None:
posthog.start(**_telemetry_start_kwargs)
scarf.start(**_telemetry_start_kwargs)
from strix.report.state import get_global_report_state
exit_reason = "user_exit"
try:
if args.non_interactive:
from strix.interface.cli import run_cli
asyncio.run(run_cli(args))
else:
from strix.interface.tui import run_tui
asyncio.run(run_tui(args))
except KeyboardInterrupt:
exit_reason = "interrupted"
+27 -5
View File
@@ -15,6 +15,7 @@ from typing import TYPE_CHECKING, Any, ClassVar
if TYPE_CHECKING:
from pygments.token import _TokenType
from textual.timer import Timer
from rich.align import Align
@@ -33,8 +34,8 @@ from textual.widgets.tree import TreeNode
from strix.config import load_settings
from strix.config.models import is_recommended_or_frontier_model
from strix.config.settings import DEFAULT_MAX_TURNS
from strix.core.hooks import BudgetExceededError
from strix.core.inputs import DEFAULT_MAX_TURNS
from strix.core.runner import run_strix_scan
from strix.interface.tui.live_view import TuiLiveView
from strix.interface.tui.messages import send_user_message_to_agent
@@ -352,7 +353,7 @@ class VulnerabilityDetailScreen(ModalScreen): # type: ignore[misc]
if not token_value:
continue
color = None
tt = token_type
tt: _TokenType | None = token_type
while tt:
if tt in colors:
color = colors[tt]
@@ -814,6 +815,8 @@ class StrixTUIApp(App): # type: ignore[misc]
self._scan_stop_event = threading.Event()
self._scan_completed = threading.Event()
self._scan_error: BaseException | None = None
self._startup_status = "Starting up"
self._startup_status_step = 0
self._error_noted_agents: set[str] = set()
self._budget_pause_notified = False
@@ -1040,9 +1043,9 @@ class StrixTUIApp(App): # type: ignore[misc]
name=names.get(agent_id, agent_id),
parent_id=parent_of.get(agent_id),
status=status,
error_message=error,
error_message=error or "",
)
if status in {"failed", "crashed"} and error:
if error:
if agent_id not in self._error_noted_agents:
self._error_noted_agents.add(agent_id)
self.live_view.record_agent_error(agent_id, error)
@@ -1110,7 +1113,9 @@ class StrixTUIApp(App): # type: ignore[misc]
self,
) -> tuple[Any, str | None]:
if not self.selected_agent_id:
return self._get_chat_placeholder_content("Loading...", "placeholder-no-agent")
return self._get_chat_placeholder_content(
f"{self._startup_status}...", f"placeholder-no-agent-{self._startup_status_step}"
)
events = self._gather_agent_events(self.selected_agent_id)
@@ -1292,6 +1297,10 @@ class StrixTUIApp(App): # type: ignore[misc]
text.append("Send a message to continue", style="dim")
keymap = keymap_styled([("ctrl-q", "quit")])
else:
error_msg = agent_data.get("error_message") or ""
if error_msg:
text.append(error_msg, style="red")
text.append(" \u00b7 ", style="dim")
text.append("Send message to resume", style="dim")
return (text, keymap, False)
@@ -1520,6 +1529,7 @@ class StrixTUIApp(App): # type: ignore[misc]
max_budget_usd=getattr(self.args, "max_budget_usd", None),
max_turns=getattr(self.args, "max_turns", DEFAULT_MAX_TURNS),
event_sink=self._capture_sdk_event,
status_sink=self._capture_startup_status,
),
)
@@ -1551,6 +1561,18 @@ class StrixTUIApp(App): # type: ignore[misc]
self._scan_thread = threading.Thread(target=scan_target, daemon=True)
self._scan_thread.start()
def _capture_startup_status(self, phase: str) -> None:
try:
self.call_from_thread(self._record_startup_status, phase)
except RuntimeError:
self._record_startup_status(phase)
def _record_startup_status(self, phase: str) -> None:
self._startup_status = phase
self._startup_status_step += 1
if not self.show_splash and not self.selected_agent_id:
self.call_later(self._update_chat_view)
def _capture_sdk_event(self, agent_id: str, event: Any) -> None:
try:
self.call_from_thread(self._record_sdk_event, agent_id, event)
+8 -6
View File
@@ -20,7 +20,7 @@ class TuiLiveView:
self.events: list[dict[str, Any]] = []
self._next_event_id = 1
self._open_assistant_event_by_agent: dict[str, dict[str, Any]] = {}
self._tool_event_by_call_id: dict[str, dict[str, Any]] = {}
self._tool_event_by_agent_and_call_id: dict[tuple[str, str], dict[str, Any]] = {}
def hydrate_from_run_dir(self, run_dir: Path) -> None:
state_dir = runtime_state_dir(run_dir)
@@ -82,7 +82,7 @@ class TuiLiveView:
current["parent_id"] = parent_id
if status is not None:
current["status"] = status
if error_message:
if error_message is not None:
current["error_message"] = error_message
current["updated_at"] = now
@@ -223,7 +223,8 @@ class TuiLiveView:
timestamp: str | None = None,
) -> None:
call_id = call["call_id"]
existing = self._tool_event_by_call_id.get(call_id)
event_key = (agent_id, call_id)
existing = self._tool_event_by_agent_and_call_id.get(event_key)
tool_data = {
"tool_name": call["tool_name"],
"args": call["args"],
@@ -233,7 +234,7 @@ class TuiLiveView:
}
if existing is None:
event = self._append_event(agent_id, "tool", tool_data, timestamp=timestamp)
self._tool_event_by_call_id[call_id] = event
self._tool_event_by_agent_and_call_id[event_key] = event
else:
existing["data"].update(tool_data)
self._bump_event(existing, timestamp=timestamp)
@@ -249,7 +250,8 @@ class TuiLiveView:
timestamp: str | None = None,
) -> None:
call_id = output["call_id"]
event = self._tool_event_by_call_id.get(call_id)
event_key = (agent_id, call_id)
event = self._tool_event_by_agent_and_call_id.get(event_key)
if event is None:
event = self._append_event(
agent_id,
@@ -263,7 +265,7 @@ class TuiLiveView:
},
timestamp=timestamp,
)
self._tool_event_by_call_id[call_id] = event
self._tool_event_by_agent_and_call_id[event_key] = event
result = _parse_json_value(output["output"])
event["data"]["result"] = result
@@ -6,6 +6,7 @@ from . import (
notes_renderer,
proxy_renderer,
reporting_renderer,
respond_renderer,
shell_renderer,
thinking_renderer,
todo_renderer,
@@ -23,6 +24,7 @@ __all__ = [
"proxy_renderer",
"render_tool_widget",
"reporting_renderer",
"respond_renderer",
"shell_renderer",
"thinking_renderer",
"todo_renderer",
@@ -117,8 +117,8 @@ class AgentFinishRenderer(BaseToolRenderer):
@register_tool_renderer
class WaitForMessageRenderer(BaseToolRenderer):
tool_name: ClassVar[str] = "wait_for_message"
class WaitForAgentsRenderer(BaseToolRenderer):
tool_name: ClassVar[str] = "wait_for_agents"
css_classes: ClassVar[list[str]] = ["tool-call", "agents-graph-tool"]
@classmethod
@@ -0,0 +1,35 @@
from typing import Any, ClassVar
from rich.text import Text
from textual.widgets import Static
from .agent_message_renderer import AgentMessageRenderer
from .base_renderer import BaseToolRenderer
from .registry import register_tool_renderer
@register_tool_renderer
class RespondToUserRenderer(BaseToolRenderer):
"""Render a reply as the agent's own prose, not as a tool call.
``respond_to_user`` carries the message the user is meant to read, so it
gets the same markdown treatment as a plain assistant turn.
"""
tool_name: ClassVar[str] = "respond_to_user"
css_classes: ClassVar[list[str]] = ["tool-call", "respond-tool"]
@classmethod
def render(cls, tool_data: dict[str, Any]) -> Static:
args = tool_data.get("args", {})
message = args.get("message", "")
text = Text()
if message:
text.append_text(AgentMessageRenderer.render_simple(message))
text.append("\n\n")
text.append("", style="#6b7280")
text.append("waiting for your reply", style="dim")
css_classes = cls.get_css_classes(tool_data.get("status", "unknown"))
return Static(text, classes=css_classes)
+107 -100
View File
@@ -11,11 +11,10 @@ import tempfile
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from urllib.error import HTTPError, URLError
from urllib.parse import urlparse
from urllib.request import Request, urlopen
import docker
import requests
from docker.errors import DockerException, ImageNotFound
from rich.console import Console
from rich.panel import Panel
@@ -291,6 +290,11 @@ def _detail_value(usage: dict[str, Any], detail_key: str, value_key: str) -> int
return _int_stat(details, value_key)
def has_model_response(report_state: Any) -> bool:
usage = _llm_usage(report_state)
return bool(usage) and _int_stat(usage, "requests") > 0
def _build_llm_usage_stats(
stats_text: Text,
report_state: Any,
@@ -1088,13 +1092,12 @@ def resolve_diff_scope_context(
def _is_http_git_repo(url: str) -> bool:
check_url = f"{url.rstrip('/')}/info/refs?service=git-upload-pack"
try:
req = Request(check_url, headers={"User-Agent": "git/strix"}) # noqa: S310
with urlopen(req, timeout=10) as resp: # noqa: S310 # nosec B310
return "x-git-upload-pack-advertisement" in resp.headers.get("Content-Type", "")
except HTTPError as e:
return e.code == 401
except (URLError, OSError, ValueError):
resp = requests.get(check_url, headers={"User-Agent": "git/strix"}, timeout=10)
except (requests.RequestException, ValueError):
return False
if resp.status_code >= 400:
return resp.status_code == 401
return "x-git-upload-pack-advertisement" in resp.headers.get("Content-Type", "")
def infer_target_type(target: str) -> tuple[str, dict[str, str]]: # noqa: PLR0911
@@ -1133,6 +1136,7 @@ def infer_target_type(target: str) -> tuple[str, dict[str, str]]: # noqa: PLR09
try:
if path.exists():
if path.is_dir():
check_mountable_dir(path)
return "local_code", {"target_path": str(path.resolve())}
raise ValueError(f"Path exists but is not a directory: {target}")
except (OSError, RuntimeError) as e:
@@ -1261,7 +1265,7 @@ def collect_local_sources(targets_info: list[dict[str, Any]]) -> list[dict[str,
{
"source_path": details["target_path"],
"workspace_subdir": workspace_subdir,
"mount": bool(details.get("mount", False)),
"protect_metadata": True,
}
)
@@ -1270,123 +1274,126 @@ def collect_local_sources(targets_info: list[dict[str, Any]]) -> list[dict[str,
{
"source_path": details["cloned_repo_path"],
"workspace_subdir": workspace_subdir,
"mount": False,
"protect_metadata": False,
}
)
return local_sources
def directory_size_bytes(path: Path) -> int:
"""Total size in bytes of regular files under ``path`` (symlinks not followed).
# Refused along with everything under them.
_FORBIDDEN_MOUNT_TREES = frozenset(
{
"/bin",
"/sbin",
"/usr",
"/etc",
"/lib",
"/lib64",
"/nix/store",
"/run/current-system/sw",
"/Applications",
"/Library",
"/System",
"/dev",
"/boot",
"/proc",
"/sys",
}
)
Best-effort: files that disappear or can't be stat'd mid-walk are skipped.
Used as a cheap (stat-only) pre-flight to estimate the cost of streaming a
local target into the sandbox before we actually try to copy it.
# Refused themselves, but they hold projects too, so their contents are fine.
_FORBIDDEN_MOUNT_ROOTS = frozenset(
{
"/",
"/private",
"/var",
"/opt",
"/home",
"/root",
"/srv",
"/Users",
"/Volumes",
}
)
Directories that can't be listed (e.g. permission denied) are logged and
skipped rather than silently dropped — so an under-count is at least
visible — but the returned total then excludes their contents.
"""
_FORBIDDEN_WINDOWS_TREE_NAMES = frozenset(
{"windows", "program files", "program files (x86)", "programdata"}
)
def _on_walk_error(error: OSError) -> None:
logger.warning("Could not read %s while measuring size: %s", error.filename, error)
total = 0
for root, _dirs, files in os.walk(path, followlinks=False, onerror=_on_walk_error):
for name in files:
file_path = os.path.join(root, name) # noqa: PTH118
try:
if os.path.islink(file_path): # noqa: PTH114
continue
total += os.path.getsize(file_path) # noqa: PTH202
except OSError:
continue
return total
_FORBIDDEN_MOUNT_DIR_NAMES = frozenset(
{
".ssh",
".tsh",
".brev",
".gnupg",
".aws",
".azure",
".kube",
".docker",
".config",
".npm",
".pki",
".terraform.d",
}
)
def find_oversized_local_targets(
targets_info: list[dict[str, Any]], max_bytes: int
) -> list[tuple[str, int]]:
"""Return ``(path, size_bytes)`` for non-mounted local targets over ``max_bytes``.
Mounted targets are bind-mounted rather than copied, so their size is
irrelevant and they are excluded. A ``max_bytes`` of zero or less disables
the check entirely (returns no targets).
"""
if max_bytes <= 0:
return []
oversized: list[tuple[str, int]] = []
for target in targets_info:
if target.get("type") != "local_code":
continue
details = target.get("details") or {}
if details.get("mount"):
continue
target_path = details.get("target_path")
if not target_path:
continue
size = directory_size_bytes(Path(target_path))
if size > max_bytes:
oversized.append((target_path, size))
return oversized
def _is_within(path: Path, ancestor: Path) -> bool:
ancestor_parts = [part.casefold() for part in ancestor.parts]
path_parts = [part.casefold() for part in path.parts]
return path_parts[: len(ancestor_parts)] == ancestor_parts
def build_mount_targets_info(mount_paths: list[str]) -> list[dict[str, Any]]:
"""Build ``targets_info`` entries for ``--mount`` directories.
def check_mountable_dir(path: Path) -> None:
resolved = path.resolve()
if not resolved.is_dir():
raise ValueError(f"'{path}' is not an existing directory.")
Each path must be an existing local directory; it is bind-mounted into the
sandbox (read-only) instead of being copied file-by-file. Raises
``ValueError`` for an empty path, or one that does not exist or is not a
directory.
"""
targets_info: list[dict[str, Any]] = []
for raw in mount_paths:
if not raw or not raw.strip():
raise ValueError("--mount path must not be empty.")
path = Path(raw).expanduser()
try:
resolved = path.resolve()
is_dir = resolved.is_dir()
except (OSError, RuntimeError) as e:
raise ValueError(f"Invalid mount path '{raw}': {e!s}") from e
if not is_dir:
raise ValueError(
f"Mount path '{raw}' is not an existing directory. "
"--mount requires a path to a local directory."
)
targets_info.append(
{
"type": "local_code",
"details": {"target_path": str(resolved), "mount": True},
"original": str(resolved),
}
# Both the literal and the resolved form: macOS reaches /etc through the
# /private/etc symlink, and only the resolved path is compared below.
exact = {str(Path(root)).casefold() for root in _FORBIDDEN_MOUNT_ROOTS}
exact |= {str(Path(root).resolve()).casefold() for root in _FORBIDDEN_MOUNT_ROOTS}
exact.add(str(Path.home().resolve()).casefold())
tree_roots = set(_FORBIDDEN_MOUNT_TREES)
if os.name == "nt":
drive = Path(resolved.anchor)
tree_roots |= {str(drive / name) for name in _FORBIDDEN_WINDOWS_TREE_NAMES}
exact.add(str(drive / "Users").casefold())
trees = [Path(root) for root in tree_roots] + [Path(root).resolve() for root in tree_roots]
if (
str(resolved).casefold() in exact
or resolved.parent == resolved
or any(_is_within(resolved, tree) for tree in trees)
):
raise ValueError(
f"Refusing to mount '{resolved}' into the sandbox: it is a system "
"or home directory, not a codebase. Point the target at the "
"project directory you want tested."
)
credential = next(
(part for part in resolved.parts if part.casefold() in _FORBIDDEN_MOUNT_DIR_NAMES), None
)
if credential is not None:
raise ValueError(
f"Refusing to mount '{resolved}' into the sandbox: '{credential}' "
"holds credentials, not code."
)
return targets_info
def dedupe_local_targets(targets_info: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Collapse local_code targets that resolve to the same path.
When a directory is supplied both as a copied ``--target`` and via
``--mount`` (or as duplicate values of either), keep one entry and prefer
the bind-mounted one — so the same tree is never both streamed in and
mounted. Order is preserved; non-local targets pass through untouched.
"""
result: list[dict[str, Any]] = []
index_by_path: dict[str, int] = {}
seen_paths: set[str] = set()
for target in targets_info:
details = target.get("details") or {}
path = details.get("target_path")
if target.get("type") != "local_code" or not path:
result.append(target)
continue
existing = index_by_path.get(path)
if existing is None:
index_by_path[path] = len(result)
if path not in seen_paths:
seen_paths.add(path)
result.append(target)
elif details.get("mount") and not (result[existing].get("details") or {}).get("mount"):
result[existing] = target # bind mount supersedes the copied entry
return result
+10 -14
View File
@@ -15,12 +15,12 @@ import base64
import contextlib
import json
import logging
import urllib.error
import urllib.request
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
import requests
from strix.config.loader import load_settings
@@ -147,21 +147,17 @@ def _post_json(path: str, payload: dict[str, Any], *, timeout: int) -> tuple[int
map, not raised.
"""
url = f"{_app_url()}{path}"
body = json.dumps(payload).encode("utf-8")
request = urllib.request.Request( # noqa: S310 - fixed https relay URL
url,
data=body,
headers={"Content-Type": "application/json", "Accept": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(request, timeout=timeout) as response: # noqa: S310 # nosec B310
return response.status, _parse_body(response.read())
except urllib.error.HTTPError as exc:
return exc.code, _parse_body(exc.read())
except (urllib.error.URLError, TimeoutError, OSError) as exc:
response = requests.post(
url,
json=payload,
headers={"Accept": "application/json"},
timeout=timeout,
)
except requests.RequestException as exc:
logger.warning("relay request to %s failed: %s", path, exc)
raise RelayError("unavailable") from exc
return response.status_code, _parse_body(response.content)
def _parse_body(raw: bytes) -> dict[str, Any]:
@@ -54,7 +54,7 @@ export default function AgentCommsRenderer({ toolName, args }: ToolRendererProps
);
}
if (toolName === "wait_for_message") {
if (toolName === "wait_for_agents") {
const reason = (args.reason as string) ?? "";
return (
<div className="flex items-center gap-2">
@@ -0,0 +1,20 @@
"use client";
import type { ToolRendererProps } from "@/types/events";
import Markdown from "./Markdown";
/**
* `respond_to_user` carries the message the user is meant to read, so it renders
* as the agent's own prose rather than as a tool call.
*/
export default function RespondRenderer({ args }: ToolRendererProps) {
const message = (args.message as string) ?? "";
if (!message) return null;
return (
<div>
<Markdown text={message} />
<div className="mt-1.5 text-[#888] text-[13px]">waiting for your reply</div>
</div>
);
}
@@ -24,6 +24,7 @@ import NotesRenderer from "./NotesRenderer";
import TodoRenderer from "./TodoRenderer";
import FallbackRenderer from "./FallbackRenderer";
import LoadSkillRenderer from "./LoadSkillRenderer";
import RespondRenderer from "./RespondRenderer";
/**
* Tool-renderer mapping data-driven, keyed by the engine's tool *family*.
@@ -104,10 +105,10 @@ const CATEGORY_TOOLS: Record<ToolCategory, readonly string[]> = {
proxy: ["list_requests", "view_request", "repeat_request", "list_sitemap", "view_sitemap_entry", "scope_rules", "send_request"],
reporting: ["create_vulnerability_report", "list_reports", "get_report"],
thinking: ["think"],
agents: ["create_agent", "agent_finish", "send_message_to_agent", "wait_for_message", "view_agent_graph", "stop_agent"],
agents: ["create_agent", "agent_finish", "send_message_to_agent", "wait_for_agents", "view_agent_graph", "stop_agent"],
search: ["web_search"],
// scan_start_info / subagent_start_info are strix-app synthetic events; finish_scan is the engine's
lifecycle: ["scan_start_info", "subagent_start_info", "finish_scan"],
lifecycle: ["scan_start_info", "subagent_start_info", "finish_scan", "respond_to_user"],
notes: ["create_note", "delete_note", "update_note", "list_notes", "get_note"],
skills: ["load_skill"],
todos: ["create_todo", "list_todos", "update_todo", "mark_todo_done", "mark_todo_pending", "delete_todo"],
@@ -127,6 +128,7 @@ const TOOL_CATEGORY: Record<string, ToolCategory> = Object.fromEntries(
*/
const RENDERER_OVERRIDES: Partial<Record<string, ComponentType<ToolRendererProps>>> = {
finish_scan: FinishRenderer,
respond_to_user: RespondRenderer,
apply_patch: ApplyPatchRenderer,
view_image: ViewImageRenderer,
list_reports: ReportListRenderer,
@@ -140,7 +142,8 @@ const RENDERER_OVERRIDES: Partial<Record<string, ComponentType<ToolRendererProps
const ICON_OVERRIDES: Partial<Record<string, ToolIconMeta>> = {
agent_finish: { icon: Flag, color: "text-cyan-400" },
send_message_to_agent: { icon: MessageCircle, color: "text-cyan-400" },
wait_for_message: { icon: MessageCircle, color: "text-cyan-400" },
wait_for_agents: { icon: MessageCircle, color: "text-cyan-400" },
respond_to_user: { icon: MessageCircle, color: "text-emerald-400" },
view_agent_graph: { icon: Eye, color: "text-cyan-400" },
stop_agent: { icon: Ban, color: "text-red-400" },
scan_start_info: { icon: Crosshair, color: "text-emerald-400" },
+11 -4
View File
@@ -107,8 +107,11 @@ def resolve_run_dir(base_dir: Path, run_param: str | None, default_run_dir: Path
return candidate
# Name of the cookie carrying the per-process session capability.
SESSION_COOKIE = "strix_viewer_session"
# Prefix of the cookie carrying the per-process session capability. The bound
# port is appended (``strix_viewer_session_<port>``) because browsers scope
# cookies by host only, never by port: concurrent viewers on 127.0.0.1 would
# otherwise share one cookie slot and clobber each other's session.
SESSION_COOKIE_PREFIX = "strix_viewer_session"
class _ViewerState:
@@ -135,6 +138,9 @@ class _ViewerState:
# enough to steer a live scan, trigger a report, or browse history --
# the token is never handed to a caller who merely reaches ``/``.
self.session_token = secrets.token_urlsafe(32)
# Finalized in ``serve()`` once the port is known (the server binds
# after this state is constructed); see SESSION_COOKIE_PREFIX.
self.cookie_name = SESSION_COOKIE_PREFIX
def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]:
@@ -476,7 +482,7 @@ def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]:
the browser this process handed the page to can pass. A direct
caller on an exposed port has no cookie and is rejected.
"""
supplied = self._cookies().get(SESSION_COOKIE, "")
supplied = self._cookies().get(state.cookie_name, "")
return bool(supplied) and secrets.compare_digest(supplied, state.session_token)
def _token_presented(self, query: dict[str, list[str]]) -> bool:
@@ -512,7 +518,7 @@ def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]:
# SameSite=Strict (never sent from a cross-site context).
self.send_header(
"Set-Cookie",
f"{SESSION_COOKIE}={state.session_token}; Path=/; HttpOnly; SameSite=Strict",
f"{state.cookie_name}={state.session_token}; Path=/; HttpOnly; SameSite=Strict",
)
self.end_headers()
self.wfile.write(content)
@@ -586,6 +592,7 @@ def serve(
httpd.daemon_threads = True
bound_port = int(httpd.server_address[1])
state.cookie_name = f"{SESSION_COOKIE_PREFIX}_{bound_port}"
url = f"http://{host}:{bound_port}"
thread = threading.Thread(target=httpd.serve_forever, name="strix-viewer", daemon=True)
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -6,7 +6,7 @@
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<meta name="color-scheme" content="dark" />
<title>Strix Results</title>
<script type="module" crossorigin src="./assets/index-DzvI_0HX.js"></script>
<script type="module" crossorigin src="./assets/index-CGvQq6oe.js"></script>
<link rel="stylesheet" crossorigin href="./assets/index-C3kQ5kk8.css">
</head>
<body>
+44 -12
View File
@@ -12,15 +12,20 @@ from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Any
import litellm
from agents.model_settings import ModelSettings
from agents.models.interface import ModelTracing
from litellm.exceptions import BadRequestError, ContextWindowExceededError
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from strix.config import load_settings
from strix.config.models import StrixProvider
from strix.core.inputs import make_model_settings
from strix.core.sessions import replace_session_items, session_write_lock
from strix.llm.context_budget import context_window, count_tokens, output_limit
if TYPE_CHECKING:
from agents.items import ModelResponse
from agents.memory import Session
@@ -268,26 +273,53 @@ def _checkpoint_item(summary: str) -> dict[str, Any]:
}
def _extract_text(response: ModelResponse) -> str:
parts: list[str] = []
for item in response.output:
if not isinstance(item, ResponseOutputMessage):
continue
parts.extend(
chunk.text
for chunk in item.content
if isinstance(chunk, ResponseOutputText) and chunk.text
)
return "".join(parts)
async def _summarize(model: str, prompt: str, max_tokens: int) -> str | None:
llm = load_settings().llm
model_settings = make_model_settings(
None,
model_name=model,
request_timeout=llm.timeout,
prompt_cache=False,
extra_headers=llm.extra_headers,
).resolve(ModelSettings(max_tokens=max_tokens))
try:
response = await litellm.acompletion(
model=model,
messages=[{"role": "user", "content": prompt}],
max_tokens=max_tokens,
api_key=llm.api_key,
api_base=llm.api_base,
timeout=llm.timeout,
response = (
await StrixProvider()
.get_model(model)
.get_response(
system_instructions=None,
input=prompt,
model_settings=model_settings,
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
previous_response_id=None,
conversation_id=None,
prompt=None,
)
)
except Exception:
logger.exception("compaction summary call failed for model %s", model)
return None
try:
content = response.choices[0].message.content
except (AttributeError, IndexError, KeyError):
content = _extract_text(response).strip()
if not content:
logger.warning("compaction summary returned no content")
return None
return content.strip() if isinstance(content, str) and content.strip() else None
return content
async def maybe_compact(
+13 -2
View File
@@ -17,7 +17,14 @@ logger = logging.getLogger(__name__)
# LiteLLM keys models without the routing prefix users type (``openai/``,
# ``litellm/``, ``ollama/`` ...). Strip a leading provider segment on lookup.
_STRIPPABLE_PREFIXES = ("openai/", "litellm/", "any-llm/", "ollama/", "ollama_chat/")
_STRIPPABLE_PREFIXES = (
"openai/",
"chatgpt/",
"litellm/",
"any-llm/",
"ollama/",
"ollama_chat/",
)
_DEFAULT_OUTPUT_TOKENS = 8_192
@@ -38,7 +45,11 @@ def _safe_get_model_info(model: str) -> dict[str, Any] | None:
@lru_cache(maxsize=128)
def _model_info(model: str) -> dict[str, int]:
for candidate in (model, _lookup_key(model)):
lookup_key = _lookup_key(model)
# Provider-qualified ChatGPT lookups may start a synchronous device-login
# poll. LiteLLM keys the metadata by the underlying model slug.
candidates = (lookup_key,) if model.startswith("chatgpt/") else (model, lookup_key)
for candidate in candidates:
info = _safe_get_model_info(candidate)
if info is not None:
return {
+8 -3
View File
@@ -51,17 +51,24 @@ def _dedupe_extra_args(dedupe: DedupeSettings) -> dict[str, str]:
def _dedupe_model_settings(
dedupe: DedupeSettings, model_name: str, request_timeout: float | None
) -> ModelSettings:
llm = load_settings().llm
settings = make_model_settings(
dedupe.reasoning_effort,
model_name=model_name,
force_required_tool_choice=False,
request_timeout=request_timeout,
# The main model's headers apply only when dedupe falls back to the main
# model; a dedicated dedupe model may route to another provider, which
# 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,
)
extra = _dedupe_extra_args(dedupe)
if extra:
settings = settings.resolve(ModelSettings(extra_args=extra))
return settings
DEDUPE_SYSTEM_PROMPT = """You are an expert vulnerability report deduplication judge.
Your task is to determine if a candidate vulnerability report describes the SAME vulnerability
as any existing report.
@@ -347,9 +354,7 @@ async def check_duplicate(
response = await model.get_response(
system_instructions=DEDUPE_SYSTEM_PROMPT,
input=user_msg,
model_settings=_dedupe_model_settings(
dedupe, resolved_model, settings.llm.timeout
),
model_settings=_dedupe_model_settings(dedupe, resolved_model, settings.llm.timeout),
tools=[],
output_schema=None,
handoffs=[],
+74
View File
@@ -1,6 +1,7 @@
import json
import logging
import subprocess
import threading
from collections.abc import Callable
from datetime import UTC, datetime
from importlib.metadata import PackageNotFoundError, version
@@ -95,6 +96,8 @@ def get_global_report_state() -> Optional["ReportState"]:
def set_global_report_state(report_state: "ReportState") -> None:
global _global_report_state # noqa: PLW0603
_global_report_state = report_state
# New run: drop any streamed-cost entries a prior run left unconsumed.
streamed_openrouter_costs.clear()
class ReportState:
@@ -507,6 +510,72 @@ class ReportState:
self._sync_llm_usage_record()
def openrouter_stream_cost(usage: Any) -> float | None:
"""Total OpenRouter-reported cost from a raw stream ``usage`` block, or None.
Non-BYOK responses bill everything to ``usage.cost``. BYOK responses put the
OpenRouter fee in ``usage.cost`` (often 0) and the provider charge in
``usage.cost_details.upstream_inference_cost``, so BYOK totals sum the two.
"""
if not isinstance(usage, dict):
return None
total = 0.0
cost = usage.get("cost")
if isinstance(cost, int | float) and cost > 0:
total += float(cost)
if bool(usage.get("is_byok")):
details = usage.get("cost_details")
upstream = details.get("upstream_inference_cost") if isinstance(details, dict) else None
if isinstance(upstream, int | float) and upstream > 0:
total += float(upstream)
return total if total > 0 else None
def _response_id(completion_response: Any) -> str | None:
response_id = getattr(completion_response, "id", None)
if response_id is None and isinstance(completion_response, dict):
response_id = cast("dict[str, Any]", completion_response).get("id")
return response_id if isinstance(response_id, str) and response_id else None
class StreamedOpenRouterCosts:
"""Correlates OpenRouter's per-stream cost from the parser to the cost callback.
LiteLLM rebuilds streamed responses from token-only chunks and drops the
``usage.cost`` OpenRouter reports in its final stream chunk (its non-streamed
path preserves it; streaming snapshots hidden params at stream start). Every
scan streams, so the OpenRouter streaming handler (see strix.config.models)
records the cost here keyed by response id, and the callback takes it back out
for the matching rebuilt response. Entries are removed on read; ``clear()``
runs per scan so nothing accumulates across runs.
"""
def __init__(self) -> None:
self._costs: dict[str, float] = {}
self._lock = threading.Lock()
def remember(self, response_id: Any, usage: Any) -> None:
cost = openrouter_stream_cost(usage)
if cost is None or not (isinstance(response_id, str) and response_id):
return
with self._lock:
self._costs[response_id] = cost
def take(self, completion_response: Any) -> float | None:
response_id = _response_id(completion_response)
if response_id is None:
return None
with self._lock:
return self._costs.pop(response_id, None)
def clear(self) -> None:
with self._lock:
self._costs.clear()
streamed_openrouter_costs = StreamedOpenRouterCosts()
def litellm_cost_callback(
kwargs: Any,
completion_response: Any,
@@ -541,6 +610,11 @@ def litellm_cost_callback(
if cost is None:
cost = _usage_reported_cost(completion_response)
# Recover the exact OpenRouter cost the streaming handler stashed for this
# response — LiteLLM drops it from streamed usage, so nothing above sees it.
if cost is None:
cost = streamed_openrouter_costs.take(completion_response)
if cost is None:
cost = _estimate_response_cost(kwargs, completion_response)
+25 -13
View File
@@ -31,16 +31,11 @@ async def _docker_backend(
``docker`` lazily so deployments that target a non-Docker
backend don't need the docker-py library installed.
``session.start()`` is what materializes the manifest entries
(LocalDir copies and manifest-declared volume/FUSE mounts) into the
running container the SDK's ``client.create()`` only builds the inner
session object without applying the manifest. ``async with session:``
would call it too, but Strix manages session lifetime explicitly via
``client.delete()`` so we trigger ``start()`` ourselves.
``bind_mounts`` are host directories (e.g. large repos passed via
``--mount``) bind-mounted read-only; unlike manifest entries they are
applied by Docker at container-create time, not by ``start()``.
``session.start()`` is what materializes the manifest into the running
container the SDK's ``client.create()`` only builds the inner session
object without applying it. ``async with session:`` would call it too, but
Strix manages session lifetime explicitly via ``client.delete()`` so we
trigger ``start()`` ourselves.
"""
import docker
from agents.sandbox.sandboxes.docker import DockerSandboxClientOptions
@@ -59,6 +54,8 @@ _BACKENDS: dict[str, SandboxBackend] = {
"docker": _docker_backend,
}
_BIND_MOUNT_BACKENDS: set[str] = {"docker"}
def get_backend(name: str) -> SandboxBackend:
"""Return the backend factory for ``name`` or raise.
@@ -78,15 +75,30 @@ def get_backend(name: str) -> SandboxBackend:
return backend
def register_backend(name: str, backend: SandboxBackend) -> None:
def register_backend(
name: str,
backend: SandboxBackend,
*,
supports_bind_mounts: bool = False,
) -> None:
"""Register a custom backend under ``name``.
Intended for downstream users who ship their own runtime register
before any ``session_manager.create_or_reuse`` call. Re-registering
an existing name overwrites the prior entry.
an existing name overwrites the prior entry. ``supports_bind_mounts``
defaults to False: a remote runtime cannot see the caller's filesystem, so
it is handed local sources as manifest entries to upload instead.
"""
_BACKENDS[name] = backend
logger.info("Registered sandbox backend: %s", name)
if supports_bind_mounts:
_BIND_MOUNT_BACKENDS.add(name)
else:
_BIND_MOUNT_BACKENDS.discard(name)
logger.info("Registered sandbox backend: %s (bind mounts: %s)", name, supports_bind_mounts)
def backend_supports_bind_mounts(name: str) -> bool:
return name in _BIND_MOUNT_BACKENDS
def supported_backends() -> list[str]:
+19 -5
View File
@@ -110,6 +110,19 @@ def _apply_log_limits(create_kwargs: dict[str, Any]) -> None:
)
def _apply_run_labels(create_kwargs: dict[str, Any]) -> None:
run_id = os.getenv("STRIX_RUN_ID")
if not run_id:
return
labels = create_kwargs.setdefault("labels", {})
if not isinstance(labels, dict):
return
labels["strix-run-id"] = run_id
run_type = os.getenv("STRIX_RUN_TYPE")
if run_type:
labels["strix-run-type"] = run_type
class StrixDockerSandboxSession(DockerSandboxSession):
sandbox_network: str = ""
@@ -222,19 +235,20 @@ class StrixDockerSandboxClient(DockerSandboxClient):
_apply_sandbox_network(create_kwargs)
_apply_resource_limits(create_kwargs)
_apply_log_limits(create_kwargs)
_apply_run_labels(create_kwargs)
# Strix injection: host bind mounts (e.g. large repos passed via --mount)
# that bypass the SDK's file-by-file LocalDir copy.
bind_mounts = getattr(self, "strix_bind_mounts", ())
# Strix injection: local source trees, sorted shallowest-first so a
# nested spec lands on top of the tree it covers.
bind_mounts = self.strix_bind_mounts or ()
if bind_mounts:
mounts = create_kwargs.setdefault("mounts", [])
for spec in bind_mounts:
for spec in sorted(bind_mounts, key=lambda s: str(s["target"]).count("/")):
mounts.append(
DockerSDKMount(
target=spec["target"],
source=spec["source"],
type="bind",
read_only=spec.get("read_only", True),
read_only=spec.get("read_only", False),
)
)
-120
View File
@@ -1,120 +0,0 @@
"""Symlink-safe staging for ``LocalDir`` manifest uploads.
The sandbox SDK's ``LocalDir`` walker refuses to copy symlinks at all — it
raises ``LocalDirReadError(reason="symlink_not_supported")`` on the first one
as a path-escape / TOCTOU safeguard. Real source trees (especially JS/TS
monorepos with workspace or shared-config links) routinely commit symlinks, so
handing such a tree straight to ``LocalDir`` aborts the upload before the agent
even starts.
:func:`stage_symlink_safe_dir` returns a path that is always safe to hand to
``LocalDir``:
* a tree with no symlinks is used as-is (no copy);
* otherwise the tree is copied into a temp directory with symlinks resolved:
- a link whose target stays inside the tree is *dereferenced* (its target
content is materialized in place), so the agent still sees the file;
- a link that escapes the tree, dangles, or forms a cycle is *dropped* and
never followed. Refusing to follow out-of-tree links preserves the walker's
path-escape safety and keeps host/out-of-tree content from leaking into the
(hostile) sandbox.
Regular files are hard-linked when possible (falling back to a copy across
devices), so the staged tree adds negligible disk for the non-symlink bulk.
"""
from __future__ import annotations
import logging
import os
import shutil
import tempfile
from pathlib import Path
logger = logging.getLogger(__name__)
_STAGING_PREFIX = "strix-localdir-"
def _is_within(target: Path, root: Path) -> bool:
"""Return whether ``target`` is ``root`` itself or nested under it."""
if target == root:
return True
try:
target.relative_to(root)
except ValueError:
return False
return True
def tree_has_symlink(root: Path) -> bool:
"""Return whether ``root`` contains any symlink (file or directory)."""
for dirpath, dirnames, filenames in os.walk(root, followlinks=False):
base = Path(dirpath)
for name in (*dirnames, *filenames):
if (base / name).is_symlink():
return True
return False
def _link_or_copy(src: Path, dst: Path) -> None:
"""Hard-link ``src`` to ``dst``, falling back to a content copy."""
try:
os.link(src, dst)
except OSError:
shutil.copy2(src, dst, follow_symlinks=True)
def _stage_dir(src: Path, dst: Path, root: Path, seen: frozenset[Path]) -> None:
dst.mkdir(parents=True, exist_ok=True)
for entry in os.scandir(src):
entry_path = Path(entry.path)
dest_path = dst / entry.name
if entry.is_symlink():
target = Path(os.path.realpath(entry_path))
if not _is_within(target, root):
logger.warning("staging: dropping out-of-tree symlink %s -> %s", entry_path, target)
continue
if not target.exists():
logger.warning("staging: dropping dangling symlink %s", entry_path)
continue
if target in seen:
logger.warning("staging: dropping cyclic symlink %s -> %s", entry_path, target)
continue
if target.is_dir():
_stage_dir(target, dest_path, root, seen | {target})
else:
_link_or_copy(target, dest_path)
elif entry.is_dir(follow_symlinks=False):
_stage_dir(entry_path, dest_path, root, seen)
elif entry.is_file(follow_symlinks=False):
_link_or_copy(entry_path, dest_path)
else:
# Sockets, FIFOs, devices — not part of a source tree; skip.
logger.debug("staging: skipping non-regular entry %s", entry_path)
def stage_symlink_safe_dir(src_root: Path) -> tuple[Path, Path | None]:
"""Return ``(upload_path, staged_temp)`` for uploading ``src_root``.
``upload_path`` is safe to hand to ``LocalDir``. When the tree contains no
symlinks it is ``src_root`` itself and ``staged_temp`` is ``None``.
Otherwise a symlink-safe copy is materialized in a temp directory and both
returned values point at it; the caller owns removing ``staged_temp`` once
the upload completes.
"""
root = src_root.resolve()
if not tree_has_symlink(root):
return root, None
staged = Path(tempfile.mkdtemp(prefix=_STAGING_PREFIX)).resolve()
try:
_stage_dir(root, staged, root, frozenset({root}))
except OSError:
shutil.rmtree(staged, ignore_errors=True)
raise
logger.info("staging: materialized symlink-safe copy of %s at %s", root, staged)
return staged, staged
+89 -47
View File
@@ -3,17 +3,21 @@
from __future__ import annotations
import logging
import shutil
import os
import sys
from pathlib import Path
from typing import Any
from typing import TYPE_CHECKING, Any
from agents.sandbox.entries import BaseEntry, LocalDir
from agents.sandbox.manifest import Environment, Manifest
from strix.config import load_settings
from strix.runtime.backends import get_backend
from strix.runtime.backends import backend_supports_bind_mounts, get_backend
from strix.runtime.caido_bootstrap import bootstrap_caido
from strix.runtime.local_dir_staging import stage_symlink_safe_dir
if TYPE_CHECKING:
from strix.runtime.status import StatusSink
logger = logging.getLogger(__name__)
@@ -28,43 +32,72 @@ _SESSION_CACHE: dict[str, dict[str, Any]] = {}
# Manifest root inside the container; entry keys hang off this path.
_WORKSPACE_ROOT = "/workspace"
_PROTECTED_METADATA_NAMES = (".git", ".agents", ".codex")
def build_session_entries(
local_sources: list[dict[str, Any]],
) -> tuple[dict[str | Path, BaseEntry], list[dict[str, Any]], list[Path]]:
"""Split local sources into copied manifest entries and host bind mounts.
Sources flagged ``mount`` are bind-mounted read-only at
``/workspace/<workspace_subdir>`` (not added to the manifest, so the SDK
does not stream them in file-by-file). Every other source becomes a
``LocalDir`` entry copied into the container as before. Trees containing
symlinks (which the SDK's ``LocalDir`` walker refuses outright) are first
staged into a symlink-safe temp copy; those temp dirs are returned so the
caller can remove them once the upload completes.
"""
entries: dict[str | Path, BaseEntry] = {}
def _host_identity_env() -> dict[str, str]:
if sys.platform != "linux":
return {}
return {"STRIX_HOST_UID": str(os.getuid()), "STRIX_HOST_GID": str(os.getgid())}
def build_bind_mounts(local_sources: list[dict[str, Any]]) -> list[dict[str, Any]]:
bind_mounts: list[dict[str, Any]] = []
staged_dirs: list[Path] = []
for src in local_sources:
ws_subdir = src.get("workspace_subdir") or ""
host_path = src.get("source_path") or ""
if not ws_subdir or not host_path:
continue
resolved = Path(host_path).expanduser().resolve()
if src.get("mount"):
bind_mounts.append(
{
"source": str(resolved),
"target": f"{_WORKSPACE_ROOT}/{ws_subdir}",
"read_only": True,
}
target = f"{_WORKSPACE_ROOT}/{ws_subdir}"
bind_mounts.append({"source": str(resolved), "target": target, "read_only": False})
if src.get("protect_metadata"):
bind_mounts.extend(_metadata_mounts(resolved, target))
return bind_mounts
def build_manifest_entries(local_sources: list[dict[str, Any]]) -> dict[str | Path, BaseEntry]:
entries: dict[str | Path, BaseEntry] = {}
for src in local_sources:
ws_subdir = src.get("workspace_subdir") or ""
host_path = src.get("source_path") or ""
if not ws_subdir or not host_path:
continue
entries[ws_subdir] = LocalDir(src=Path(host_path).expanduser().resolve())
return entries
def _metadata_mounts(tree: Path, target: str) -> list[dict[str, Any]]:
mounts: list[dict[str, Any]] = []
for name in _PROTECTED_METADATA_NAMES:
metadata = tree / name
if not metadata.is_dir() and not metadata.is_file():
continue
if not metadata.resolve().is_relative_to(tree):
continue
mounts.append({"source": str(metadata), "target": f"{target}/{name}", "read_only": True})
gitdir = _gitdir_from_pointer(metadata) if metadata.is_file() else None
if gitdir is not None and gitdir.exists() and gitdir.is_relative_to(tree):
relative = gitdir.relative_to(tree).as_posix()
mounts.append(
{"source": str(gitdir), "target": f"{target}/{relative}", "read_only": True}
)
else:
upload_path, staged = stage_symlink_safe_dir(resolved)
if staged is not None:
staged_dirs.append(staged)
entries[ws_subdir] = LocalDir(src=upload_path)
return entries, bind_mounts, staged_dirs
return mounts
def _gitdir_from_pointer(git_file: Path) -> Path | None:
try:
content = git_file.read_text(encoding="utf-8", errors="replace")
except OSError:
return None
for line in content.splitlines():
prefix, _, value = line.partition(":")
if prefix.strip() == "gitdir" and value.strip():
candidate = Path(value.strip()).expanduser()
if not candidate.is_absolute():
candidate = git_file.parent / candidate
return candidate.resolve()
return None
async def create_or_reuse(
@@ -72,19 +105,32 @@ async def create_or_reuse(
*,
image: str,
local_sources: list[dict[str, Any]],
status_sink: StatusSink | None = None,
) -> dict[str, Any]:
"""Return the existing session bundle for ``scan_id`` or create a new one.
Each ``local_sources`` entry exposes its host ``source_path`` at
``/workspace/<workspace_subdir>`` inside the container copied in, or
bind-mounted read-only when the entry is flagged ``mount``.
``/workspace/<workspace_subdir>`` inside the container.
"""
def report(phase: str) -> None:
if status_sink is not None:
status_sink(phase)
cached = _SESSION_CACHE.get(scan_id)
if cached is not None:
logger.info("Reusing existing sandbox session for scan %s", scan_id)
return cached
entries, bind_mounts, staged_dirs = build_session_entries(local_sources)
backend_name = load_settings().runtime.backend
backend = get_backend(backend_name)
if backend_supports_bind_mounts(backend_name):
bind_mounts = build_bind_mounts(local_sources)
entries: dict[str | Path, BaseEntry] = {}
else:
bind_mounts = []
entries = build_manifest_entries(local_sources)
# Caido runs as an in-container sidecar; HTTP(S) traffic from any
# process started via ``session.exec`` (the SDK's Shell tool, etc.)
@@ -98,6 +144,7 @@ async def create_or_reuse(
value={
"PYTHONUNBUFFERED": "1",
"HOST_GATEWAY": "host.docker.internal",
**_host_identity_env(),
"http_proxy": container_caido_url,
"https_proxy": container_caido_url,
"ALL_PROXY": container_caido_url,
@@ -106,26 +153,21 @@ async def create_or_reuse(
),
)
backend_name = load_settings().runtime.backend
backend = get_backend(backend_name)
logger.info(
"Creating sandbox session for scan %s (backend=%s, image=%s)",
scan_id,
backend_name,
image,
)
try:
client, session = await backend(
image=image,
manifest=manifest,
exposed_ports=(_CONTAINER_CAIDO_PORT,),
bind_mounts=bind_mounts,
)
finally:
for staged in staged_dirs:
shutil.rmtree(staged, ignore_errors=True)
report("Starting sandbox container")
client, session = await backend(
image=image,
manifest=manifest,
exposed_ports=(_CONTAINER_CAIDO_PORT,),
bind_mounts=bind_mounts,
)
report("Setting up the proxy")
caido_endpoint = await session.resolve_exposed_port(_CONTAINER_CAIDO_PORT)
scheme = "https" if caido_endpoint.tls else "http"
host_caido_url = f"{scheme}://{caido_endpoint.host}:{caido_endpoint.port}"
+8
View File
@@ -0,0 +1,8 @@
"""Startup phase reporting."""
from __future__ import annotations
from collections.abc import Callable
StatusSink = Callable[[str], None]
+5 -5
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import contextlib
import logging
import os
import sys
import warnings
from contextvars import ContextVar
from pathlib import Path # noqa: TC003 used at runtime by ``setup_scan_logging``
@@ -78,11 +79,10 @@ class _StdoutQuietFilter(logging.Filter):
def configure_dependency_logging() -> None:
"""Quiet dependency logging/warnings that obscure Strix scan logs."""
with contextlib.suppress(Exception):
import litellm
litellm_logging = litellm._logging
litellm_logging._disable_debugging() # type: ignore[no-untyped-call]
litellm = sys.modules.get("litellm")
if litellm is not None:
with contextlib.suppress(Exception):
litellm._logging._disable_debugging()
logging.getLogger("asyncio").setLevel(logging.CRITICAL)
logging.getLogger("asyncio").propagate = False
+3 -9
View File
@@ -1,9 +1,9 @@
import json
import logging
import urllib.request
from datetime import datetime
from typing import TYPE_CHECKING, Any
import requests
from strix.config import load_settings
from strix.telemetry._common import (
SESSION_ID,
@@ -37,13 +37,7 @@ def _send(event: str, properties: dict[str, Any]) -> bool:
"distinct_id": SESSION_ID,
"properties": properties,
}
req = urllib.request.Request( # noqa: S310
f"{_POSTHOG_HOST}/capture/",
data=json.dumps(payload).encode(),
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=10): # noqa: S310 # nosec B310
pass
requests.post(f"{_POSTHOG_HOST}/capture/", json=payload, timeout=10)
except Exception: # noqa: BLE001
logger.debug("posthog send failed for event %s", event, exc_info=True)
return False
+3 -4
View File
@@ -2,10 +2,11 @@ from __future__ import annotations
import logging
import urllib.parse
import urllib.request
from datetime import datetime
from typing import TYPE_CHECKING, Any
import requests
from strix.config import load_settings
from strix.telemetry._common import (
SESSION_ID,
@@ -42,9 +43,7 @@ def _send(event: str, properties: dict[str, Any]) -> bool:
url = f"{_SCARF_ENDPOINT}{path}"
if query:
url = f"{url}?{query}"
req = urllib.request.Request(url, method="POST") # noqa: S310
with urllib.request.urlopen(req, timeout=10): # noqa: S310 # nosec B310
pass
requests.post(url, timeout=10)
except Exception: # noqa: BLE001
logger.debug("scarf send failed for event %s", event, exc_info=True)
return False
+51 -22
View File
@@ -13,6 +13,7 @@ from typing import Any, Literal, get_args
from agents import RunContextWrapper, function_tool
from strix.core.agents import Status, coordinator_from_context
from strix.core.execution import notify_parent_on_terminal
from strix.skills import validate_requested_skills
@@ -218,25 +219,43 @@ def _session_items_payload(items: list[Any]) -> list[dict[str, Any]]:
return payload
@function_tool(timeout=601)
async def wait_for_message( # noqa: PLR0911
_WAIT_DEFAULT_TIMEOUT_S = 300
# Enforced by the SDK around the whole tool call, so it caps an oversized
# ``timeout_seconds`` the model asks for. One second of headroom lets the
# tool's own timeout fire first and return a clean result.
_WAIT_HARD_CEILING_S = _WAIT_DEFAULT_TIMEOUT_S + 1
@function_tool(timeout=_WAIT_HARD_CEILING_S)
async def wait_for_agents( # noqa: PLR0911
ctx: RunContextWrapper,
reason: str = "Waiting for messages from other agents",
timeout_seconds: int = 600,
timeout_seconds: int = _WAIT_DEFAULT_TIMEOUT_S,
) -> str:
"""Pause this agent until a message lands in its inbox (or timeout).
"""Pause until another AGENT messages you (or the timeout elapses).
Use when you have nothing useful to do until a child/peer responds
typically after spawning subagents and you want to wait for
their completion reports. The agent automatically resumes when any
message arrives, so pick a ``timeout_seconds`` proportional to the
work you're awaiting.
Use when you have nothing useful to do until a child or peer
responds typically after spawning subagents and you want their
completion reports. You resume the instant any message arrives, so
size ``timeout_seconds`` to the work you're awaiting.
**This tool is only for waiting on other agents.** Two things it is
NOT for:
- **Talking to the user.** Use ``respond_to_user``, which delivers
your message and hands control back in one call.
- **Waiting for a long-running command.** This tool does not watch
processes at all it sleeps until a *message* arrives, so it
burns the full timeout even if your command finished a second
later. Poll the process instead: ``exec_command`` returns a
session/process id, and ``write_stdin`` with ``chars=""`` returns
as soon as there is new output or the process exits.
**Critical caveats:**
- **Never** call this if you finished your own task and have **no**
child agents running that's a permanent stall. Call
``finish_scan`` (root) or ``agent_finish`` (subagent) instead.
- **Never** call this if you have no agents left to hear from
that just strands you until the timeout. Call ``finish_scan``
(root) or ``agent_finish`` (subagent) instead.
- If you're waiting on an agent that **isn't your child**, message
it first asking it to ping you when done otherwise it has no
reason to send to your inbox and you'll wait the full timeout.
@@ -247,7 +266,8 @@ async def wait_for_message( # noqa: PLR0911
reason: One-line note shown in graph snapshots while you're
waiting (helps a human or sibling agent debug who's stuck
on what).
timeout_seconds: Max seconds to wait (default 600). This is only
timeout_seconds: Max seconds to wait (default 300, and values above
that are cut short by a hard ceiling). This is only
a cap the tool returns the INSTANT a message arrives, so a
larger value never makes you wait longer when the reply does
come. Right-size it to what you're waiting on: a short wait
@@ -257,9 +277,7 @@ async def wait_for_message( # noqa: PLR0911
bites when the expected message never arrives so an oversized
timeout on a trivial wait just strands you idle until it
elapses. On timeout the tool returns and you decide whether to
keep working or wait again. (Applies to autonomous multi-agent
runs; in interactive/chat sessions the agent instead parks until
a message arrives and this cap is not enforced.)
keep working or wait again.
"""
inner = _ctx(ctx)
coordinator = coordinator_from_context(inner)
@@ -302,7 +320,7 @@ async def wait_for_message( # noqa: PLR0911
)
if interactive:
await coordinator.park_waiting(me)
await coordinator.park_waiting(me, wait_kind="agents")
return json.dumps(
{
"success": True,
@@ -314,7 +332,7 @@ async def wait_for_message( # noqa: PLR0911
default=str,
)
await coordinator.park_waiting(me)
await coordinator.park_waiting(me, wait_kind="agents")
try:
await asyncio.wait_for(coordinator.wait_for_message(me), timeout_seconds)
except TimeoutError:
@@ -373,7 +391,7 @@ async def create_agent(
Decompose complex pentests by handing focused subtasks to dedicated
children. The child runs asynchronously the parent continues
immediately and can ``wait_for_message`` later (or just keep
immediately and can ``wait_for_agents`` later (or just keep
working in parallel). When the child calls ``agent_finish``, its
completion report lands in the parent's inbox.
@@ -541,7 +559,7 @@ async def agent_finish(
)
parent_notified = False
if report_to_parent:
if report_to_parent and await coordinator.claim_parent_notice(me):
async with coordinator._lock:
agent_name = coordinator.names.get(me, me)
report = _render_completion_report(
@@ -565,6 +583,11 @@ async def agent_finish(
)
parent_notified = True
await coordinator.set_status(me, "completed")
if not parent_notified:
# Silence here would leave a parent waiting on a report that is never coming.
await notify_parent_on_terminal(coordinator, me, "completed")
logger.info(
"agent_finish: %s success=%s findings=%d parent_notified=%s",
me,
@@ -572,7 +595,6 @@ async def agent_finish(
len(findings or []),
parent_notified,
)
await coordinator.set_status(me, "completed")
return json.dumps(
{
@@ -663,9 +685,16 @@ async def stop_agent(
)
if cascade:
await coordinator.cancel_descendants_graceful(target_agent_id)
stopped = await coordinator.cancel_descendants_graceful(target_agent_id)
else:
await coordinator.request_stop(target_agent_id)
stopped = [target_agent_id]
# The stopper knows what it just did; anyone else waiting on those agents does not.
async with coordinator._lock:
orphaned = [aid for aid in stopped if coordinator.parent_of.get(aid) not in (None, me)]
for aid in orphaned:
await notify_parent_on_terminal(coordinator, aid, "stopped")
logger.info(
"stop_agent: target=%s cascade=%s reason=%r",
+2 -2
View File
@@ -101,7 +101,7 @@ async def finish_scan(
execution stops. There is no draft mode and no second chance: never
submit placeholder, provisional, or "checking if done" text in any
field, and never call ``finish_scan`` to poll whether subagents are
done (use ``view_agent_graph`` / ``wait_for_message`` for that).
done (use ``view_agent_graph`` / ``wait_for_agents`` for that).
Call it exactly ONCE, only when every field holds genuine, finished
assessment prose.
@@ -111,7 +111,7 @@ async def finish_scan(
summary. If ANY agent is in ``running`` / ``waiting`` state,
you MUST NOT call ``finish_scan`` yet
wrap them up first via ``send_message_to_agent`` (ask them to
finish), ``wait_for_message`` (block until their report
finish), ``wait_for_agents`` (block until their report
arrives), or ``stop_agent`` (graceful cancel). Only ``completed``
/ ``crashed`` / ``stopped`` agents are safe to leave behind.
Calling ``finish_scan`` while children are alive orphans their
+6
View File
@@ -0,0 +1,6 @@
"""User-facing reply tool for interactive sessions."""
from strix.tools.respond.tool import respond_to_user
__all__ = ["respond_to_user"]
+110
View File
@@ -0,0 +1,110 @@
"""``respond_to_user`` — deliver a reply and hand control back to the user."""
from __future__ import annotations
import json
from typing import Any
from agents import RunContextWrapper, function_tool
from strix.core.agents import coordinator_from_context
def _ctx(ctx: RunContextWrapper) -> dict[str, Any]:
return ctx.context if isinstance(ctx.context, dict) else {}
@function_tool
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
yielding are the same call on purpose: there is no way to answer and
then forget to stop, and no way to stop without having answered.
Call it when you have something for the user and nothing to do until
they reply you answered their question, you need a decision or a
credential only they can give, or you finished a chunk of work and
want direction. You resume exactly where you left off when they
reply, with everything you have done so far intact.
Do NOT call it to narrate progress or to think out loud. Plain text
is still shown to the user as you work, so say whatever you like
mid-task without stopping; ``respond_to_user`` is specifically the
act of *waiting* for them. Every call costs the user their attention.
Not for these:
- **Waiting on another agent** (a child's report, a peer's reply)
use ``wait_for_agents``.
- **Ending the engagement** use ``finish_scan`` (root) or
``agent_finish`` (subagent). Those are terminal; this is a pause.
Args:
message: What to say to the user. Self-contained: they may not
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.
"""
inner = _ctx(ctx)
coordinator = coordinator_from_context(inner)
me = inner.get("agent_id")
interactive = bool(inner.get("interactive", False))
if coordinator is None or me is None:
return json.dumps(
{"success": False, "error": "Agent coordinator or agent_id missing in context"},
ensure_ascii=False,
default=str,
)
if not interactive:
return json.dumps(
{
"success": False,
"error": (
"No user is attached to an autonomous run. Keep working, and call "
"finish_scan (root) or agent_finish (subagent) when the task is done."
),
},
ensure_ascii=False,
default=str,
)
async with coordinator._lock:
stopped = coordinator.statuses.get(me) == "stopped"
if stopped:
return json.dumps(
{"success": True, "wait_outcome": "stopped", "message": message},
ensure_ascii=False,
default=str,
)
# A message that arrived while this turn was running is the user already
# talking: take it now instead of parking for one they have sent.
pending, _ = await coordinator.consume_pending(me)
if pending > 0:
await coordinator.mark_running(me)
return json.dumps(
{
"success": True,
"wait_outcome": "message_arrived",
"pending_messages": pending,
"message": message,
"note": "Your reply was delivered; the user had already sent a new message.",
},
ensure_ascii=False,
default=str,
)
await coordinator.park_waiting(me, wait_kind="user")
return json.dumps(
{
"success": True,
"wait_outcome": "waiting",
"message": message,
"note": "Reply delivered; parked until the user responds.",
},
ensure_ascii=False,
default=str,
)
+124
View File
@@ -0,0 +1,124 @@
"""Tests for tool-argument shape coercion in the agent factory."""
from __future__ import annotations
import json
from typing import Any, cast
import pytest
from agents.tool import FunctionTool
from strix.agents import factory
def _capturing_tool(captured: dict[str, str], schema: dict[str, Any]) -> FunctionTool:
async def invoke(_ctx: Any, raw_input: str) -> str:
captured["raw_input"] = raw_input
return "ok"
return FunctionTool(
name="probe",
description="test tool",
params_json_schema={"type": "object", "properties": schema},
on_invoke_tool=invoke,
)
async def _roundtrip(schema: dict[str, Any], payload: dict[str, Any]) -> dict[str, Any]:
captured: dict[str, str] = {}
wrapped = factory._with_coerced_arguments(_capturing_tool(captured, schema))
assert await wrapped.on_invoke_tool(cast("Any", None), json.dumps(payload)) == "ok"
return cast("dict[str, Any]", json.loads(captured["raw_input"]))
_STRING = {"todos": {"type": "string"}}
_ARRAY = {"tags": {"type": "array", "items": {"type": "string"}}}
_NULLABLE_ARRAY = {
"tags": {"anyOf": [{"type": "array", "items": {"type": "string"}}, {"type": "null"}]}
}
_OBJECT = {"modifications": {"type": "object"}}
@pytest.mark.asyncio
async def test_structured_value_is_encoded_for_a_string_parameter() -> None:
parsed = await _roundtrip(_STRING, {"todos": [{"title": "Phase 1: recon"}]})
assert parsed["todos"] == '[{"title": "Phase 1: recon"}]'
@pytest.mark.asyncio
async def test_string_parameter_keeps_an_already_encoded_value() -> None:
parsed = await _roundtrip(_STRING, {"todos": '[{"title": "a"}]'})
assert parsed["todos"] == '[{"title": "a"}]'
@pytest.mark.asyncio
@pytest.mark.parametrize("schema", [_ARRAY, _NULLABLE_ARRAY])
async def test_encoded_list_is_decoded_for_an_array_parameter(schema: dict[str, Any]) -> None:
parsed = await _roundtrip(schema, {"tags": '["auth", "idor"]'})
assert parsed["tags"] == ["auth", "idor"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"value",
[
"auth, idor",
"auth\nidor",
"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:
parsed = await _roundtrip(_ARRAY, {"tags": value})
assert parsed["tags"] == value
@pytest.mark.asyncio
async def test_encoded_mapping_is_decoded_for_an_object_parameter() -> None:
parsed = await _roundtrip(_OBJECT, {"modifications": '{"method": "POST"}'})
assert parsed["modifications"] == {"method": "POST"}
@pytest.mark.asyncio
async def test_a_decoded_container_of_the_wrong_kind_is_not_substituted() -> None:
parsed = await _roundtrip(_OBJECT, {"modifications": '["POST"]'})
assert parsed["modifications"] == '["POST"]'
@pytest.mark.asyncio
async def test_values_matching_the_schema_are_left_alone() -> None:
parsed = await _roundtrip({**_ARRAY, **_OBJECT}, {"tags": ["auth"], "modifications": {"a": 1}})
assert parsed == {"tags": ["auth"], "modifications": {"a": 1}}
@pytest.mark.asyncio
async def test_unknown_and_null_arguments_are_untouched() -> None:
parsed = await _roundtrip(_NULLABLE_ARRAY, {"tags": None, "other": ["x"]})
assert parsed == {"tags": None, "other": ["x"]}
@pytest.mark.asyncio
async def test_non_object_payloads_pass_through_unchanged() -> None:
captured: dict[str, str] = {}
wrapped = factory._with_coerced_arguments(_capturing_tool(captured, _ARRAY))
assert await wrapped.on_invoke_tool(cast("Any", None), "not json") == "ok"
assert captured["raw_input"] == "not json"
@pytest.mark.asyncio
async def test_coercion_is_applied_once_per_tool() -> None:
captured: dict[str, str] = {}
tool = factory._with_coerced_arguments(_capturing_tool(captured, _ARRAY))
assert factory._with_coerced_arguments(tool) is tool
+27 -1
View File
@@ -2,18 +2,29 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any
import pytest
from agents.tool import FunctionTool
from strix.agents import factory
if TYPE_CHECKING:
from agents.tool_context import ToolContext
def _tool(name: str) -> FunctionTool:
# A per-tool closure keeps two same-named tools unequal, which is what the
# duplicate-name tests exercise.
async def invoke(_ctx: ToolContext[Any], _input: str) -> str:
return "ok"
return FunctionTool(
name=name,
description="test tool",
params_json_schema={"type": "object", "properties": {}, "additionalProperties": False},
on_invoke_tool=lambda _ctx, _inp: "ok",
on_invoke_tool=invoke,
)
@@ -86,3 +97,18 @@ def test_no_override_renders_builtin_prompt() -> None:
assert isinstance(agent.instructions, str)
assert agent.instructions != ""
def test_respond_to_user_is_interactive_only() -> None:
"""Yielding to the user is meaningless when no user is attached."""
interactive = factory.build_strix_agent(is_root=True, interactive=True)
autonomous = factory.build_strix_agent(is_root=True, interactive=False)
assert "respond_to_user" in [t.name for t in interactive.tools]
assert "respond_to_user" not in [t.name for t in autonomous.tools]
def test_wait_for_agents_is_available_in_both_modes() -> None:
for interactive in (True, False):
agent = factory.build_strix_agent(is_root=True, interactive=interactive)
assert "wait_for_agents" in [t.name for t in agent.tools]
+2 -7
View File
@@ -30,9 +30,7 @@ def test_parse_arguments_accepts_target_list_file(
) -> None:
target_list = tmp_path / "targets.txt"
target_list.write_text(
"https://test1.com/\n"
"\n"
"http://test2.com:5789/\n",
"https://test1.com/\n\nhttp://test2.com:5789/\n",
encoding="utf-8",
)
_stub_settings(monkeypatch)
@@ -84,7 +82,4 @@ def test_parse_arguments_rejects_resume_with_target_list(
with pytest.raises(SystemExit):
cli_main.parse_arguments()
assert (
"Cannot combine --resume with --target/--target-list/--mount"
in capsys.readouterr().err
)
assert "Cannot combine --resume with --target/--target-list" in capsys.readouterr().err
+14
View File
@@ -7,8 +7,10 @@ import hashlib
import json
import time
from typing import TYPE_CHECKING, Any
from unittest import mock
import pytest
import requests
from strix.config import codex
@@ -52,6 +54,18 @@ def test_authorize_url_carries_pkce_and_client() -> None:
assert "state=st8" in url
def test_post_form_returns_parsed_body() -> None:
resp = mock.MagicMock()
resp.status_code = 200
resp.content = b'{"access_token": "tok"}'
with mock.patch.object(requests, "post", return_value=resp) as post:
data = codex._post_form({"grant_type": "refresh_token"})
assert data == {"access_token": "tok"}
assert post.call_args.kwargs["timeout"] == codex._TOKEN_TIMEOUT
@pytest.mark.parametrize(
("value", "expected"),
[
+69 -42
View File
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any
import pytest
from litellm.exceptions import BadRequestError, ContextWindowExceededError, RateLimitError
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from strix.config import ContextSettings
from strix.llm import compaction
@@ -146,17 +147,35 @@ def _patch_budget(monkeypatch: pytest.MonkeyPatch, *, keep_tokens: int, window:
context.auto_compact = True
settings = SimpleNamespace(
context=context,
llm=SimpleNamespace(api_key=None, api_base=None, timeout=1),
llm=SimpleNamespace(api_key=None, api_base=None, timeout=1, extra_headers=None),
)
monkeypatch.setattr(compaction, "load_settings", lambda: settings)
def _patch_summary(monkeypatch: pytest.MonkeyPatch, text: str) -> None:
async def fake_acompletion(**_kwargs: Any) -> Any:
message = SimpleNamespace(content=text)
return SimpleNamespace(choices=[SimpleNamespace(message=message)])
def _model_response(text: str) -> Any:
chunk = ResponseOutputText(annotations=[], text=text, type="output_text")
message = ResponseOutputMessage(
id="msg", content=[chunk], role="assistant", status="completed", type="message"
)
return SimpleNamespace(output=[message])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
def _patch_summary(
monkeypatch: pytest.MonkeyPatch, text: str, captured: dict[str, Any] | None = None
) -> None:
class FakeModel:
async def get_response(self, **kwargs: Any) -> Any:
if captured is not None:
captured.update(kwargs)
return _model_response(text)
class FakeProvider:
def get_model(self, model_name: str | None) -> Any:
if captured is not None:
captured["model"] = model_name
return FakeModel()
monkeypatch.setattr(compaction, "StrixProvider", FakeProvider)
@pytest.mark.asyncio
@@ -189,19 +208,38 @@ async def test_maybe_compact_rewrites_and_keeps_pairs(monkeypatch: pytest.Monkey
async def test_maybe_compact_updates_previous_summary(monkeypatch: pytest.MonkeyPatch) -> None:
# Window large enough to leave real room for the summary instructions.
_patch_budget(monkeypatch, keep_tokens=30, window=4_000)
captured: dict[str, str] = {}
async def fake_acompletion(**kwargs: Any) -> Any:
captured["prompt"] = kwargs["messages"][0]["content"]
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="NEW"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "NEW", captured)
prior = compaction._checkpoint_item("OLD SUMMARY TEXT")
session = FakeSession([prior, *_turns(12)])
assert await compaction.maybe_compact(session, model="m", force=True) is True
assert "OLD SUMMARY TEXT" in captured["prompt"]
assert "OLD SUMMARY TEXT" in captured["input"]
@pytest.mark.asyncio
async def test_summarize_routes_through_provider_with_settings(
monkeypatch: pytest.MonkeyPatch,
) -> None:
_patch_budget(monkeypatch, keep_tokens=30, window=4_000)
monkeypatch.setattr(
compaction,
"load_settings",
lambda: SimpleNamespace(
llm=SimpleNamespace(
api_key=None, api_base=None, timeout=1, extra_headers={"X-Feature-Key": "svc"}
)
),
)
captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
assert await compaction._summarize("litellm/openai/some-model", "p", 64) == "S"
assert captured["model"] == "litellm/openai/some-model"
settings = captured["model_settings"]
assert settings.extra_headers == {"X-Feature-Key": "svc"}
assert settings.max_tokens == 64
def test_fit_to_tokens_truncates_oversized_text(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -233,19 +271,14 @@ def test_summary_output_tokens_capped_at_model_limit(monkeypatch: pytest.MonkeyP
async def test_maybe_compact_bounds_summary_prompt(monkeypatch: pytest.MonkeyPatch) -> None:
# A tiny window with a huge head must not send an oversized summary request.
_patch_budget(monkeypatch, keep_tokens=30, window=4_000)
captured: dict[str, str] = {}
async def fake_acompletion(**kwargs: Any) -> Any:
captured["prompt"] = kwargs["messages"][0]["content"]
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
big_turns = [{"role": "user", "content": "y" * 2_000} for _ in range(50)]
session = FakeSession(big_turns)
assert await compaction.maybe_compact(session, model="m") is True
# count_tokens==len(chars); prompt must fit the model window.
assert len(captured["prompt"]) <= 4_000
assert len(captured["input"]) <= 4_000
@pytest.mark.asyncio
@@ -256,27 +289,27 @@ async def test_summary_request_fits_when_room_is_below_old_floor(
instructions = len(compaction._SUMMARY_INSTRUCTIONS)
window = instructions + 64 + 256 + 300 # summary_max(64)+slack(256)+room(300)
_patch_budget(monkeypatch, keep_tokens=30, window=window)
captured: dict[str, str] = {}
async def fake_acompletion(**kwargs: Any) -> Any:
captured["prompt"] = kwargs["messages"][0]["content"]
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
session = FakeSession([{"role": "user", "content": "y" * 5_000} for _ in range(20)])
assert await compaction.maybe_compact(session, model="m") is True
assert len(captured["prompt"]) <= window
assert len(captured["input"]) <= window
@pytest.mark.asyncio
async def test_maybe_compact_skips_when_summary_fails(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_budget(monkeypatch, keep_tokens=30, window=4_000)
async def fake_acompletion(**_kwargs: Any) -> Any:
raise RuntimeError("boom")
class BoomModel:
async def get_response(self, **_kwargs: Any) -> Any:
raise RuntimeError("boom")
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
class BoomProvider:
def get_model(self, _model_name: str | None) -> Any:
return BoomModel()
monkeypatch.setattr(compaction, "StrixProvider", BoomProvider)
session = FakeSession(_turns(12))
before = await session.get_items()
@@ -290,17 +323,11 @@ async def test_maybe_compact_skips_when_no_room_to_summarise(
) -> None:
# No room for any head -> no (doomed) summary is attempted.
_patch_budget(monkeypatch, keep_tokens=30, window=200)
called = False
async def fake_acompletion(**_kwargs: Any) -> Any:
nonlocal called
called = True
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
session = FakeSession(_turns(12))
before = await session.get_items()
assert await compaction.maybe_compact(session, model="m", force=True) is False
assert called is False
assert not captured
assert await session.get_items() == before
-1
View File
@@ -33,7 +33,6 @@ _LLM_ENV_KEYS = [
# RuntimeSettings
"STRIX_IMAGE",
"STRIX_RUNTIME_BACKEND",
"STRIX_MAX_LOCAL_COPY_MB",
# TelemetrySettings
"STRIX_TELEMETRY",
]
+18
View File
@@ -21,6 +21,24 @@ def test_context_window_strips_provider_prefix() -> None:
assert context_budget.context_window("openai/gpt-4o") == 128_000
def test_context_window_chatgpt_prefix_skips_provider_auth(
monkeypatch: pytest.MonkeyPatch,
) -> None:
context_budget._model_info.cache_clear()
calls: list[str] = []
def _model_info(model: str) -> dict[str, int]:
calls.append(model)
return {"max_input_tokens": 1_050_000, "max_output_tokens": 128_000}
monkeypatch.setattr("strix.llm.context_budget.litellm.get_model_info", _model_info)
try:
assert context_budget.context_window("chatgpt/gpt-5.6-luna") == 1_050_000
assert calls == ["gpt-5.6-luna"]
finally:
context_budget._model_info.cache_clear()
def test_context_window_unmapped_uses_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
context_budget._model_info.cache_clear()
+105 -2
View File
@@ -7,9 +7,25 @@ from unittest.mock import MagicMock, patch
import litellm
import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
from strix.config.models import _configure_litellm_compatibility
from strix.report.state import litellm_cost_callback
from strix.config.models import (
_configure_litellm_compatibility,
_install_openrouter_stream_cost_capture,
)
from strix.report.state import (
ReportState,
litellm_cost_callback,
openrouter_stream_cost,
set_global_report_state,
streamed_openrouter_costs,
)
@pytest.fixture(autouse=True)
def _clear_streamed_costs() -> None:
streamed_openrouter_costs.clear()
def test_streaming_logging_stays_enabled_for_cost_callback() -> None:
@@ -151,3 +167,90 @@ def test_cost_callback_records_nothing_when_no_cost_available() -> None:
litellm_cost_callback({"response_cost": None, "model": "x/y"}, response)
report_state.record_observed_llm_cost.assert_not_called()
def test_openrouter_stream_cost_extracts_plain_and_byok_totals() -> None:
assert openrouter_stream_cost({"cost": 0.003168}) == pytest.approx(0.003168)
assert openrouter_stream_cost(
{"cost": 0.01, "is_byok": True, "cost_details": {"upstream_inference_cost": 0.2}}
) == pytest.approx(0.21)
# Upstream cost is only added for BYOK responses.
assert openrouter_stream_cost(
{"cost": 0.05, "is_byok": False, "cost_details": {"upstream_inference_cost": 0.04}}
) == pytest.approx(0.05)
assert openrouter_stream_cost({"prompt_tokens": 10}) is None
assert openrouter_stream_cost(None) is None
def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> None:
report_state = MagicMock()
streamed_openrouter_costs.remember("gen-abc", {"cost": 0.42})
# LiteLLM strips cost from the rebuilt streamed usage; only the id survives.
response = SimpleNamespace(id="gen-abc", usage=SimpleNamespace(cost=None), _hidden_params={})
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=ValueError("unknown model")),
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.42)
# The entry is consumed so a later response cannot double-count it.
assert streamed_openrouter_costs.take(response) is None
def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None:
report_state = MagicMock()
streamed_openrouter_costs.remember("gen-xyz", {"cost": 0.9})
response = SimpleNamespace(
id="gen-xyz",
usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15),
_hidden_params={},
)
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", return_value=0.1) as estimate,
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.9)
estimate.assert_not_called()
def test_streamed_openrouter_costs_ignores_entries_without_cost() -> None:
streamed_openrouter_costs.remember("gen-none", {"prompt_tokens": 10})
streamed_openrouter_costs.remember("", {"cost": 0.5})
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-none")) is None
def test_streamed_openrouter_costs_cleared_on_new_run() -> None:
streamed_openrouter_costs.remember("gen-stale", {"cost": 0.7})
set_global_report_state(ReportState.__new__(ReportState))
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stale")) is None
def test_openrouter_stream_handler_records_cost() -> None:
_install_openrouter_stream_cost_capture()
# Resolve the config the way LiteLLM does in production so we prove the
# override is actually reachable through provider resolution, not just as a
# directly-constructed class.
config = ProviderConfigManager.get_provider_chat_config(
model="moonshotai/kimi-k3", provider=LlmProviders.OPENROUTER
)
assert config is not None
assert type(config).__name__ == "_StrixOpenrouterConfig"
handler = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True)
chunk = {
"id": "gen-stream",
"created": 1,
"model": "moonshotai/kimi-k3",
"choices": [{"index": 0, "delta": {"content": None}}],
"usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055},
}
handler.chunk_parser(chunk)
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx(
0.0035055
)
+32
View File
@@ -44,6 +44,38 @@ def test_dedupe_endpoint_sent_per_call() -> None:
assert (settings.extra_args or {})["api_key"] == "dedupe-key"
def test_dedicated_dedupe_model_uses_own_headers_not_main() -> None:
dedupe = DedupeSettings(
STRIX_DEDUPE_MODEL="deepseek/cheap",
DEDUPE_LLM_EXTRA_HEADERS={"X-Dedupe": "yes"},
)
settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300)
assert settings.extra_headers == {"X-Dedupe": "yes"}
def test_dedicated_dedupe_model_gets_no_main_headers_by_default(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Main": "secret"}))
loader._cached = None
try:
dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap")
settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300)
assert settings.extra_headers is None
finally:
loader._cached = None
def test_fallback_dedupe_inherits_main_headers(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Main": "svc"}))
loader._cached = None
try:
settings = _dedupe_model_settings(DedupeSettings(), "openai/main-model", 300)
assert settings.extra_headers == {"X-Main": "svc"}
finally:
loader._cached = None
def test_dedupe_defaults_are_empty() -> None:
settings = DedupeSettings()
assert settings.model is None
+326
View File
@@ -0,0 +1,326 @@
"""Tests for LLM_DISABLE_STREAMING: serve the streamed run loop without SSE.
A gateway that rejects ``stream:true`` (or delivers SSE unreliably) breaks the
SDK run loop, which only issues streamed requests. ``_NonStreamingModel`` wraps
the resolved model so each turn makes one non-streaming ``get_response`` and
replays the completed result as a single terminal stream event. A local server
that rejects streamed requests but answers non-streamed ones including a
structured tool call proves the wrapper works where the stock model fails.
"""
from __future__ import annotations
import json
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents import Agent, Runner, function_tool
from agents.model_settings import ModelSettings
from agents.models.interface import Model, ModelProvider, ModelTracing
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from agents.run import RunConfig
from openai import AsyncOpenAI, BadRequestError
from openai.types.responses import (
ResponseCompletedEvent,
ResponseFunctionToolCall,
ResponseOutputMessage,
ResponseOutputText,
)
from strix.config import codex, loader
from strix.config.loader import load_settings
from strix.config.models import StrixProvider, _NonStreamingModel
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterator
def _tool_call_completion() -> dict[str, Any]:
return {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "do_thing", "arguments": '{"n": 1}'},
}
],
},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
}
def _text_completion() -> dict[str, Any]:
return {
"id": "chatcmpl-2",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "stop",
"message": {"role": "assistant", "content": "hello from gateway"},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
}
_CAPTURED: dict[str, Any] = {}
_PAYLOAD: dict[str, dict[str, Any]] = {"value": _tool_call_completion()}
class _Handler(BaseHTTPRequestHandler):
"""A gateway that only speaks non-streaming Chat Completions."""
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
body = json.loads(self.rfile.read(length) or b"{}")
_CAPTURED.clear()
_CAPTURED.update(body)
if body.get("stream"):
payload = json.dumps(
{"error": {"message": "streaming is not supported by this endpoint"}}
).encode()
self.send_response(400)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
return
payload = json.dumps(_PAYLOAD["value"]).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
@pytest.fixture
def gateway_url() -> Iterator[str]:
_PAYLOAD["value"] = _tool_call_completion()
server = HTTPServer(("127.0.0.1", 0), _Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
server.shutdown()
server.server_close()
def _model(base_url: str) -> OpenAIChatCompletionsModel:
client = AsyncOpenAI(api_key="tok", base_url=base_url)
return OpenAIChatCompletionsModel(model="gw-model", openai_client=client)
def _call_kwargs() -> dict[str, Any]:
return {
"system_instructions": "s",
"input": "hi",
"model_settings": ModelSettings(),
"tools": [],
"output_schema": None,
"handoffs": [],
"tracing": ModelTracing.DISABLED,
"previous_response_id": None,
"conversation_id": None,
"prompt": None,
}
async def _drain(gen: AsyncIterator[Any]) -> list[Any]:
return [event async for event in gen]
@pytest.mark.asyncio
async def test_stock_model_streaming_fails_on_non_streaming_gateway(gateway_url: str) -> None:
# The stock model issues stream:true and the gateway rejects it.
model = _model(gateway_url)
with pytest.raises(BadRequestError, match="streaming is not supported"):
await _drain(model.stream_response(**_call_kwargs()))
assert _CAPTURED["stream"] is True
@pytest.mark.asyncio
async def test_wrapper_streams_tool_call_without_streaming_request(gateway_url: str) -> None:
# The wrapper turns the streamed run-loop call into one non-streaming
# request and replays the completed result as a terminal stream event.
model = _NonStreamingModel(_model(gateway_url))
events = await _drain(model.stream_response(**_call_kwargs()))
assert _CAPTURED.get("stream") is not True
assert len(events) == 1
completed = events[0]
assert isinstance(completed, ResponseCompletedEvent)
tool_call = completed.response.output[0]
assert isinstance(tool_call, ResponseFunctionToolCall)
assert tool_call.name == "do_thing"
assert json.loads(tool_call.arguments) == {"n": 1}
assert completed.response.usage is not None
assert completed.response.usage.total_tokens == 7
@pytest.mark.asyncio
async def test_wrapper_streams_plain_text(gateway_url: str) -> None:
_PAYLOAD["value"] = _text_completion()
model = _NonStreamingModel(_model(gateway_url))
events = await _drain(model.stream_response(**_call_kwargs()))
assert _CAPTURED.get("stream") is not True
message = events[0].response.output[0]
assert isinstance(message, ResponseOutputMessage)
text = message.content[0]
assert isinstance(text, ResponseOutputText)
assert text.text == "hello from gateway"
@pytest.mark.asyncio
async def test_wrapper_get_response_stays_non_streaming(gateway_url: str) -> None:
# The non-streaming path is a plain pass-through to the inner model.
model = _NonStreamingModel(_model(gateway_url))
response = await model.get_response(**_call_kwargs())
assert _CAPTURED.get("stream") is not True
tool_call = response.output[0]
assert isinstance(tool_call, ResponseFunctionToolCall)
assert tool_call.name == "do_thing"
_TURN_STREAM_FLAGS: list[bool] = []
class _MultiTurnHandler(BaseHTTPRequestHandler):
"""Non-streaming gateway: a tool call on turn 1, a final answer on turn 2."""
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
body = json.loads(self.rfile.read(length) or b"{}")
_TURN_STREAM_FLAGS.append(bool(body.get("stream")))
completion = _tool_call_completion() if len(_TURN_STREAM_FLAGS) == 1 else _text_completion()
if len(_TURN_STREAM_FLAGS) > 1:
completion["choices"][0]["message"]["content"] = "all done"
payload = json.dumps(completion).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
@pytest.fixture
def multiturn_url() -> Iterator[str]:
_TURN_STREAM_FLAGS.clear()
server = HTTPServer(("127.0.0.1", 0), _MultiTurnHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
server.shutdown()
server.server_close()
@pytest.mark.asyncio
async def test_run_loop_executes_tool_and_completes_without_streaming(multiturn_url: str) -> None:
# The whole streamed agent loop runs against a non-streaming gateway: the
# synthetic terminal event feeds the runner, which executes the tool and
# continues the turn until a final answer.
calls: list[int] = []
@function_tool
def do_thing(n: int) -> str:
calls.append(n)
return f"did {n}"
class _Provider(ModelProvider):
def get_model(self, model_name: str | None) -> Model: # noqa: ARG002
return _NonStreamingModel(_model(multiturn_url))
agent = Agent(name="t", instructions="use the tool", tools=[do_thing], model="gw-model")
result = Runner.run_streamed(
agent, input="please", run_config=RunConfig(model_provider=_Provider())
)
async for _ in result.stream_events():
pass
assert calls == [1] # tool executed with the streamed tool-call args
assert result.final_output == "all done"
assert len(_TURN_STREAM_FLAGS) == 2 # two turns, both...
assert not any(_TURN_STREAM_FLAGS) # ...issued as non-streaming requests
class _DummyModel(Model):
async def get_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
def stream_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
@pytest.fixture
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING"):
monkeypatch.delenv(key, raising=False)
monkeypatch.setattr(loader, "_cached", None)
monkeypatch.setattr(loader, "_override", None)
yield
def test_get_model_wraps_when_disabled(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
inner = _DummyModel()
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: inner)
monkeypatch.setenv("LLM_DISABLE_STREAMING", "true")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _NonStreamingModel)
def test_get_model_unwrapped_by_default(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
inner = _DummyModel()
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: inner)
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert model is inner
def test_get_model_does_not_wrap_subscription_model(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
# Subscription (ChatGPT) models are always streamed and must not be wrapped.
monkeypatch.setattr(codex, "subscription_model", lambda *_: "gpt-5.5")
monkeypatch.setattr(codex, "get_subscription_client", lambda: AsyncOpenAI(api_key="x"))
monkeypatch.setenv("LLM_DISABLE_STREAMING", "true")
load_settings()
model = StrixProvider().get_model("gpt-5.5")
assert not isinstance(model, _NonStreamingModel)
+25 -10
View File
@@ -7,7 +7,7 @@ from unittest.mock import MagicMock, patch
import pytest
from strix.core import execution
from strix.core.agents import AgentCoordinator
from strix.core.agents import AgentCoordinator, WaitKind
from strix.core.execution import _start_child_runner, run_agent_loop
from strix.core.hooks import BudgetExceededError, ReportUsageHooks
from strix.core.sessions import open_agent_session
@@ -42,23 +42,32 @@ class _FakeStream:
hooks: ReportUsageHooks,
context: dict[str, Any],
agent: Any,
coordinator: AgentCoordinator,
) -> None:
self._ledger = ledger
self._hooks = hooks
self._context = context
self._agent = agent
self._coordinator = coordinator
self.run_loop_exception: BaseException | None = None
self.final_output = None
async def stream_events(self) -> AsyncIterator[Any]:
agent_id = str(self._context.get("agent_id"))
self._ledger.cost += COST_PER_CALL
self._ledger.calls.append(str(self._context.get("agent_id")))
self._ledger.calls.append(agent_id)
ctx_wrapper = MagicMock()
ctx_wrapper.context = self._context
try:
await self._hooks.on_llm_end(ctx_wrapper, self._agent, MagicMock())
except Exception as exc: # noqa: BLE001
self.run_loop_exception = exc
# Stand in for the explicit yield tool a real turn ends with. Without it
# every turn looks like a forgotten tool call and burns the recovery
# budget, which is a different scenario from the one under test here.
if self._coordinator.statuses.get(agent_id) == "running":
wait_kind: WaitKind = "user" if self._context.get("parent_id") is None else "agents"
await self._coordinator.park_waiting(agent_id, wait_kind=wait_kind)
items: tuple[Any, ...] = ()
for item in items:
yield item
@@ -67,7 +76,7 @@ class _FakeStream:
return
def _fake_runner(ledger: _FakeLedger) -> Any:
def _fake_runner(ledger: _FakeLedger, coordinator: AgentCoordinator) -> Any:
class _FakeRunner:
@staticmethod
def run_streamed(
@@ -80,7 +89,13 @@ def _fake_runner(ledger: _FakeLedger) -> Any:
session: Any, # noqa: ARG004
hooks: ReportUsageHooks,
) -> _FakeStream:
return _FakeStream(ledger=ledger, hooks=hooks, context=context, agent=agent)
return _FakeStream(
ledger=ledger,
hooks=hooks,
context=context,
agent=agent,
coordinator=coordinator,
)
return _FakeRunner
@@ -103,10 +118,10 @@ async def test_full_budget_lifecycle_reserve_then_cap( # noqa: PLR0915
) -> None:
ledger = _FakeLedger()
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
coordinator = AgentCoordinator()
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger, coordinator))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
db_path = tmp_path / "agents.sqlite"
sessions: list[Any] = []
run_config = MagicMock()
@@ -218,7 +233,6 @@ async def test_respawned_children_after_reserve_never_spend(
ledger = _FakeLedger()
ledger.cost = 9.5
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
@@ -230,6 +244,7 @@ async def test_respawned_children_after_reserve_never_spend(
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.reserve_stopped is True
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger, restored))
sessions: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
@@ -264,7 +279,6 @@ async def test_resumed_parked_root_after_reserve_is_renotified_and_finalizes(
ledger = _FakeLedger()
ledger.cost = 9.0
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
@@ -276,6 +290,7 @@ async def test_resumed_parked_root_after_reserve_is_renotified_and_finalizes(
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.reserve_stopped is True
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger, restored))
root_session = open_agent_session("root", tmp_path / "agents.sqlite")
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
@@ -312,10 +327,10 @@ async def test_interactive_budget_pause_then_user_message_extends_and_resumes(
ledger = _FakeLedger()
ledger.cost = 9.0
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET, interactive=True)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
coordinator = AgentCoordinator()
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger, coordinator))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
coordinator.set_budget_extender(hooks.extend_budget)
await coordinator.register("root", "strix", parent_id=None)
root_session = open_agent_session("root", tmp_path / "agents.sqlite")
+738 -8
View File
@@ -5,17 +5,54 @@ from __future__ import annotations
import asyncio
import contextlib
import json
from typing import Any
from typing import Any, cast
from unittest.mock import MagicMock
import pytest
from agents.exceptions import MaxTurnsExceeded
from agents.items import MessageOutputItem
from agents.memory import SQLiteSession
from agents.tool_context import ToolContext
from openai.types.responses import ResponseOutputMessage, ResponseOutputRefusal
from strix.core import execution
from strix.core.agents import AgentCoordinator
from strix.core.execution import _notify_parent_on_terminal, _notify_root_on_budget_reserve
from strix.core.execution import (
_notify_root_on_budget_reserve,
notify_parent_on_terminal,
)
from strix.core.sessions import seed_initial_input
from strix.tools.agents_graph.tools import agent_finish, stop_agent
from strix.tools.finish.tool import finish_scan
_NO_STREAM_EVENTS: list[Any] = []
class _StructuredRefusalStream:
def __init__(self, refusal: str) -> None:
self.run_loop_exception: BaseException | None = None
self.new_items = [
MessageOutputItem(
agent=MagicMock(),
raw_item=ResponseOutputMessage(
id="msg-refusal",
content=[ResponseOutputRefusal(type="refusal", refusal=refusal)],
role="assistant",
status="completed",
type="message",
),
)
]
async def stream_events(self) -> Any:
for event in _NO_STREAM_EVENTS:
yield event
def cancel(self, mode: str = "immediate") -> None: # noqa: ARG002
return
async def _call_finish_scan(
coordinator: AgentCoordinator, agent_id: str, parent_id: str | None
) -> dict[str, Any]:
@@ -31,6 +68,43 @@ async def _call_finish_scan(
return parsed
async def _call_agent_finish(
coordinator: AgentCoordinator,
agent_id: str,
parent_id: str | None,
*,
report_to_parent: bool,
) -> dict[str, Any]:
ctx = ToolContext(
context={"coordinator": coordinator, "agent_id": agent_id, "parent_id": parent_id},
tool_name="agent_finish",
tool_call_id="call-1",
tool_arguments="{}",
)
result: str = await agent_finish.on_invoke_tool(
ctx,
json.dumps({"result_summary": "done", "report_to_parent": report_to_parent}),
)
parsed: dict[str, Any] = json.loads(result)
return parsed
async def _call_stop_agent(
coordinator: AgentCoordinator, agent_id: str, target_agent_id: str
) -> dict[str, Any]:
ctx = ToolContext(
context={"coordinator": coordinator, "agent_id": agent_id},
tool_name="stop_agent",
tool_call_id="call-1",
tool_arguments="{}",
)
result: str = await stop_agent.on_invoke_tool(
ctx, json.dumps({"target_agent_id": target_agent_id})
)
parsed: dict[str, Any] = json.loads(result)
return parsed
@pytest.mark.asyncio
async def test_reserve_stop_notifies_root_once(monkeypatch: pytest.MonkeyPatch) -> None:
coordinator = AgentCoordinator()
@@ -430,11 +504,11 @@ async def test_snapshot_round_trip_preserves_budget_pause() -> None:
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["stopped", "failed", "crashed"])
@pytest.mark.parametrize("status", ["completed", "stopped", "failed", "crashed"])
async def test_terminal_child_wakes_parked_parent(tmp_path: Any, status: str) -> None:
# Regression for #870: a child reaching a terminal state (e.g. MaxTurnsExceeded
# -> "stopped") must wake the parent parked in wait_for_message, so the root can
# finalize the scan instead of hanging for a completion report that never arrives.
# Regression for #870 and #947: a child reaching any terminal state - including a
# plain "completed" - must wake the parent parked in wait_for_agents, so the root
# can finalize the scan instead of hanging for a report that never arrives.
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "SQL Injection", parent_id="root")
@@ -446,13 +520,82 @@ async def test_terminal_child_wakes_parked_parent(tmp_path: Any, status: str) ->
assert not root_waiter.done()
await coordinator.set_status("child", status, error="Max turns (500) exceeded")
await _notify_parent_on_terminal(coordinator, "child", status)
await notify_parent_on_terminal(coordinator, "child", status)
await asyncio.wait_for(root_waiter, timeout=1.0)
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_agent_finish_without_report_still_wakes_parent(tmp_path: Any) -> None:
# Regression for #947: a child that completes with report_to_parent=False owes its
# parent a terminal notice, otherwise the parent waits out its full timeout.
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
root_waiter = asyncio.create_task(coordinator.wait_for_message("root"))
await asyncio.sleep(0)
await _call_agent_finish(coordinator, "child", "root", report_to_parent=False)
await asyncio.wait_for(root_waiter, timeout=1.0)
assert coordinator.statuses["child"] == "completed"
assert coordinator.pending_counts.get("root", 0) == 1
session.close()
@pytest.mark.asyncio
async def test_agent_finish_report_suppresses_the_terminal_notice(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
await _call_agent_finish(coordinator, "child", "root", report_to_parent=True)
# The exit backstop must not duplicate the report the child already delivered.
await execution._notify_parent_on_exit(coordinator, "child")
assert coordinator.pending_counts.get("root", 0) == 1
session.close()
@pytest.mark.asyncio
async def test_stop_agent_notifies_a_parent_that_is_not_the_stopper(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
await coordinator.register("grandchild", "sqli", parent_id="child")
session = SQLiteSession("child", tmp_path / "agents.db")
await coordinator.attach_runtime("child", session=session)
await _call_stop_agent(coordinator, "root", "grandchild")
assert coordinator.statuses["grandchild"] == "stopped"
assert coordinator.pending_counts.get("child", 0) == 1
# The stopper already knows; only the waiting parent needs telling.
assert coordinator.pending_counts.get("root", 0) == 0
session.close()
@pytest.mark.asyncio
async def test_stop_agent_does_not_notify_the_stopping_parent(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
await _call_stop_agent(coordinator, "root", "child")
assert coordinator.pending_counts.get("root", 0) == 0
session.close()
@pytest.mark.asyncio
async def test_notify_parent_on_terminal_ignores_non_terminal_status(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
@@ -461,7 +604,594 @@ async def test_notify_parent_on_terminal_ignores_non_terminal_status(tmp_path: A
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
await _notify_parent_on_terminal(coordinator, "child", "waiting")
await notify_parent_on_terminal(coordinator, "child", "waiting")
assert coordinator.pending_counts.get("root", 0) == 0
session.close()
class _RecordingStream:
def __init__(self) -> None:
self.cancelled = False
self.cancel_mode: str | None = None
def cancel(self, mode: str = "immediate") -> None:
self.cancelled = True
self.cancel_mode = mode
@pytest.mark.asyncio
async def test_terminal_notice_does_not_cancel_parent_stream(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
stream = _RecordingStream()
await coordinator.attach_runtime("root", session=session, interrupt_on_message=True)
await coordinator.attach_stream("root", stream)
await notify_parent_on_terminal(coordinator, "child", "crashed")
assert stream.cancelled is False
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_send_queues_without_session_and_drains_on_consume(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
assert await coordinator.send("root", {"from": "user", "content": "hello"}) is True
assert coordinator.pending_counts["root"] == 1
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
count, items = await coordinator.consume_pending("root", include_items=True)
assert count == 1
assert items[0]["content"] == "hello"
stored = await session.get_items()
last = cast("dict[str, Any]", stored[-1])
assert last["content"] == "hello"
session.close()
@pytest.mark.asyncio
async def test_error_parked_agent_only_released_by_user_message(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("child", tmp_path / "agents.db")
await coordinator.attach_runtime("child", session=session)
await coordinator.set_status("child", "crashed", error="boom")
await coordinator.send("child", {"from": "root", "content": "peer nudge"})
waiter = asyncio.create_task(coordinator.wait_for_message("child"))
await asyncio.sleep(0.05)
assert not waiter.done()
await coordinator.send("child", {"from": "user", "content": "wake up"})
assert await asyncio.wait_for(waiter, timeout=1.0) is True
count, items = await coordinator.consume_pending("child", include_items=True)
assert count == 2
assert items[0]["content"].endswith("peer nudge")
assert items[1]["content"] == "wake up"
session.close()
@pytest.mark.asyncio
async def test_wait_for_message_timeout_returns_false() -> None:
coordinator = AgentCoordinator()
await coordinator.register("child", "recon", parent_id="root")
assert await coordinator.wait_for_message("child", timeout=0.05) is False
@pytest.mark.asyncio
async def test_snapshot_round_trip_preserves_mailboxes() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.send("root", {"from": "user", "content": "queued"})
snap = await coordinator.snapshot()
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.pending_counts["root"] == 1
assert restored.runtimes["root"].mailbox == [{"from": "user", "content": "queued"}]
@pytest.mark.asyncio
async def test_run_cycle_parked_parks_instead_of_raising(
monkeypatch: pytest.MonkeyPatch,
) -> None:
async def _boom(*_args: Any, **_kwargs: Any) -> Any:
raise RuntimeError("unexpected explosion")
monkeypatch.setattr(execution, "_run_cycle", _boom)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
result = await execution._run_cycle_parked(
object(),
coordinator,
"root",
input_data=[],
run_config=None, # type: ignore[arg-type]
context={},
max_turns=5,
session=None,
event_sink=None,
hooks=None,
)
assert result is None
assert coordinator.statuses["root"] == "failed"
assert coordinator.errors["root"] == "unexpected explosion"
class _SalvageStream:
def __init__(self, replay: list[dict[str, Any]]) -> None:
self._replay = replay
def to_input_list(self) -> list[dict[str, Any]]:
return self._replay
@pytest.mark.asyncio
async def test_salvage_stream_to_session_preserves_full_history(tmp_path: Any) -> None:
session = SQLiteSession("child", tmp_path / "agents.db")
await session.add_items([{"role": "user", "content": "identity + task"}])
pre_run = list(await session.get_items())
# A crash mid-run: the stream produced two turns the SDK never committed.
stream = _SalvageStream(
[
{"role": "assistant", "content": "recon turn 1"},
{"role": "assistant", "content": "recon turn 2"},
]
)
await execution._salvage_stream_to_session(session, pre_run, stream, "child")
stored = [cast("dict[str, Any]", i) for i in await session.get_items()]
assert [i["content"] for i in stored] == [
"identity + task",
"recon turn 1",
"recon turn 2",
]
# A crash with nothing new to salvage leaves the session untouched.
await execution._salvage_stream_to_session(
session, list(await session.get_items()), _SalvageStream([]), "child"
)
assert len(await session.get_items()) == 3
session.close()
@pytest.mark.asyncio
async def test_seed_initial_input_persists_and_is_idempotent(tmp_path: Any) -> None:
session = SQLiteSession("child", tmp_path / "agents.db")
identity = [{"role": "user", "content": "You are agent recon (abc); do X."}]
assert await seed_initial_input(session, identity) is True
assert len(await session.get_items()) == 1
# A populated session is left untouched (no duplicate identity message).
assert await seed_initial_input(session, identity) is False
assert len(await session.get_items()) == 1
assert await seed_initial_input(session, []) is False
session.close()
@pytest.mark.asyncio
async def test_structured_provider_refusal_fails_interactive_agent(
monkeypatch: pytest.MonkeyPatch,
) -> None:
refusal = "This request was blocked under the provider's usage policy."
stream = _StructuredRefusalStream(refusal)
monkeypatch.setattr(
"strix.core.execution.Runner.run_streamed", lambda *_args, **_kwargs: stream
)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
result = await execution._run_cycle(
MagicMock(),
coordinator,
"root",
input_data="task",
run_config=MagicMock(),
context={},
max_turns=5,
session=None,
interactive=True,
event_sink=None,
hooks=None,
)
assert result is None
assert coordinator.statuses["root"] == "failed"
assert coordinator.errors["root"] == refusal
@pytest.mark.asyncio
async def test_structured_provider_refusal_fails_noninteractive_child(
tmp_path: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
refusal = "This request was blocked under the provider's usage policy."
stream = _StructuredRefusalStream(refusal)
monkeypatch.setattr(
"strix.core.execution.Runner.run_streamed", lambda *_args, **_kwargs: stream
)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
result = await execution._run_cycle(
MagicMock(),
coordinator,
"child",
input_data="task",
run_config=MagicMock(),
context={"parent_id": "root"},
max_turns=5,
session=None,
interactive=False,
event_sink=None,
hooks=None,
)
assert result is None
assert coordinator.statuses["child"] == "failed"
assert coordinator.errors["child"] == refusal
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_run_agent_loop_seeds_identity_before_first_cycle(
tmp_path: Any, monkeypatch: pytest.MonkeyPatch
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("child", tmp_path / "agents.db")
captured: dict[str, Any] = {}
async def _crash_first_turn(*_args: Any, **kwargs: Any) -> Any:
captured["input_data"] = kwargs.get("input_data")
captured["items_at_start"] = await session.get_items()
raise RuntimeError("first-turn crash")
monkeypatch.setattr(execution, "_run_cycle", _crash_first_turn)
identity = [{"role": "user", "content": "You are agent recon (abc); maintain your identity."}]
with pytest.raises(RuntimeError, match="first-turn crash"):
await execution.run_agent_loop(
agent=object(),
initial_input=identity,
run_config=None, # type: ignore[arg-type]
context={"agent_id": "child", "parent_id": "root"},
max_turns=5,
coordinator=coordinator,
agent_id="child",
interactive=False,
session=session,
)
# The first cycle ran with an empty input against the pre-seeded session.
assert captured["input_data"] == []
assert captured["items_at_start"]
# The identity/task survives the first-turn crash, so a revival can resume it.
stored = await session.get_items()
assert any("recon" in str(cast("dict[str, Any]", i).get("content", "")) for i in stored)
session.close()
def _scripted_cycle(
coordinator: AgentCoordinator,
agent_id: str,
statuses: list[str],
calls: list[Any],
) -> Any:
"""Fake run cycle that leaves ``agent_id`` in a scripted status per call."""
async def _cycle(*_args: Any, **kwargs: Any) -> Any:
calls.append(kwargs.get("input_data"))
status = statuses[min(len(calls) - 1, len(statuses) - 1)]
await coordinator.set_status(agent_id, status)
return MagicMock(final_output="plain text, no tool call")
return _cycle
async def _drive(
coordinator: AgentCoordinator,
agent_id: str,
*,
interactive: bool,
max_turns: int = 5,
) -> Any:
return await execution._run_until_lifecycle(
MagicMock(),
coordinator,
agent_id,
initial_input=[],
run_config=MagicMock(),
context={"agent_id": agent_id, "parent_id": None},
max_turns=max_turns,
session=None,
interactive=interactive,
event_sink=None,
hooks=None,
)
@pytest.mark.asyncio
async def test_interactive_text_only_turn_is_nudged_instead_of_parking(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A no-tool-call turn must not silently hand control back to the user."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "root", ["running", "completed"], calls),
)
await _drive(coordinator, "root", interactive=True)
assert len(calls) == 2
# The retry carries an explicit "call a tool" nudge rather than empty input.
nudge = calls[1][0]["content"]
assert "without a tool call" in nudge
assert "respond_to_user" in nudge
assert coordinator.statuses["root"] == "completed"
@pytest.mark.asyncio
async def test_interactive_explicit_park_gets_no_nudge(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""``waiting`` is only reachable via respond_to_user / wait_for_agents."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "root", ["waiting"], calls),
)
await _drive(coordinator, "root", interactive=True)
assert len(calls) == 1
assert coordinator.statuses["root"] == "waiting"
@pytest.mark.asyncio
async def test_interactive_recovery_exhaustion_parks_instead_of_crashing(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A human can resume an interactive scan, so exhaustion parks rather than dies."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "root", ["running"], calls),
)
await _drive(coordinator, "root", interactive=True)
assert len(calls) == execution._INTERACTIVE_TOOL_RECOVERY_LIMIT
assert coordinator.statuses["root"] == "waiting"
@pytest.mark.asyncio
async def test_interactive_subagent_exhaustion_tells_its_parent(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A parked child must report up so its parent stops waiting on it.
The parent is an agent, not a watching human, so a parent blocked in
wait_for_agents otherwise burns its whole timeout on a completion
report the child can no longer send.
"""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "child", ["running"], calls),
)
await _drive(coordinator, "child", interactive=True)
assert coordinator.statuses["child"] == "waiting"
pending, items = await coordinator.consume_pending("root", include_items=True)
assert pending == 1
notice = str(items[0])
assert "child" in notice
assert "parked" in notice
@pytest.mark.asyncio
async def test_interactive_root_exhaustion_notifies_nobody(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The root has no parent to report to, so parking stays silent."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "root", ["running"], []),
)
await _drive(coordinator, "root", interactive=True)
pending, _ = await coordinator.consume_pending("root")
assert pending == 0
@pytest.mark.asyncio
async def test_noninteractive_recovery_exhaustion_crashes(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""No user is present to resume an autonomous run, so it still fails loudly."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle",
_scripted_cycle(coordinator, "root", ["running"], calls),
)
with pytest.raises(MaxTurnsExceeded):
await _drive(coordinator, "root", interactive=False, max_turns=2)
assert len(calls) == 2
assert coordinator.statuses["root"] == "crashed"
@pytest.mark.asyncio
async def test_tool_required_message_is_persisted_to_the_session(tmp_path: Any) -> None:
session = SQLiteSession("root", tmp_path / "agents.db")
assert (
await execution._append_tool_required_message(
session=session,
context={"parent_id": None},
attempt=1,
limit=3,
interactive=True,
)
== []
)
stored = [cast("dict[str, Any]", i) for i in await session.get_items()]
assert "finish_scan" in stored[0]["content"]
assert "respond_to_user" in stored[0]["content"]
session.close()
@pytest.mark.asyncio
async def test_recovery_count_survives_a_snapshot_round_trip(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A resumed agent must not earn a fresh nudge budget and loop forever."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "root", ["running"], calls),
)
await _drive(coordinator, "root", interactive=True)
assert coordinator.recovery_counts["root"] == execution._INTERACTIVE_TOOL_RECOVERY_LIMIT
restored = AgentCoordinator()
await restored.restore(await coordinator.snapshot())
assert restored.recovery_counts["root"] == execution._INTERACTIVE_TOOL_RECOVERY_LIMIT
# The restored agent is already at its cap, so it parks after a single
# further text-only cycle instead of starting the whole budget over.
resumed_calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(restored, "root", ["running"], resumed_calls),
)
await _drive(restored, "root", interactive=True)
assert len(resumed_calls) == 1
assert restored.statuses["root"] == "waiting"
@pytest.mark.asyncio
async def test_recovery_count_is_cleared_by_a_lifecycle_tool(
monkeypatch: pytest.MonkeyPatch,
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
calls: list[Any] = []
monkeypatch.setattr(
execution,
"_run_cycle_parked",
_scripted_cycle(coordinator, "root", ["running", "completed"], calls),
)
await _drive(coordinator, "root", interactive=True)
assert "root" not in coordinator.recovery_counts
@pytest.mark.asyncio
async def test_agent_awaiting_a_human_is_never_auto_resumed() -> None:
"""The user can message any agent, so parking for one is not root-only."""
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
for agent_id in ("root", "child"):
await coordinator.park_waiting(agent_id, wait_kind="user")
assert await execution._plain_waiting_timeout(coordinator, agent_id) is None
@pytest.mark.asyncio
async def test_agent_awaiting_other_agents_is_re_checked_on_a_timer() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.park_waiting("root", wait_kind="agents")
timeout = await execution._plain_waiting_timeout(coordinator, "root")
assert timeout == execution._WAITING_AUTO_RESUME_TIMEOUT_S
@pytest.mark.asyncio
async def test_idle_auto_resumes_stop_after_their_budget() -> None:
"""A wedged agent must not burn a model turn per timeout for the whole scan."""
coordinator = AgentCoordinator()
await coordinator.register("child", "recon", parent_id="root")
await coordinator.park_waiting("child", wait_kind="agents")
for _ in range(execution._MAX_IDLE_AUTO_RESUMES):
assert await execution._plain_waiting_timeout(coordinator, "child") is not None
await coordinator.record_idle_resume("child")
assert await execution._plain_waiting_timeout(coordinator, "child") is None
# A real message is real progress, so the budget starts over.
await coordinator.reset_idle_resumes("child")
assert await execution._plain_waiting_timeout(coordinator, "child") is not None
@pytest.mark.asyncio
async def test_wait_kind_survives_a_snapshot_round_trip() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.park_waiting("root", wait_kind="user")
await coordinator.record_idle_resume("root")
restored = AgentCoordinator()
await restored.restore(await coordinator.snapshot())
assert restored.wait_kinds["root"] == "user"
assert restored.idle_resume_counts["root"] == 1
assert await execution._plain_waiting_timeout(restored, "root") is None
+19 -2
View File
@@ -15,6 +15,7 @@ from openai import (
RateLimitError,
)
from strix.config import codex
from strix.core import execution
from strix.core.agents import AgentCoordinator
@@ -55,11 +56,27 @@ def test_server_errors_are_transient() -> None:
assert execution._is_transient_model_error(_status_error(status)) is True
def test_rate_limit_is_not_retried_here() -> None:
def test_rate_limit_is_retried() -> None:
rate_limited = RateLimitError(
"slow down", response=httpx.Response(429, request=_request()), body=None
)
assert execution._is_transient_model_error(rate_limited) is False
assert execution._is_transient_model_error(rate_limited) is True
def test_dns_and_connection_errors_are_transient() -> None:
assert execution._is_transient_model_error(OSError("nodename nor servname provided")) is True
assert execution._is_transient_model_error(ConnectionError("reset")) is True
assert execution._is_transient_model_error(TimeoutError("timed out")) is True
def test_content_guardrail_is_not_retried() -> None:
guardrail = APIError(
"This content was flagged for possible cybersecurity risk",
_request(),
body=None,
)
assert codex.is_content_guardrail_error(guardrail) is True
assert execution._is_transient_model_error(guardrail) is False
def test_client_errors_are_not_transient() -> None:
+35
View File
@@ -129,6 +129,17 @@ def test_prompt_cache_kept_for_non_bedrock_claude_even_if_unmapped(monkeypatch:
]
def test_max_reasoning_effort_sent_as_raw_body_field() -> None:
# "max" is absent from the OpenAI SDK's Reasoning enum, and LiteLLM's DeepSeek
# mapping collapses every effort to thinking-enabled, so it has to ride along
# as a raw body field to reach the provider.
settings = make_model_settings(
"max", model_name="deepseek/deepseek-v4-flash", request_timeout=30
)
assert settings.reasoning is None
assert settings.extra_args == {"timeout": 30, "extra_body": {"reasoning_effort": "max"}}
def test_conversation_tail_breakpoint_moves_with_appended_transcript() -> None:
# LiteLLM must place the index=-1 cache_control on the last message however
# long the transcript grows.
@@ -272,6 +283,30 @@ def test_make_model_settings_omits_timeout_when_unset() -> None:
assert settings.extra_args is None
def test_make_model_settings_sets_extra_headers() -> None:
settings = make_model_settings(
"none",
model_name="openai/some-model",
extra_headers={"X-Feature-Key": "svc", "X-Tenant": "acme"},
)
assert settings.extra_headers == {"X-Feature-Key": "svc", "X-Tenant": "acme"}
def test_make_model_settings_omits_extra_headers_when_unset() -> None:
assert make_model_settings("none", model_name="gpt-4o").extra_headers is None
def test_make_model_settings_extra_headers_survive_reasoning_resolve() -> None:
settings = make_model_settings(
"high",
model_name="openai/o3",
extra_headers={"X-Feature-Key": "svc"},
)
assert settings.extra_headers == {"X-Feature-Key": "svc"}
def test_make_model_settings_timeout_survives_reasoning_resolve() -> None:
# Reasoning is resolved via ModelSettings.resolve(); the timeout in extra_args
# must not be dropped when a reasoning override is merged in.
+97
View File
@@ -0,0 +1,97 @@
"""Tests for LLM_EXTRA_HEADERS: custom default headers on OpenAI-compatible endpoints."""
from __future__ import annotations
import json
from typing import TYPE_CHECKING
import litellm
import pytest
from agents.models import _openai_shared
from strix.config import loader
from strix.config.loader import load_settings
from strix.config.models import configure_sdk_model_defaults
if TYPE_CHECKING:
from collections.abc import Iterator
_ENV_KEYS = ["STRIX_LLM", "LLM_API_KEY", "LLM_API_BASE", "LLM_EXTRA_HEADERS"]
@pytest.fixture(autouse=True)
def _reset(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in _ENV_KEYS:
monkeypatch.delenv(key, raising=False)
monkeypatch.setattr(loader, "_cached", None)
monkeypatch.setattr(loader, "_override", None)
saved_headers = litellm.headers
saved_client = _openai_shared.get_default_openai_client()
litellm.headers = None
try:
yield
finally:
litellm.headers = saved_headers
_openai_shared.set_default_openai_client(saved_client) # type: ignore[arg-type]
def test_extra_headers_parsed_from_json_env(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-A": "1", "X-B": "2"}))
settings = load_settings()
assert settings.llm.extra_headers == {"X-A": "1", "X-B": "2"}
def test_extra_headers_merged_into_litellm_headers(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("STRIX_LLM", "litellm/openai/some-model")
monkeypatch.setenv("LLM_API_BASE", "https://gateway.example/v1")
monkeypatch.setenv("LLM_API_KEY", "token")
headers = {"X-Feature-Key": "svc", "X-Tenant": "acme"}
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps(headers))
configure_sdk_model_defaults(load_settings())
current: object = litellm.headers
assert isinstance(current, dict)
assert current["X-Feature-Key"] == "svc"
assert current["X-Tenant"] == "acme"
def test_extra_headers_applied_to_native_openai_client(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("STRIX_LLM", "openai/some-model")
monkeypatch.setenv("LLM_API_BASE", "https://gateway.example/v1")
monkeypatch.setenv("LLM_API_KEY", "token")
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Feature-Key": "svc"}))
configure_sdk_model_defaults(load_settings())
client = _openai_shared.get_default_openai_client()
assert client is not None
assert client.default_headers.get("X-Feature-Key") == "svc"
assert str(client.base_url).rstrip("/") == "https://gateway.example/v1"
def test_extra_headers_applied_to_native_openai_without_custom_base(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("STRIX_LLM", "openai/gpt-5")
monkeypatch.setenv("LLM_API_KEY", "token")
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Feature-Key": "svc"}))
configure_sdk_model_defaults(load_settings())
client = _openai_shared.get_default_openai_client()
assert client is not None
assert client.default_headers.get("X-Feature-Key") == "svc"
def test_no_extra_headers_leaves_litellm_headers_untouched(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("STRIX_LLM", "openai/some-model")
monkeypatch.setenv("LLM_API_BASE", "https://gateway.example/v1")
monkeypatch.setenv("LLM_API_KEY", "token")
configure_sdk_model_defaults(load_settings())
assert litellm.headers is None
-143
View File
@@ -1,143 +0,0 @@
"""Tests for symlink-safe LocalDir staging."""
from __future__ import annotations
from typing import TYPE_CHECKING
from strix.runtime.local_dir_staging import stage_symlink_safe_dir, tree_has_symlink
if TYPE_CHECKING:
from pathlib import Path
def _make_repo(tmp_path: Path) -> Path:
repo = tmp_path / "repo"
(repo / "pkg").mkdir(parents=True)
(repo / "pkg" / "mod.py").write_text("x = 1\n")
(repo / "README.md").write_text("readme\n")
return repo
def test_tree_without_symlinks_used_as_is(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
upload_path, staged = stage_symlink_safe_dir(repo)
assert staged is None
assert upload_path == repo.resolve()
assert not tree_has_symlink(repo)
def test_in_tree_file_symlink_is_dereferenced(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
(repo / "link.py").symlink_to(repo / "pkg" / "mod.py")
upload_path, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert upload_path == staged
assert not (staged / "link.py").is_symlink()
assert (staged / "link.py").read_text() == "x = 1\n"
assert (staged / "pkg" / "mod.py").read_text() == "x = 1\n"
assert not tree_has_symlink(staged)
def test_in_tree_relative_dir_symlink_is_dereferenced(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
(repo / "pkg_alias").symlink_to("pkg")
_upload, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert (staged / "pkg_alias" / "mod.py").read_text() == "x = 1\n"
assert not tree_has_symlink(staged)
def test_out_of_tree_symlink_is_dropped(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
outside = tmp_path / "outside.txt"
outside.write_text("secret\n")
(repo / "escape.txt").symlink_to(outside)
(repo / "abs_escape").symlink_to("/etc")
_upload, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert not (staged / "escape.txt").exists()
assert not (staged / "abs_escape").exists()
assert (staged / "README.md").exists()
def test_dangling_symlink_is_dropped(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
(repo / "dangling").symlink_to(repo / "does-not-exist")
_upload, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert not (staged / "dangling").exists()
assert not (staged / "dangling").is_symlink()
def test_cyclic_symlink_terminates(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
(repo / "self").symlink_to(repo)
(repo / "pkg" / "up").symlink_to("..")
_upload, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert (staged / "README.md").exists()
assert not tree_has_symlink(staged)
def test_nested_symlinks_inside_linked_dir(tmp_path: Path) -> None:
repo = _make_repo(tmp_path)
shared = repo / "shared"
shared.mkdir()
(shared / "conf.json").write_text("{}\n")
(shared / "escape").symlink_to("/etc/passwd")
(repo / "pkg" / "shared_link").symlink_to(shared)
_upload, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert (staged / "pkg" / "shared_link" / "conf.json").read_text() == "{}\n"
assert not (staged / "pkg" / "shared_link" / "escape").exists()
assert not (staged / "shared" / "escape").exists()
def test_staged_path_has_no_symlink_ancestor(tmp_path: Path, monkeypatch) -> None: # noqa: ANN001
"""The staging directory itself must never sit behind a symlink.
``tempfile.mkdtemp()`` honors ``$TMPDIR``, and on macOS the default
``$TMPDIR`` resolves through ``/var``, which is itself a symlink to
``/private/var``. ``LocalDir`` rejects any symlink component in its
source path, so returning the raw ``mkdtemp()`` result breaks every
local-dir upload on macOS whenever the source tree contains a symlink.
This reproduces that shape without depending on the host OS layout.
"""
repo = _make_repo(tmp_path)
(repo / "link.py").symlink_to(repo / "pkg" / "mod.py")
real_tmp_root = tmp_path / "real_tmp"
real_tmp_root.mkdir()
symlinked_tmp_root = tmp_path / "tmp_symlink"
symlinked_tmp_root.symlink_to(real_tmp_root)
def fake_mkdtemp(prefix: str = "") -> str:
real_dir = real_tmp_root / f"{prefix}fake"
real_dir.mkdir()
return str(symlinked_tmp_root / real_dir.name)
monkeypatch.setattr(
"strix.runtime.local_dir_staging.tempfile.mkdtemp", fake_mkdtemp
)
upload_path, staged = stage_symlink_safe_dir(repo)
assert staged is not None
assert upload_path == staged
for path in (staged, *staged.parents):
assert not path.is_symlink(), f"staged path has a symlink ancestor: {path}"
+105 -157
View File
@@ -1,171 +1,138 @@
"""Tests for local-source sizing and ``--mount`` target helpers in interface.utils."""
"""Tests for local-source collection and mount policy in interface.utils."""
from __future__ import annotations
import logging
import os
import sys
from typing import TYPE_CHECKING, Any
from pathlib import Path
from typing import Any
import pytest
if TYPE_CHECKING:
from pathlib import Path
from strix.interface.utils import (
build_mount_targets_info,
check_mountable_dir,
collect_local_sources,
dedupe_local_targets,
directory_size_bytes,
find_oversized_local_targets,
infer_target_type,
read_target_list_file,
)
def _write_file(path: Path, size: int) -> None:
path.write_bytes(b"x" * size)
def _local_target(target_path: str) -> dict[str, Any]:
return {
"type": "local_code",
"details": {"target_path": target_path, "workspace_subdir": "repo"},
"original": target_path,
}
def _local_target(target_path: str, *, mount: bool = False) -> dict[str, Any]:
details: dict[str, Any] = {"target_path": target_path, "workspace_subdir": "repo"}
if mount:
details["mount"] = True
return {"type": "local_code", "details": details, "original": target_path}
def test_collect_local_sources_protects_the_users_own_git() -> None:
sources = collect_local_sources([_local_target("/code")])
assert sources == [
{"source_path": "/code", "workspace_subdir": "repo", "protect_metadata": True}
]
def test_directory_size_empty_dir_is_zero(tmp_path: Path) -> None:
assert directory_size_bytes(tmp_path) == 0
def test_directory_size_sums_flat_and_nested_files(tmp_path: Path) -> None:
_write_file(tmp_path / "a.txt", 100)
nested = tmp_path / "sub" / "deep"
nested.mkdir(parents=True)
_write_file(nested / "b.txt", 250)
assert directory_size_bytes(tmp_path) == 350
def test_directory_size_skips_symlinks(tmp_path: Path) -> None:
_write_file(tmp_path / "real.txt", 100)
(tmp_path / "link.txt").symlink_to(tmp_path / "real.txt")
# The symlink target is counted once via the real file, not doubled.
assert directory_size_bytes(tmp_path) == 100
@pytest.mark.skipif(sys.platform == "win32", reason="relies on POSIX permissions")
def test_directory_size_logs_and_skips_unreadable_subdir(
tmp_path: Path, caplog: pytest.LogCaptureFixture
) -> None:
if hasattr(os, "geteuid") and os.geteuid() == 0:
pytest.skip("root bypasses directory permissions")
_write_file(tmp_path / "top.txt", 100)
locked = tmp_path / "locked"
locked.mkdir()
_write_file(locked / "secret.bin", 9999)
locked.chmod(0o000)
try:
with caplog.at_level(logging.WARNING):
size = directory_size_bytes(tmp_path)
finally:
locked.chmod(0o755)
# The unreadable subtree is excluded (not silently treated as readable) and
# the omission is logged rather than vanishing without a trace.
assert size == 100
assert any("Could not read" in record.message for record in caplog.records)
def test_find_oversized_returns_nothing_under_limit(tmp_path: Path) -> None:
_write_file(tmp_path / "a.txt", 100)
targets = [_local_target(str(tmp_path))]
assert find_oversized_local_targets(targets, max_bytes=1000) == []
def test_find_oversized_returns_target_over_limit(tmp_path: Path) -> None:
_write_file(tmp_path / "big.bin", 500)
targets = [_local_target(str(tmp_path))]
result = find_oversized_local_targets(targets, max_bytes=100)
assert result == [(str(tmp_path), 500)]
def test_find_oversized_ignores_mounted_targets(tmp_path: Path) -> None:
_write_file(tmp_path / "big.bin", 500)
targets = [_local_target(str(tmp_path), mount=True)]
assert find_oversized_local_targets(targets, max_bytes=100) == []
def test_find_oversized_ignores_non_local_targets() -> None:
targets = [{"type": "web_application", "details": {"target_url": "https://x"}}]
assert find_oversized_local_targets(targets, max_bytes=1) == []
@pytest.mark.parametrize("disabled", [0, -1])
def test_find_oversized_disabled_for_non_positive_limit(tmp_path: Path, disabled: int) -> None:
_write_file(tmp_path / "big.bin", 500)
targets = [_local_target(str(tmp_path))]
assert find_oversized_local_targets(targets, max_bytes=disabled) == []
def test_collect_local_sources_propagates_mount_flag() -> None:
copied = _local_target("/copied")
copied["details"]["workspace_subdir"] = "copied"
mounted = _local_target("/mounted", mount=True)
mounted["details"]["workspace_subdir"] = "mounted"
sources = collect_local_sources([copied, mounted])
by_path = {s["source_path"]: s for s in sources}
assert by_path["/copied"]["mount"] is False
assert by_path["/mounted"]["mount"] is True
def test_collect_local_sources_repository_is_never_mounted() -> None:
def test_collect_local_sources_leaves_a_clone_writable() -> None:
repo = {
"type": "repository",
"details": {"cloned_repo_path": "/clone", "workspace_subdir": "clone"},
}
sources = collect_local_sources([repo])
assert sources == [{"source_path": "/clone", "workspace_subdir": "clone", "mount": False}]
assert sources == [
{"source_path": "/clone", "workspace_subdir": "clone", "protect_metadata": False}
]
def test_build_mount_targets_info_for_valid_dir(tmp_path: Path) -> None:
result = build_mount_targets_info([str(tmp_path)])
assert len(result) == 1
entry = result[0]
assert entry["type"] == "local_code"
assert entry["details"]["mount"] is True
assert entry["details"]["target_path"] == str(tmp_path.resolve())
def test_check_mountable_dir_accepts_a_project_dir(tmp_path: Path) -> None:
check_mountable_dir(tmp_path)
def test_build_mount_targets_info_rejects_missing_path(tmp_path: Path) -> None:
missing = tmp_path / "does-not-exist"
def test_check_mountable_dir_rejects_missing_path(tmp_path: Path) -> None:
with pytest.raises(ValueError, match="not an existing directory"):
build_mount_targets_info([str(missing)])
check_mountable_dir(tmp_path / "nope")
def test_build_mount_targets_info_rejects_file(tmp_path: Path) -> None:
file_path = tmp_path / "a-file.txt"
_write_file(file_path, 10)
with pytest.raises(ValueError, match="not an existing directory"):
build_mount_targets_info([str(file_path)])
def test_check_mountable_dir_rejects_filesystem_root() -> None:
with pytest.raises(ValueError, match="Refusing to mount"):
check_mountable_dir(Path("/"))
@pytest.mark.parametrize("empty", ["", " "])
def test_build_mount_targets_info_rejects_empty_path(empty: str) -> None:
# An empty path would otherwise resolve to the current working directory
# and silently bind-mount it into the sandbox.
with pytest.raises(ValueError, match="must not be empty"):
build_mount_targets_info([empty])
def test_check_mountable_dir_rejects_home(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
home = tmp_path / "home"
home.mkdir()
monkeypatch.setenv("HOME", str(home))
monkeypatch.setattr(Path, "home", classmethod(lambda _cls: home))
with pytest.raises(ValueError, match="Refusing to mount"):
check_mountable_dir(home)
def test_check_mountable_dir_rejects_system_root() -> None:
etc = Path("/etc")
if not etc.is_dir():
pytest.skip("no /etc on this platform")
with pytest.raises(ValueError, match="Refusing to mount"):
check_mountable_dir(etc)
def test_check_mountable_dir_rejects_the_shared_home_root() -> None:
home_root = Path("/home")
if not home_root.is_dir():
pytest.skip("no /home on this platform")
with pytest.raises(ValueError, match="Refusing to mount"):
check_mountable_dir(home_root)
def test_check_mountable_dir_matches_forbidden_names_case_insensitively(tmp_path: Path) -> None:
ssh_dir = tmp_path / ".SSH"
ssh_dir.mkdir()
with pytest.raises(ValueError, match="holds credentials"):
check_mountable_dir(ssh_dir)
def test_check_mountable_dir_rejects_credential_dirs(tmp_path: Path) -> None:
ssh_dir = tmp_path / ".ssh"
ssh_dir.mkdir()
with pytest.raises(ValueError, match="holds credentials"):
check_mountable_dir(ssh_dir)
def test_check_mountable_dir_rejects_credential_subdirs(tmp_path: Path) -> None:
keys = tmp_path / ".ssh" / "keys"
keys.mkdir(parents=True)
with pytest.raises(ValueError, match="holds credentials"):
check_mountable_dir(keys)
def test_check_mountable_dir_rejects_system_subdirs() -> None:
system_subdir = next((p for p in (Path("/etc/ssl"), Path("/usr/bin")) if p.is_dir()), None)
if system_subdir is None:
pytest.skip("no system subdirectory on this platform")
with pytest.raises(ValueError, match="Refusing to mount"):
check_mountable_dir(system_subdir)
def test_check_mountable_dir_accepts_a_project_under_the_home_root(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
project = tmp_path / "home" / "dev" / "project"
project.mkdir(parents=True)
monkeypatch.setattr(Path, "home", classmethod(lambda _cls: tmp_path / "home" / "dev"))
check_mountable_dir(project)
def test_infer_target_type_applies_the_mount_policy() -> None:
with pytest.raises(ValueError, match="Refusing to mount"):
infer_target_type("/etc")
def test_read_target_list_file_strips_blank_lines(tmp_path: Path) -> None:
target_list = tmp_path / "targets.txt"
target_list.write_text(
"\n"
" https://test1.com/ \n"
"\n"
"http://test2.com:5789/\n"
" \n",
"\n https://test1.com/ \n\nhttp://test2.com:5789/\n \n",
encoding="utf-8",
)
@@ -178,10 +145,7 @@ def test_read_target_list_file_strips_blank_lines(tmp_path: Path) -> None:
def test_read_target_list_file_ignores_comment_lines(tmp_path: Path) -> None:
target_list = tmp_path / "targets.txt"
target_list.write_text(
"# production targets\n"
"https://test1.com/\n"
" # staging targets\n"
"http://test2.com:5789/\n",
"# production targets\nhttps://test1.com/\n # staging targets\nhttp://test2.com:5789/\n",
encoding="utf-8",
)
@@ -222,28 +186,12 @@ def test_dedupe_keeps_distinct_targets_in_order() -> None:
targets = [
_local_target("/a"),
{"type": "web_application", "details": {"target_url": "https://x"}},
_local_target("/b", mount=True),
_local_target("/b"),
]
assert dedupe_local_targets(targets) == targets
def test_dedupe_mount_supersedes_copied_same_path() -> None:
copied = _local_target("/repo")
mounted = _local_target("/repo", mount=True)
# Copied first, then mounted: the single surviving entry is the mount.
result = dedupe_local_targets([copied, mounted])
assert len(result) == 1
assert result[0]["details"]["mount"] is True
# Order-independent: mounted first, copied second also yields the mount.
result_rev = dedupe_local_targets([mounted, copied])
assert len(result_rev) == 1
assert result_rev[0]["details"]["mount"] is True
def test_dedupe_collapses_duplicate_mounts() -> None:
result = dedupe_local_targets(
[_local_target("/repo", mount=True), _local_target("/repo", mount=True)]
)
assert len(result) == 1
def test_dedupe_collapses_the_same_path() -> None:
assert dedupe_local_targets([_local_target("/repo"), _local_target("/repo")]) == [
_local_target("/repo")
]
+66
View File
@@ -0,0 +1,66 @@
"""Tests for the ``respond_to_user`` yield tool."""
from __future__ import annotations
import json
from typing import Any
import pytest
from agents.tool_context import ToolContext
from strix.core.agents import AgentCoordinator
from strix.tools.respond.tool import respond_to_user
async def _call(context: dict[str, Any], message: str = "here is what I found") -> 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, json.dumps({"message": message}))
return json.loads(raw) # type: ignore[no-any-return]
async def _context(*, interactive: bool, agent_id: str = "root") -> dict[str, Any]:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
return {"coordinator": coordinator, "agent_id": agent_id, "interactive": interactive}
@pytest.mark.asyncio
async def test_parks_the_agent_and_carries_the_message() -> None:
context = await _context(interactive=True)
result = await _call(context)
coordinator = context["coordinator"]
assert result["success"] is True
assert result["wait_outcome"] == "waiting"
assert result["message"] == "here is what I found"
assert coordinator.statuses["root"] == "waiting"
# Recorded as a human wait, so the driver never auto-resumes it.
assert coordinator.wait_kinds["root"] == "user"
@pytest.mark.asyncio
async def test_rejected_in_an_autonomous_run() -> None:
context = await _context(interactive=False)
result = await _call(context)
assert result["success"] is False
assert "finish_scan" in result["error"]
assert context["coordinator"].statuses["root"] == "running"
@pytest.mark.asyncio
async def test_a_message_that_already_arrived_is_taken_instead_of_parking() -> None:
context = await _context(interactive=True)
coordinator = context["coordinator"]
await coordinator.send("root", {"from": "user", "content": "wait, one more thing"})
result = await _call(context)
assert result["wait_outcome"] == "message_arrived"
assert result["pending_messages"] == 1
assert coordinator.statuses["root"] == "running"
+1
View File
@@ -40,6 +40,7 @@ async def test_persistent_rate_limit_stops_gracefully(
force_required_tool_choice=False,
timeout=300,
prompt_cache=True,
extra_headers=None,
),
runtime=types.SimpleNamespace(max_context_images=3),
)
+1
View File
@@ -48,6 +48,7 @@ def _patch_engine_scaffold(
force_required_tool_choice=False,
timeout=300,
prompt_cache=True,
extra_headers=None,
),
runtime=types.SimpleNamespace(max_context_images=3),
)
+147 -54
View File
@@ -1,4 +1,4 @@
"""Tests for build_session_entries: splitting copied vs bind-mounted sources."""
"""Tests for how local sources reach the sandbox: bind mounts or manifest upload."""
from __future__ import annotations
@@ -6,82 +6,175 @@ from typing import TYPE_CHECKING, Any
from agents.sandbox.entries import LocalDir
from strix.runtime.session_manager import build_session_entries
from strix.runtime.backends import (
_BACKENDS,
_BIND_MOUNT_BACKENDS,
backend_supports_bind_mounts,
register_backend,
)
from strix.runtime.session_manager import build_bind_mounts, build_manifest_entries
if TYPE_CHECKING:
from pathlib import Path
def _source(subdir: str, path: str, *, mount: bool = False) -> dict[str, Any]:
return {"source_path": path, "workspace_subdir": subdir, "mount": mount}
def _source(subdir: str, path: str, *, protect_metadata: bool = False) -> dict[str, Any]:
return {"source_path": path, "workspace_subdir": subdir, "protect_metadata": protect_metadata}
def test_copied_source_becomes_localdir_entry(tmp_path: Path) -> None:
entries, bind_mounts, staged_dirs = build_session_entries([_source("repo", str(tmp_path))])
assert bind_mounts == []
assert staged_dirs == []
assert isinstance(entries["repo"], LocalDir)
assert entries["repo"].src == tmp_path.resolve()
def test_mounted_source_becomes_bind_mount(tmp_path: Path) -> None:
entries, bind_mounts, _staged = build_session_entries(
[_source("repo", str(tmp_path), mount=True)]
)
assert entries == {}
assert bind_mounts == [
def test_source_becomes_writable_bind_mount(tmp_path: Path) -> None:
assert build_bind_mounts([_source("repo", str(tmp_path))]) == [
{
"source": str(tmp_path.resolve()),
"target": "/workspace/repo",
"read_only": True,
"read_only": False,
}
]
def test_mixed_sources_split_correctly(tmp_path: Path) -> None:
copied = tmp_path / "copied"
mounted = tmp_path / "mounted"
copied.mkdir()
mounted.mkdir()
def test_git_dir_is_remounted_read_only_when_protected(tmp_path: Path) -> None:
(tmp_path / ".git").mkdir()
entries, bind_mounts, _staged = build_session_entries(
[
_source("copied", str(copied)),
_source("mounted", str(mounted), mount=True),
]
)
mounts = build_bind_mounts([_source("repo", str(tmp_path), protect_metadata=True)])
assert list(entries) == ["copied"]
assert isinstance(entries["copied"], LocalDir)
assert [m["target"] for m in bind_mounts] == ["/workspace/mounted"]
assert mounts == [
{"source": str(tmp_path.resolve()), "target": "/workspace/repo", "read_only": False},
{
"source": str((tmp_path / ".git").resolve()),
"target": "/workspace/repo/.git",
"read_only": True,
},
]
def test_agent_instruction_dirs_are_protected_too(tmp_path: Path) -> None:
(tmp_path / ".agents").mkdir()
(tmp_path / ".codex").mkdir()
mounts = build_bind_mounts([_source("repo", str(tmp_path), protect_metadata=True)])
assert [(m["target"], m["read_only"]) for m in mounts] == [
("/workspace/repo", False),
("/workspace/repo/.agents", True),
("/workspace/repo/.codex", True),
]
def test_worktree_git_pointer_file_is_protected(tmp_path: Path) -> None:
gitdir = tmp_path / "nested" / "gitdir"
gitdir.mkdir(parents=True)
(tmp_path / ".git").write_text(f"gitdir: {gitdir}\n", encoding="utf-8")
mounts = build_bind_mounts([_source("repo", str(tmp_path), protect_metadata=True)])
assert [(m["target"], m["read_only"]) for m in mounts] == [
("/workspace/repo", False),
("/workspace/repo/.git", True),
("/workspace/repo/nested/gitdir", True),
]
def test_git_pointer_to_a_missing_gitdir_is_not_mounted(tmp_path: Path) -> None:
(tmp_path / ".git").write_text(f"gitdir: {tmp_path / 'gone'}\n", encoding="utf-8")
mounts = build_bind_mounts([_source("repo", str(tmp_path), protect_metadata=True)])
assert [m["target"] for m in mounts] == ["/workspace/repo", "/workspace/repo/.git"]
def test_git_pointer_outside_the_tree_needs_no_nested_mount(tmp_path: Path) -> None:
tree = tmp_path / "worktree"
tree.mkdir()
(tree / ".git").write_text(f"gitdir: {tmp_path / 'main' / '.git'}\n", encoding="utf-8")
mounts = build_bind_mounts([_source("repo", str(tree), protect_metadata=True)])
assert [m["target"] for m in mounts] == ["/workspace/repo", "/workspace/repo/.git"]
def test_metadata_symlinked_outside_the_tree_is_not_mounted(tmp_path: Path) -> None:
outside = tmp_path / "elsewhere"
outside.mkdir()
tree = tmp_path / "repo"
tree.mkdir()
(tree / ".git").symlink_to(outside, target_is_directory=True)
mounts = build_bind_mounts([_source("repo", str(tree), protect_metadata=True)])
assert [m["target"] for m in mounts] == ["/workspace/repo"]
def test_no_git_guard_without_a_git_dir(tmp_path: Path) -> None:
mounts = build_bind_mounts([_source("repo", str(tmp_path), protect_metadata=True)])
assert [m["target"] for m in mounts] == ["/workspace/repo"]
def test_clone_keeps_its_git_writable(tmp_path: Path) -> None:
(tmp_path / ".git").mkdir()
mounts = build_bind_mounts([_source("clone", str(tmp_path), protect_metadata=False)])
assert [m["target"] for m in mounts] == ["/workspace/clone"]
def test_multiple_sources_each_get_a_mount(tmp_path: Path) -> None:
first = tmp_path / "first"
second = tmp_path / "second"
first.mkdir()
second.mkdir()
mounts = build_bind_mounts([_source("first", str(first)), _source("second", str(second))])
assert [m["target"] for m in mounts] == ["/workspace/first", "/workspace/second"]
assert all(m["read_only"] is False for m in mounts)
def test_incomplete_sources_are_skipped() -> None:
entries, bind_mounts, staged_dirs = build_session_entries(
[
{"source_path": "", "workspace_subdir": "x"},
{"source_path": "/p", "workspace_subdir": ""},
]
assert (
build_bind_mounts(
[
{"source_path": "", "workspace_subdir": "x"},
{"source_path": "/p", "workspace_subdir": ""},
]
)
== []
)
assert entries == {}
assert bind_mounts == []
assert staged_dirs == []
def test_symlink_tree_is_staged(tmp_path: Path) -> None:
repo = tmp_path / "repo"
repo.mkdir()
(repo / "real.txt").write_text("content")
(repo / "link.txt").symlink_to(repo / "real.txt")
def test_manifest_entries_upload_sources_for_backends_without_bind_mounts(
tmp_path: Path,
) -> None:
entries = build_manifest_entries([_source("repo", str(tmp_path), protect_metadata=True)])
entries, _mounts, staged_dirs = build_session_entries([_source("repo", str(repo))])
assert len(staged_dirs) == 1
assert set(entries) == {"repo"}
entry = entries["repo"]
assert isinstance(entry, LocalDir)
assert entry.src == staged_dirs[0]
assert not (staged_dirs[0] / "link.txt").is_symlink()
assert (staged_dirs[0] / "link.txt").read_text() == "content"
assert entry.src == tmp_path.resolve()
def test_manifest_entries_skip_incomplete_sources() -> None:
assert (
build_manifest_entries(
[
{"source_path": "", "workspace_subdir": "x"},
{"source_path": "/p", "workspace_subdir": ""},
]
)
== {}
)
def test_only_bind_mount_capable_backends_are_registered_as_such() -> None:
assert backend_supports_bind_mounts("docker")
assert not backend_supports_bind_mounts("e2b")
async def _remote_backend(**_kwargs: Any) -> tuple[Any, Any]:
return object(), object()
try:
register_backend("e2b", _remote_backend)
assert not backend_supports_bind_mounts("e2b")
register_backend("e2b", _remote_backend, supports_bind_mounts=True)
assert backend_supports_bind_mounts("e2b")
finally:
_BACKENDS.pop("e2b", None)
_BIND_MOUNT_BACKENDS.discard("e2b")
+125 -2
View File
@@ -4,9 +4,11 @@ from __future__ import annotations
import json
import os
import sqlite3
import urllib.error
import urllib.request
from typing import TYPE_CHECKING
from urllib.parse import urlsplit
from strix.core.paths import latest_run_dir, runs_base_dir
from strix.interface.viewer.server import serve
@@ -86,6 +88,75 @@ def test_build_run_state_from_agents_json(tmp_path: Path) -> None:
assert state["events"] == []
def test_build_run_state_keeps_same_call_id_separate_per_agent(tmp_path: Path) -> None:
run_dir = _make_run(tmp_path, "tools", status="completed", end_time=None)
agents_db = run_dir / ".state" / "agents.db"
rows = [
(
"root",
{
"type": "function_call",
"call_id": "exec_command_0",
"name": "exec_command",
"arguments": json.dumps({"cmd": "echo root"}),
},
),
(
"root",
{
"type": "function_call_output",
"call_id": "exec_command_0",
"output": json.dumps({"success": True, "output": "root"}),
},
),
(
"child",
{
"type": "function_call",
"call_id": "exec_command_0",
"name": "exec_command",
"arguments": json.dumps({"cmd": "echo child"}),
},
),
(
"child",
{
"type": "function_call_output",
"call_id": "exec_command_0",
"output": json.dumps({"success": True, "output": "child"}),
},
),
]
with sqlite3.connect(agents_db) as conn:
conn.execute(
"""
create table agent_messages (
id integer primary key,
session_id text not null,
message_data text not null,
created_at text not null
)
"""
)
conn.executemany(
"""
insert into agent_messages (session_id, message_data, created_at)
values (?, ?, '2026-01-01T00:00:00+00:00')
""",
[(agent_id, json.dumps(message)) for agent_id, message in rows],
)
state = build_run_state(run_dir)
tools = [event for event in state["events"] if event["type"] == "tool"]
assert len(tools) == 2
by_agent = {event["agent_id"]: event for event in tools}
assert by_agent["root"]["data"]["args"] == {"cmd": "echo root"}
assert by_agent["root"]["data"]["result"]["output"] == "root"
assert by_agent["child"]["data"]["args"] == {"cmd": "echo child"}
assert by_agent["child"]["data"]["result"]["output"] == "child"
def _get(url: str, *, cookie: str | None = None) -> tuple[int, str, bytes]:
headers = {"Cookie": cookie} if cookie else {}
req = urllib.request.Request(url, headers=headers) # noqa: S310 - localhost test server
@@ -271,6 +342,11 @@ def _session_cookie(url: str, token: str) -> str:
return raw.split(";", 1)[0]
def _cookie_name(url: str) -> str:
"""The per-server session cookie name, derived from the bound port."""
return f"strix_viewer_session_{urlsplit(url).port}"
def _get_status(url: str, *, cookie: str | None = None) -> int:
headers = {"Cookie": cookie} if cookie else {}
req = urllib.request.Request(url, headers=headers) # noqa: S310 - localhost test server
@@ -311,7 +387,7 @@ def test_capability_issued_only_for_tokened_bootstrap(
# Only the correct bootstrap token mints the session cookie.
with urllib.request.urlopen(f"{url}/?token={token}") as resp: # noqa: S310 # nosec B310
cookie = str(resp.headers.get("Set-Cookie", ""))
assert "strix_viewer_session=" in cookie
assert f"{_cookie_name(url)}=" in cookie
assert "HttpOnly" in cookie and "SameSite=Strict" in cookie
# Static assets never carry it.
@@ -344,7 +420,7 @@ def test_unauthorized_client_cannot_acquire_capability(
url,
"/api/agents/steer",
{"agent_id": "root", "message": "pwn"},
cookie="strix_viewer_session=",
cookie=f"{_cookie_name(url)}=",
)
assert status == 403
assert delivered == []
@@ -541,6 +617,53 @@ def test_runs_list_requires_session_and_verification(
httpd.server_close()
def test_concurrent_servers_use_distinct_cookies(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Cookies are host-scoped, not port-scoped: two viewers on 127.0.0.1 must
not share a cookie slot, and one server's cookie must not pass the other's
session gate."""
run_a = _make_run(tmp_path / "a", "run-a", status="running", end_time=None)
run_b = _make_run(tmp_path / "b", "run-b", status="running", end_time=None)
_bundle(tmp_path, monkeypatch)
monkeypatch.setattr(
"strix.interface.viewer.auth.read_auth", lambda: {"email": "a@b.com", "token": "t"}
)
monkeypatch.setattr("strix.interface.viewer.auth.is_verified", lambda: True)
httpd_a, url_a, token_a = serve(run_a, open_browser=False)
httpd_b, url_b, token_b = serve(run_b, open_browser=False)
try:
cookie_a = _session_cookie(url_a, token_a)
cookie_b = _session_cookie(url_b, token_b)
# The two servers mint differently named cookies, so a browser stores both.
assert cookie_a.split("=", 1)[0] == _cookie_name(url_a)
assert cookie_b.split("=", 1)[0] == _cookie_name(url_b)
assert cookie_a.split("=", 1)[0] != cookie_b.split("=", 1)[0]
def _status(url: str, cookie: str) -> dict[str, object]:
_, _, body = _get(f"{url}/api/auth/status", cookie=cookie)
return dict(json.loads(body))
# Each server honors its own cookie...
assert _status(url_a, cookie_a)["verified"] is True
assert _status(url_b, cookie_b)["verified"] is True
# ...but treats the other server's cookie as session-less.
assert _status(url_a, cookie_b)["verified"] is False
assert _status(url_b, cookie_a)["verified"] is False
# Even both cookies together (what a real browser would send) only
# match the token minted by the receiving server.
both = f"{cookie_a}; {cookie_b}"
assert _status(url_a, both)["verified"] is True
assert _status(url_b, both)["verified"] is True
finally:
httpd_a.shutdown()
httpd_a.server_close()
httpd_b.shutdown()
httpd_b.server_close()
def test_server_rejects_path_traversal(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
run_dir = _make_run(tmp_path, "guard", status="completed", end_time="2026-01-01T00:00:00Z")
secret = tmp_path / "secret.txt"
Generated
+1 -1
View File
@@ -2411,7 +2411,7 @@ wheels = [
[[package]]
name = "strix-agent"
version = "1.4.0"
version = "1.4.1"
source = { editable = "." }
dependencies = [
{ name = "caido-sdk-client" },