mirror of
https://github.com/usestrix/strix.git
synced 2026-08-21 02:45:31 +02:00
Compare commits
35
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d56424090f | ||
|
|
8a458c3187 | ||
|
|
76e97e6a59 | ||
|
|
885b2ca5c5 | ||
|
|
980216860e | ||
|
|
d4e58b2cd0 | ||
|
|
e9ebdc502f | ||
|
|
ebb3a62a99 | ||
|
|
1a2fa89972 | ||
|
|
9de747d135 | ||
|
|
b313d78f60 | ||
|
|
e037d8d727 | ||
|
|
fade37025d | ||
|
|
f968f8e5a7 | ||
|
|
ac0014fe65 | ||
|
|
86282e83a8 | ||
|
|
37c7f5a6ba | ||
|
|
082d4ae62c | ||
|
|
c55a8fa4ba | ||
|
|
47617969d3 | ||
|
|
27f9750cdc | ||
|
|
427cdcd9d4 | ||
|
|
384338cf31 | ||
|
|
3b79e97f00 | ||
|
|
66a283b71b | ||
|
|
74f334cb93 | ||
|
|
8e9a6bf903 | ||
|
|
0ebd3c6230 | ||
|
|
6bda366065 | ||
|
|
1f36f5d401 | ||
|
|
a70a87f272 | ||
|
|
d2fbcb726d | ||
|
|
8169e177de | ||
|
|
589bade39a | ||
|
|
d1e8225d5f |
@@ -21,6 +21,8 @@ jobs:
|
|||||||
target: macos-x86_64
|
target: macos-x86_64
|
||||||
- os: ubuntu-22.04
|
- os: ubuntu-22.04
|
||||||
target: linux-x86_64
|
target: linux-x86_64
|
||||||
|
- os: ubuntu-22.04-arm
|
||||||
|
target: linux-arm64
|
||||||
- os: windows-latest
|
- os: windows-latest
|
||||||
target: windows-x86_64
|
target: windows-x86_64
|
||||||
|
|
||||||
@@ -43,6 +45,20 @@ jobs:
|
|||||||
uv sync --frozen
|
uv sync --frozen
|
||||||
uv run pyinstaller strix.spec --noconfirm
|
uv run pyinstaller strix.spec --noconfirm
|
||||||
|
|
||||||
|
if [[ "${{ runner.os }}" == "Windows" ]]; then
|
||||||
|
dist/strix.exe --version
|
||||||
|
else
|
||||||
|
dist/strix --version
|
||||||
|
fi
|
||||||
|
|
||||||
|
if [[ "${{ matrix.target }}" == "linux-arm64" ]]; then
|
||||||
|
file dist/strix
|
||||||
|
file dist/strix | grep -q "ARM aarch64" || {
|
||||||
|
echo "::error::linux-arm64 artifact is not an ARM aarch64 binary"
|
||||||
|
exit 1
|
||||||
|
}
|
||||||
|
fi
|
||||||
|
|
||||||
VERSION=$(grep '^version' pyproject.toml | head -1 | sed 's/.*"\(.*\)"/\1/')
|
VERSION=$(grep '^version' pyproject.toml | head -1 | sed 's/.*"\(.*\)"/\1/')
|
||||||
mkdir -p dist/release
|
mkdir -p dist/release
|
||||||
|
|
||||||
|
|||||||
+3
-3
@@ -1,8 +1,8 @@
|
|||||||
# Node / local-viewer SPA source (the built bundle in
|
# Node / local-viewer SPA source (the built bundle in
|
||||||
# strix/viewer/static/ is committed and shipped; do not ignore it)
|
# strix/interface/viewer/static/ is committed and shipped; do not ignore it)
|
||||||
node_modules/
|
node_modules/
|
||||||
strix/viewer/frontend/node_modules/
|
strix/interface/viewer/frontend/node_modules/
|
||||||
strix/viewer/frontend/.vite/
|
strix/interface/viewer/frontend/.vite/
|
||||||
|
|
||||||
# Python
|
# Python
|
||||||
__pycache__/
|
__pycache__/
|
||||||
|
|||||||
+5
-5
@@ -102,16 +102,16 @@ We welcome feature ideas! Please:
|
|||||||
## 🖥️ Local viewer SPA
|
## 🖥️ Local viewer SPA
|
||||||
|
|
||||||
`strix view` serves a prebuilt web UI whose source lives in
|
`strix view` serves a prebuilt web UI whose source lives in
|
||||||
`strix/viewer/frontend/` (a Vite + React project) and whose built output is
|
`strix/interface/viewer/frontend/` (a Vite + React project) and whose built output is
|
||||||
committed to `strix/viewer/static/` and shipped in the package. End users never
|
committed to `strix/interface/viewer/static/` and shipped in the package. End users never
|
||||||
run a JS build. If you change anything under `strix/viewer/frontend/`, rebuild
|
run a JS build. If you change anything under `strix/interface/viewer/frontend/`, rebuild
|
||||||
and commit the output:
|
and commit the output:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
make viewer # or: cd strix/viewer/frontend && npm ci && npm run build
|
make viewer # or: cd strix/interface/viewer/frontend && npm ci && npm run build
|
||||||
```
|
```
|
||||||
|
|
||||||
Commit both the source change and the regenerated `strix/viewer/static/`.
|
Commit both the source change and the regenerated `strix/interface/viewer/static/`.
|
||||||
|
|
||||||
## 🤝 Community
|
## 🤝 Community
|
||||||
|
|
||||||
|
|||||||
@@ -69,8 +69,8 @@ clean:
|
|||||||
|
|
||||||
viewer:
|
viewer:
|
||||||
@echo "🖥️ Building the local-viewer SPA..."
|
@echo "🖥️ Building the local-viewer SPA..."
|
||||||
cd strix/viewer/frontend && npm ci && npm run build
|
cd strix/interface/viewer/frontend && npm ci && npm run build
|
||||||
@echo "✅ Viewer built to strix/viewer/static/ (commit the changes)."
|
@echo "✅ Viewer built to strix/interface/viewer/static/ (commit the changes)."
|
||||||
|
|
||||||
dev: format lint type-check
|
dev: format lint type-check
|
||||||
@echo "✅ Development cycle complete!"
|
@echo "✅ Development cycle complete!"
|
||||||
|
|||||||
@@ -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`.
|
Custom API base URL. Also accepts `OPENAI_API_BASE`, `LITELLM_BASE_URL`, or `OLLAMA_API_BASE`.
|
||||||
</ParamField>
|
</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">
|
<ParamField path="LLM_TIMEOUT" default="300" type="integer">
|
||||||
Request timeout in seconds for LLM calls.
|
Request timeout in seconds for LLM calls.
|
||||||
</ParamField>
|
</ParamField>
|
||||||
@@ -55,6 +63,12 @@ affecting the agents that do the actual testing.
|
|||||||
model runs on a different endpoint than the main model.
|
model runs on a different endpoint than the main model.
|
||||||
</ParamField>
|
</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">
|
<ParamField path="STRIX_DEDUPE_REASONING_EFFORT" type="string">
|
||||||
Reasoning effort for the deduplication model. Defaults to the model's own
|
Reasoning effort for the deduplication model. Defaults to the model's own
|
||||||
baseline when unset.
|
baseline when unset.
|
||||||
|
|||||||
@@ -54,3 +54,20 @@ If you use LM Studio, vLLM, or other runners:
|
|||||||
export STRIX_LLM="openai/local-model"
|
export STRIX_LLM="openai/local-model"
|
||||||
export LLM_API_BASE="http://localhost:1234/v1" # Adjust port as needed
|
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.
|
||||||
|
|||||||
+36
-3
@@ -61,11 +61,28 @@ strix (--target <target> | --target-list <path> | --mount <path>) [options]
|
|||||||
Path to a custom config file (JSON) to use instead of `~/.strix/cli-config.json`.
|
Path to a custom config file (JSON) to use instead of `~/.strix/cli-config.json`.
|
||||||
</ParamField>
|
</ParamField>
|
||||||
|
|
||||||
<ParamField path="--max-budget-usd" type="number">
|
<ParamField path="--max-budget" type="number">
|
||||||
Maximum LLM spend in USD for the whole scan, counted cumulatively across the
|
Maximum LLM spend in USD for the whole scan, counted cumulatively across the
|
||||||
root agent and every child agent. The budget is checked after each model
|
root agent and every child agent. The budget is checked after each model
|
||||||
response; once the running cost reaches the threshold, the scan stops cleanly
|
response.
|
||||||
with a `stopped` status (not a failure) and the sandbox is torn down.
|
|
||||||
|
In non-interactive mode (`-n`), once the running cost reaches the threshold,
|
||||||
|
the scan stops cleanly with a `stopped` status (not a failure) and the sandbox
|
||||||
|
is torn down. Sub-agents are stopped early, at 90% of the budget, reserving
|
||||||
|
the final slice for the root agent to wind down and produce the final report.
|
||||||
|
|
||||||
|
In interactive mode, reaching the budget pauses the scan instead of ending
|
||||||
|
it: every agent parks, and sending any message resumes the scan with the cap
|
||||||
|
extended by the original budget amount. There is no sub-agent reserve in
|
||||||
|
interactive mode.
|
||||||
|
|
||||||
|
As the budget is approached, graduated wrap-up warnings are surfaced to
|
||||||
|
**every** agent so they can finish their work and call their lifecycle tool
|
||||||
|
before the hard stop. The bands sit just below each role's own stop point: the
|
||||||
|
root is warned at **70%, 85% and 95%** (it stops at 100%), while sub-agents are
|
||||||
|
warned at **75%, 80% and 85%** (they stop at the 90% reserve). In interactive
|
||||||
|
mode every agent uses the **70%, 85% and 95%** bands. Percentages shown in the
|
||||||
|
warnings are the real cumulative spend against the full budget.
|
||||||
|
|
||||||
Must be greater than `0`. Omit the flag for no limit.
|
Must be greater than `0`. Omit the flag for no limit.
|
||||||
|
|
||||||
@@ -84,6 +101,19 @@ strix (--target <target> | --target-list <path> | --mount <path>) [options]
|
|||||||
counts.
|
counts.
|
||||||
</ParamField>
|
</ParamField>
|
||||||
|
|
||||||
|
<ParamField path="--max-turns" type="integer" default="500">
|
||||||
|
Maximum number of turns (one model response plus its tool round) allotted to
|
||||||
|
**each** agent, applied per run. When an agent reaches this limit it is
|
||||||
|
force-stopped.
|
||||||
|
|
||||||
|
As the limit is approached, graduated wrap-up warnings (at 70%, 85% and 95%)
|
||||||
|
are injected into that agent's next model turn so it can prioritise its
|
||||||
|
remaining work and call its lifecycle tool (`finish_scan` for the root agent,
|
||||||
|
`agent_finish` for sub-agents) before the hard stop.
|
||||||
|
|
||||||
|
Must be greater than `0`.
|
||||||
|
</ParamField>
|
||||||
|
|
||||||
## Examples
|
## Examples
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
@@ -99,6 +129,9 @@ strix --target api.example.com --instruction "Focus on IDOR and auth bypass"
|
|||||||
# CI/CD mode
|
# CI/CD mode
|
||||||
strix -n --target ./ --scan-mode quick
|
strix -n --target ./ --scan-mode quick
|
||||||
|
|
||||||
|
# Cap cost and per-agent turns
|
||||||
|
strix --target https://example.com --max-budget 25 --max-turns 300
|
||||||
|
|
||||||
# Force diff-scope against a specific base ref
|
# Force diff-scope against a specific base ref
|
||||||
strix -n --target ./ --scan-mode quick --scope-mode diff --diff-base origin/main
|
strix -n --target ./ --scan-mode quick --scope-mode diff --diff-base origin/main
|
||||||
|
|
||||||
|
|||||||
+8
-7
@@ -1,6 +1,6 @@
|
|||||||
[project]
|
[project]
|
||||||
name = "strix-agent"
|
name = "strix-agent"
|
||||||
version = "1.3.1"
|
version = "1.4.1"
|
||||||
description = "Open-source AI Hackers for your apps"
|
description = "Open-source AI Hackers for your apps"
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
license = "Apache-2.0"
|
license = "Apache-2.0"
|
||||||
@@ -79,10 +79,10 @@ build-backend = "hatchling.build"
|
|||||||
|
|
||||||
[tool.hatch.build.targets.wheel]
|
[tool.hatch.build.targets.wheel]
|
||||||
packages = ["strix"]
|
packages = ["strix"]
|
||||||
# The prebuilt viewer bundle under strix/viewer/static/ ships automatically
|
# The prebuilt viewer bundle under strix/interface/viewer/static/ ships automatically
|
||||||
# (hatchling includes non-.py files under the package). The Vite SOURCE lives
|
# (hatchling includes non-.py files under the package). The Vite SOURCE lives
|
||||||
# under the package dir too (strix/viewer/frontend/) but must never ship in the wheel.
|
# under the package dir too (strix/interface/viewer/frontend/) but must never ship in the wheel.
|
||||||
exclude = ["strix/viewer/frontend", "strix/viewer/frontend/**"]
|
exclude = ["strix/interface/viewer/frontend", "strix/interface/viewer/frontend/**"]
|
||||||
|
|
||||||
# ============================================================================
|
# ============================================================================
|
||||||
# Type Checking Configuration
|
# Type Checking Configuration
|
||||||
@@ -220,12 +220,13 @@ ignore = [
|
|||||||
# Stdlib HTTP handler overrides (do_GET/do_POST).
|
# Stdlib HTTP handler overrides (do_GET/do_POST).
|
||||||
"strix/interface/auth_cli.py" = ["N802"]
|
"strix/interface/auth_cli.py" = ["N802"]
|
||||||
"tests/test_codex_streaming.py" = ["N802"]
|
"tests/test_codex_streaming.py" = ["N802"]
|
||||||
|
"tests/test_disable_streaming.py" = ["N802"]
|
||||||
"tests/test_report_pdf.py" = ["S105", "S106"]
|
"tests/test_report_pdf.py" = ["S105", "S106"]
|
||||||
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
|
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
|
||||||
# circular dependency with strix.telemetry / strix.viewer.report_pdf.
|
# circular dependency with strix.telemetry / strix.interface.viewer.report_pdf.
|
||||||
"strix/viewer/server.py" = ["N802", "PLC0415"]
|
"strix/interface/viewer/server.py" = ["N802", "PLC0415"]
|
||||||
# Lazy telemetry import to avoid importing PostHog before the viewer starts.
|
# Lazy telemetry import to avoid importing PostHog before the viewer starts.
|
||||||
"strix/viewer/cli.py" = ["PLC0415"]
|
"strix/interface/viewer/cli.py" = ["PLC0415"]
|
||||||
# Lazy imports inside functions to avoid circular dependency with
|
# Lazy imports inside functions to avoid circular dependency with
|
||||||
# strix.telemetry / strix.report.dedupe / cvss.
|
# strix.telemetry / strix.report.dedupe / cvss.
|
||||||
"strix/tools/notes/tools.py" = ["PLC0415", "TC002"]
|
"strix/tools/notes/tools.py" = ["PLC0415", "TC002"]
|
||||||
|
|||||||
+1
-1
@@ -41,7 +41,7 @@ fi
|
|||||||
|
|
||||||
combo="$os-$arch"
|
combo="$os-$arch"
|
||||||
case "$combo" in
|
case "$combo" in
|
||||||
linux-x86_64|macos-x86_64|macos-arm64|windows-x86_64)
|
linux-x86_64|linux-arm64|macos-x86_64|macos-arm64|windows-x86_64)
|
||||||
;;
|
;;
|
||||||
*)
|
*)
|
||||||
echo -e "${RED}Unsupported OS/Arch: $os/$arch${NC}"
|
echo -e "${RED}Unsupported OS/Arch: $os/$arch${NC}"
|
||||||
|
|||||||
+7
-7
@@ -26,7 +26,7 @@ for tcss_file in strix_root.rglob('*.tcss'):
|
|||||||
datas.append((str(tcss_file), str(rel_path.parent)))
|
datas.append((str(tcss_file), str(rel_path.parent)))
|
||||||
|
|
||||||
# Prebuilt local-viewer SPA (served by `strix view`).
|
# Prebuilt local-viewer SPA (served by `strix view`).
|
||||||
viewer_static = strix_root / 'viewer' / 'static'
|
viewer_static = strix_root / 'interface' / 'viewer' / 'static'
|
||||||
for asset in viewer_static.rglob('*'):
|
for asset in viewer_static.rglob('*'):
|
||||||
if asset.is_file():
|
if asset.is_file():
|
||||||
rel_path = asset.relative_to(project_root)
|
rel_path = asset.relative_to(project_root)
|
||||||
@@ -158,12 +158,12 @@ hiddenimports = [
|
|||||||
'strix.report.dedupe',
|
'strix.report.dedupe',
|
||||||
'strix.report.state',
|
'strix.report.state',
|
||||||
'strix.report.writer',
|
'strix.report.writer',
|
||||||
'strix.viewer',
|
'strix.interface.viewer',
|
||||||
'strix.viewer.auth',
|
'strix.interface.viewer.auth',
|
||||||
'strix.viewer.cli',
|
'strix.interface.viewer.cli',
|
||||||
'strix.viewer.report_pdf',
|
'strix.interface.viewer.report_pdf',
|
||||||
'strix.viewer.server',
|
'strix.interface.viewer.server',
|
||||||
'strix.viewer.transcript',
|
'strix.interface.viewer.transcript',
|
||||||
|
|
||||||
# PDF report generation + encryption
|
# PDF report generation + encryption
|
||||||
'reportlab',
|
'reportlab',
|
||||||
|
|||||||
+91
-14
@@ -16,6 +16,7 @@ from agents.tool import CustomTool, FunctionTool, Tool
|
|||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
|
|
||||||
from strix.agents.prompt import render_system_prompt
|
from strix.agents.prompt import render_system_prompt
|
||||||
|
from strix.config import load_settings
|
||||||
from strix.tools.agents_graph.tools import (
|
from strix.tools.agents_graph.tools import (
|
||||||
agent_finish,
|
agent_finish,
|
||||||
create_agent,
|
create_agent,
|
||||||
@@ -33,6 +34,7 @@ from strix.tools.notes.tools import (
|
|||||||
list_notes,
|
list_notes,
|
||||||
update_note,
|
update_note,
|
||||||
)
|
)
|
||||||
|
from strix.tools.output_store import bound_and_store, bound_text
|
||||||
from strix.tools.proxy.tools import (
|
from strix.tools.proxy.tools import (
|
||||||
list_requests,
|
list_requests,
|
||||||
list_sitemap,
|
list_sitemap,
|
||||||
@@ -41,7 +43,12 @@ from strix.tools.proxy.tools import (
|
|||||||
view_request,
|
view_request,
|
||||||
view_sitemap_entry,
|
view_sitemap_entry,
|
||||||
)
|
)
|
||||||
from strix.tools.reporting.tool import create_dependency_report, create_vulnerability_report
|
from strix.tools.reporting.tool import (
|
||||||
|
create_dependency_report,
|
||||||
|
create_vulnerability_report,
|
||||||
|
get_report,
|
||||||
|
list_reports,
|
||||||
|
)
|
||||||
from strix.tools.thinking.tool import think
|
from strix.tools.thinking.tool import think
|
||||||
from strix.tools.todo.tools import (
|
from strix.tools.todo.tools import (
|
||||||
create_todo,
|
create_todo,
|
||||||
@@ -103,8 +110,36 @@ def _extract_custom_input(tool: CustomTool, raw_input: str | dict[str, Any]) ->
|
|||||||
return value if isinstance(value, str) else ""
|
return value if isinstance(value, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_output_limits() -> tuple[int, int]:
|
||||||
|
context = load_settings().context
|
||||||
|
return context.tool_output_max_lines, context.tool_output_max_bytes
|
||||||
|
|
||||||
|
|
||||||
|
async def _bound_result(result: Any) -> Any:
|
||||||
|
if not isinstance(result, str):
|
||||||
|
return result
|
||||||
|
max_lines, max_bytes = _tool_output_limits()
|
||||||
|
return await bound_and_store(result, max_lines=max_lines, max_bytes=max_bytes)
|
||||||
|
|
||||||
|
|
||||||
def _format_tool_error(exc: Exception) -> str:
|
def _format_tool_error(exc: Exception) -> str:
|
||||||
return str(exc) or exc.__class__.__name__
|
message = str(exc) or exc.__class__.__name__
|
||||||
|
max_lines, max_bytes = _tool_output_limits()
|
||||||
|
return bound_text(message, max_lines=max_lines, max_bytes=max_bytes)
|
||||||
|
|
||||||
|
|
||||||
|
def _with_bounded_result(tool: FunctionTool) -> FunctionTool:
|
||||||
|
"""Cap a tool's result size before it enters history (idempotent)."""
|
||||||
|
if getattr(tool, "_strix_bounded", False):
|
||||||
|
return tool
|
||||||
|
invoke_tool = tool.on_invoke_tool
|
||||||
|
|
||||||
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
|
return await _bound_result(await invoke_tool(ctx, raw_input))
|
||||||
|
|
||||||
|
tool.on_invoke_tool = invoke
|
||||||
|
tool._strix_bounded = True # type: ignore[attr-defined]
|
||||||
|
return tool
|
||||||
|
|
||||||
|
|
||||||
def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
||||||
@@ -112,7 +147,7 @@ def _function_tool_with_error_result(tool: FunctionTool) -> FunctionTool:
|
|||||||
|
|
||||||
async def invoke(ctx: Any, raw_input: str) -> Any:
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
try:
|
try:
|
||||||
return await invoke_tool(ctx, raw_input)
|
return await _bound_result(await invoke_tool(ctx, raw_input))
|
||||||
except Exception as exc: # noqa: BLE001 - tool errors should be model-visible results.
|
except Exception as exc: # noqa: BLE001 - tool errors should be model-visible results.
|
||||||
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
||||||
return _format_tool_error(exc)
|
return _format_tool_error(exc)
|
||||||
@@ -127,7 +162,7 @@ def _custom_tool_as_function_tool(tool: CustomTool) -> FunctionTool:
|
|||||||
if not custom_input:
|
if not custom_input:
|
||||||
return f"`{_custom_tool_input_field(tool)}` must be a non-empty string."
|
return f"`{_custom_tool_input_field(tool)}` must be a non-empty string."
|
||||||
try:
|
try:
|
||||||
return await tool.on_invoke_tool(ctx, custom_input)
|
return await _bound_result(await tool.on_invoke_tool(ctx, custom_input))
|
||||||
except Exception as exc: # noqa: BLE001 - matches SDK CustomTool error-as-result behavior.
|
except Exception as exc: # noqa: BLE001 - matches SDK CustomTool error-as-result behavior.
|
||||||
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
logger.debug("Tool %s failed; returning error as result", tool.name, exc_info=True)
|
||||||
return _format_tool_error(exc)
|
return _format_tool_error(exc)
|
||||||
@@ -159,12 +194,35 @@ def _custom_tool_as_function_tool(tool: CustomTool) -> FunctionTool:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _configure_chat_completions_filesystem_tools(toolset: Any) -> None:
|
def _bound_custom_tool(tool: CustomTool) -> CustomTool:
|
||||||
|
"""Bound a native ``CustomTool`` result in place (Responses path)."""
|
||||||
|
invoke_tool = tool.on_invoke_tool
|
||||||
|
|
||||||
|
async def invoke(ctx: Any, raw_input: str) -> Any:
|
||||||
|
return await _bound_result(await invoke_tool(ctx, raw_input))
|
||||||
|
|
||||||
|
tool.on_invoke_tool = invoke
|
||||||
|
return tool
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_filesystem_tools(toolset: Any, *, chat_completions: bool) -> None:
|
||||||
for name, tool in vars(toolset).items():
|
for name, tool in vars(toolset).items():
|
||||||
if isinstance(tool, CustomTool):
|
if chat_completions:
|
||||||
setattr(toolset, name, _custom_tool_as_function_tool(tool))
|
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))
|
||||||
|
elif isinstance(tool, CustomTool):
|
||||||
|
setattr(toolset, name, _bound_custom_tool(tool))
|
||||||
elif isinstance(tool, FunctionTool):
|
elif isinstance(tool, FunctionTool):
|
||||||
setattr(toolset, name, _function_tool_with_error_result(tool))
|
setattr(toolset, name, _with_bounded_result(tool))
|
||||||
|
|
||||||
|
|
||||||
|
def _make_filesystem_configurator(*, chat_completions: bool) -> Any:
|
||||||
|
def configure(toolset: Any) -> None:
|
||||||
|
_configure_filesystem_tools(toolset, chat_completions=chat_completions)
|
||||||
|
|
||||||
|
return configure
|
||||||
|
|
||||||
|
|
||||||
_CHARS_ESCAPE_RE = re.compile(r"\\(?:u[0-9a-fA-F]{4}|x[0-9a-fA-F]{2}|[0abtnvfr\\])")
|
_CHARS_ESCAPE_RE = re.compile(r"\\(?:u[0-9a-fA-F]{4}|x[0-9a-fA-F]{2}|[0abtnvfr\\])")
|
||||||
@@ -205,6 +263,16 @@ def _format_validation_error(tool_name: str, exc: ValidationError) -> str:
|
|||||||
return f"{tool_name}: invalid arguments — " + "; ".join(parts)
|
return f"{tool_name}: invalid arguments — " + "; ".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_shell_output_cap(parsed: dict[str, Any]) -> None:
|
||||||
|
"""Clamp the SDK shell tools' ``max_output_tokens`` to the configured
|
||||||
|
ceiling; a smaller explicit value is respected."""
|
||||||
|
ceiling = load_settings().context.tool_output_max_tokens
|
||||||
|
requested = parsed.get("max_output_tokens")
|
||||||
|
parsed["max_output_tokens"] = (
|
||||||
|
ceiling if not isinstance(requested, int) or requested > ceiling else requested
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _wrap_exec_command(tool: FunctionTool) -> FunctionTool:
|
def _wrap_exec_command(tool: FunctionTool) -> FunctionTool:
|
||||||
invoke_tool = tool.on_invoke_tool
|
invoke_tool = tool.on_invoke_tool
|
||||||
|
|
||||||
@@ -213,8 +281,10 @@ def _wrap_exec_command(tool: FunctionTool) -> FunctionTool:
|
|||||||
parsed = json.loads(raw_input)
|
parsed = json.loads(raw_input)
|
||||||
except (json.JSONDecodeError, TypeError):
|
except (json.JSONDecodeError, TypeError):
|
||||||
parsed = None
|
parsed = None
|
||||||
if isinstance(parsed, dict) and "shell" not in parsed:
|
if isinstance(parsed, dict):
|
||||||
parsed["shell"] = "bash"
|
if "shell" not in parsed:
|
||||||
|
parsed["shell"] = "bash"
|
||||||
|
_apply_shell_output_cap(parsed)
|
||||||
raw_input = json.dumps(parsed)
|
raw_input = json.dumps(parsed)
|
||||||
try:
|
try:
|
||||||
return await invoke_tool(ctx, raw_input)
|
return await invoke_tool(ctx, raw_input)
|
||||||
@@ -240,8 +310,10 @@ def _wrap_write_stdin(tool: FunctionTool) -> FunctionTool:
|
|||||||
parsed = json.loads(raw_input)
|
parsed = json.loads(raw_input)
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
parsed = None
|
parsed = None
|
||||||
if isinstance(parsed, dict) and isinstance(parsed.get("chars"), str):
|
if isinstance(parsed, dict):
|
||||||
parsed["chars"] = _decode_chars_escape(parsed["chars"])
|
if isinstance(parsed.get("chars"), str):
|
||||||
|
parsed["chars"] = _decode_chars_escape(parsed["chars"])
|
||||||
|
_apply_shell_output_cap(parsed)
|
||||||
raw_input = json.dumps(parsed)
|
raw_input = json.dumps(parsed)
|
||||||
try:
|
try:
|
||||||
return await invoke_tool(ctx, raw_input)
|
return await invoke_tool(ctx, raw_input)
|
||||||
@@ -343,6 +415,8 @@ _BASE_TOOLS: tuple[Tool, ...] = (
|
|||||||
web_search,
|
web_search,
|
||||||
create_vulnerability_report,
|
create_vulnerability_report,
|
||||||
create_dependency_report,
|
create_dependency_report,
|
||||||
|
list_reports,
|
||||||
|
get_report,
|
||||||
list_requests,
|
list_requests,
|
||||||
view_request,
|
view_request,
|
||||||
repeat_request,
|
repeat_request,
|
||||||
@@ -440,6 +514,9 @@ def build_strix_agent(
|
|||||||
else:
|
else:
|
||||||
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
|
tools = [*_BASE_TOOLS, *agent_tools, agent_finish]
|
||||||
_ensure_unique_tool_names(tools)
|
_ensure_unique_tool_names(tools)
|
||||||
|
tools = [
|
||||||
|
_with_bounded_result(tool) if isinstance(tool, FunctionTool) else tool for tool in tools
|
||||||
|
]
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Built %s agent '%s' (skills=%d, tools=%d, scan_mode=%s, whitebox=%s)",
|
"Built %s agent '%s' (skills=%d, tools=%d, scan_mode=%s, whitebox=%s)",
|
||||||
@@ -459,8 +536,8 @@ def build_strix_agent(
|
|||||||
model=None,
|
model=None,
|
||||||
capabilities=[
|
capabilities=[
|
||||||
Filesystem(
|
Filesystem(
|
||||||
configure_tools=(
|
configure_tools=_make_filesystem_configurator(
|
||||||
_configure_chat_completions_filesystem_tools if chat_completions_tools else None
|
chat_completions=chat_completions_tools,
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
Shell(
|
Shell(
|
||||||
|
|||||||
@@ -188,7 +188,7 @@ EFFICIENCY TACTICS:
|
|||||||
script fail with `ModuleNotFoundError`.
|
script fail with `ModuleNotFoundError`.
|
||||||
- `exec_command` runs each command in a fresh non-interactive shell (plain
|
- `exec_command` runs each command in a fresh non-interactive shell (plain
|
||||||
pipes, no TTY). To drive an interactive or long-running process with
|
pipes, no TTY). To drive an interactive or long-running process with
|
||||||
`write_stdin` — REPLs, `ssh`/`nc`/`ftp`, `msfconsole`, or to send Ctrl-C —
|
`write_stdin` — REPLs, `ssh`/`nc`/`ftp`, `sqlmap`, or to send Ctrl-C —
|
||||||
you MUST start it with `exec_command(cmd="...", tty=true)` and then
|
you MUST start it with `exec_command(cmd="...", tty=true)` and then
|
||||||
`write_stdin(session_id=<id>, chars="...")`. Calling `write_stdin` on a
|
`write_stdin(session_id=<id>, chars="...")`. Calling `write_stdin` on a
|
||||||
default (non-TTY) command or on a process that has already exited fails with
|
default (non-TTY) command or on a process that has already exited fails with
|
||||||
@@ -215,6 +215,7 @@ VALIDATION REQUIREMENTS:
|
|||||||
- A vulnerability is ONLY considered reported when a reporting agent uses create_vulnerability_report (or create_dependency_report for known-CVE dependency/supply-chain findings) with full details. Mentions in agent_finish, finish_scan, or generic messages are NOT sufficient
|
- A vulnerability is ONLY considered reported when a reporting agent uses create_vulnerability_report (or create_dependency_report for known-CVE dependency/supply-chain findings) with full details. Mentions in agent_finish, finish_scan, or generic messages are NOT sufficient
|
||||||
- Reporting and fixing are ONE step, not two: when source is available, the reporting agent derives the concrete fix and files it INLINE via create_vulnerability_report (`code_locations` with `fix_before`/`fix_after` + `fix_pr_body`) — the report is not complete without it. Do NOT report first and then spawn a separate downstream agent to re-derive and re-apply the same patch; that just re-does the analysis and wastes tokens. (Do not silently patch a finding WITHOUT filing a report — the report, with its embedded fix, is the deliverable.)
|
- Reporting and fixing are ONE step, not two: when source is available, the reporting agent derives the concrete fix and files it INLINE via create_vulnerability_report (`code_locations` with `fix_before`/`fix_after` + `fix_pr_body`) — the report is not complete without it. Do NOT report first and then spawn a separate downstream agent to re-derive and re-apply the same patch; that just re-does the analysis and wastes tokens. (Do not silently patch a finding WITHOUT filing a report — the report, with its embedded fix, is the deliverable.)
|
||||||
- DEDUPLICATION: The create_vulnerability_report tool uses LLM-based deduplication. If it rejects your report as a duplicate, DO NOT attempt to re-submit the same vulnerability. Accept the rejection and move on to testing other areas. The vulnerability has already been reported by another agent
|
- DEDUPLICATION: The create_vulnerability_report tool uses LLM-based deduplication. If it rejects your report as a duplicate, DO NOT attempt to re-submit the same vulnerability. Accept the rejection and move on to testing other areas. The vulnerability has already been reported by another agent
|
||||||
|
- REVIEWING FILED FINDINGS (orchestrator/root agent): use list_reports to see every vulnerability filed so far in this scan (by any agent, root or child) — metadata-first with per-severity counts — and get_report to read one finding in full by its id. These are read-only orchestration tools: the root agent uses them to track coverage, avoid dispatching work on already-covered ground, assemble the finish_scan executive summary, and reason about attack-chaining across confirmed findings. Leaf/specialist agents should NOT call them — just do your assigned testing and file findings. Each entry shows which agent filed it (agent_name), and your own entries are flagged by_you. list_notes/get_note do the same for notes.
|
||||||
</execution_guidelines>
|
</execution_guidelines>
|
||||||
|
|
||||||
<vulnerability_focus>
|
<vulnerability_focus>
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from strix.config.loader import (
|
|||||||
persist_current,
|
persist_current,
|
||||||
)
|
)
|
||||||
from strix.config.settings import (
|
from strix.config.settings import (
|
||||||
|
ContextSettings,
|
||||||
DedupeSettings,
|
DedupeSettings,
|
||||||
IntegrationSettings,
|
IntegrationSettings,
|
||||||
LlmSettings,
|
LlmSettings,
|
||||||
@@ -27,6 +28,7 @@ from strix.config.settings import (
|
|||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"ContextSettings",
|
||||||
"DedupeSettings",
|
"DedupeSettings",
|
||||||
"IntegrationSettings",
|
"IntegrationSettings",
|
||||||
"LlmSettings",
|
"LlmSettings",
|
||||||
|
|||||||
+13
-20
@@ -18,12 +18,12 @@ import logging
|
|||||||
import secrets
|
import secrets
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
import urllib.error
|
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
import urllib.request
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import Iterator
|
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]:
|
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:
|
try:
|
||||||
with urllib.request.urlopen( # noqa: S310 # nosec B310 - fixed https endpoint
|
response = requests.post(
|
||||||
request, timeout=_TOKEN_TIMEOUT
|
TOKEN_URL,
|
||||||
) as response:
|
data=payload,
|
||||||
data = json.loads(response.read() or b"{}")
|
headers={"Accept": "application/json"},
|
||||||
except urllib.error.HTTPError as exc:
|
timeout=_TOKEN_TIMEOUT,
|
||||||
detail = exc.read().decode("utf-8", "replace")[:300]
|
)
|
||||||
raise CodexAuthError("token_http_error", f"HTTP {exc.code}: {detail}") from exc
|
except requests.RequestException as exc:
|
||||||
except (urllib.error.URLError, TimeoutError, OSError) as exc:
|
|
||||||
raise CodexAuthError("unavailable", str(exc)) from 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):
|
if not isinstance(data, dict):
|
||||||
raise CodexAuthError("bad_response", "token endpoint returned non-object")
|
raise CodexAuthError("bad_response", "token endpoint returned non-object")
|
||||||
return data
|
return data
|
||||||
|
|||||||
+266
-4
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|||||||
import contextlib
|
import contextlib
|
||||||
import inspect
|
import inspect
|
||||||
import os
|
import os
|
||||||
|
import time
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from agents import (
|
from agents import (
|
||||||
@@ -13,6 +14,8 @@ from agents import (
|
|||||||
set_tracing_disabled,
|
set_tracing_disabled,
|
||||||
)
|
)
|
||||||
from agents.model_settings import ModelSettings
|
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.multi_provider import MultiProvider
|
||||||
from agents.models.openai_responses import OpenAIResponsesModel
|
from agents.models.openai_responses import OpenAIResponsesModel
|
||||||
from agents.retry import (
|
from agents.retry import (
|
||||||
@@ -21,6 +24,8 @@ from agents.retry import (
|
|||||||
RetryPolicyContext,
|
RetryPolicyContext,
|
||||||
retry_policies,
|
retry_policies,
|
||||||
)
|
)
|
||||||
|
from openai.types.responses import Response, ResponseCompletedEvent
|
||||||
|
from openai.types.responses.response_usage import ResponseUsage
|
||||||
from openai.types.shared import Reasoning
|
from openai.types.shared import Reasoning
|
||||||
|
|
||||||
from strix.config import codex
|
from strix.config import codex
|
||||||
@@ -30,10 +35,17 @@ from strix.config.loader import load_settings
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from collections.abc import AsyncIterator
|
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 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:
|
def request_timeout_extra_args(timeout_s: float | None) -> dict[str, float] | None:
|
||||||
@@ -135,6 +147,124 @@ class _CodexResponsesModel(OpenAIResponsesModel):
|
|||||||
await result
|
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):
|
class StrixProvider(MultiProvider):
|
||||||
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
|
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
|
||||||
so users type ``deepseek/deepseek-chat`` rather than
|
so users type ``deepseek/deepseek-chat`` rather than
|
||||||
@@ -159,14 +289,21 @@ class StrixProvider(MultiProvider):
|
|||||||
return self._get_fallback_provider("litellm"), original_model_name
|
return self._get_fallback_provider("litellm"), original_model_name
|
||||||
|
|
||||||
def get_model(self, model_name: str | None) -> Model:
|
def get_model(self, model_name: str | None) -> Model:
|
||||||
|
llm = load_settings().llm
|
||||||
slug = codex.subscription_model(model_name)
|
slug = codex.subscription_model(model_name)
|
||||||
if slug:
|
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(
|
return _CodexResponsesModel(
|
||||||
slug,
|
slug,
|
||||||
codex.get_subscription_client(),
|
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(
|
DEFAULT_MODEL_RETRY = ModelRetrySettings(
|
||||||
@@ -243,6 +380,7 @@ def configure_sdk_model_defaults(settings: Settings) -> None:
|
|||||||
set_default_openai_api("chat_completions")
|
set_default_openai_api("chat_completions")
|
||||||
else:
|
else:
|
||||||
set_default_openai_api("responses")
|
set_default_openai_api("responses")
|
||||||
|
_configure_extra_headers(llm)
|
||||||
|
|
||||||
|
|
||||||
def _mirror_api_key_to_provider_env(model_name: str | None, api_key: str) -> None:
|
def _mirror_api_key_to_provider_env(model_name: str | None, api_key: str) -> None:
|
||||||
@@ -277,6 +415,51 @@ def _configure_litellm_compatibility() -> None:
|
|||||||
litellm.suppress_debug_info = True
|
litellm.suppress_debug_info = True
|
||||||
|
|
||||||
_register_litellm_cost_callback()
|
_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 = {
|
_OPENROUTER_ATTRIBUTION_HEADERS = {
|
||||||
@@ -302,6 +485,43 @@ def _configure_openrouter_attribution(model_name: str | None) -> None:
|
|||||||
litellm.headers = {**existing, **_OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment]
|
litellm.headers = {**existing, **_OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment]
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_extra_headers(llm: LlmSettings) -> None:
|
||||||
|
"""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:
|
def _register_litellm_cost_callback() -> None:
|
||||||
import litellm
|
import litellm
|
||||||
|
|
||||||
@@ -429,3 +649,45 @@ def is_known_openai_bare_model(model_name: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
entry = litellm.model_cost.get(name)
|
entry = litellm.model_cost.get(name)
|
||||||
return bool(entry and entry.get("litellm_provider") == "openai")
|
return bool(entry and entry.get("litellm_provider") == "openai")
|
||||||
|
|
||||||
|
|
||||||
|
def is_claude_model(model_name: str) -> bool:
|
||||||
|
return "claude" in (model_name or "").strip().lower()
|
||||||
|
|
||||||
|
|
||||||
|
def is_bedrock_route(model_name: str) -> bool:
|
||||||
|
name = (model_name or "").strip().lower()
|
||||||
|
return name.startswith("bedrock/") or "anthropic." in name
|
||||||
|
|
||||||
|
|
||||||
|
def _prompt_cache_name_candidates(model_name: str) -> list[str]:
|
||||||
|
# LiteLLM's model map keys the same model under several names; strip the
|
||||||
|
# route prefix, then leading dotted segments (region, provider).
|
||||||
|
name = (model_name or "").strip().lower()
|
||||||
|
for prefix in ("litellm/", "bedrock/"):
|
||||||
|
if name.startswith(prefix):
|
||||||
|
name = name[len(prefix) :]
|
||||||
|
break
|
||||||
|
candidates = [name]
|
||||||
|
rest = name
|
||||||
|
while "." in rest:
|
||||||
|
rest = rest.split(".", 1)[1]
|
||||||
|
candidates.append(rest)
|
||||||
|
return candidates
|
||||||
|
|
||||||
|
|
||||||
|
def bedrock_route_supports_prompt_caching(model_name: str) -> bool:
|
||||||
|
# Bedrock rejects the cache marker for models LiteLLM's map doesn't
|
||||||
|
# recognise as cache-capable, so callers withhold it unless confirmed here.
|
||||||
|
import litellm
|
||||||
|
|
||||||
|
checker = getattr(getattr(litellm, "utils", None), "supports_prompt_caching", None)
|
||||||
|
for cand in _prompt_cache_name_candidates(model_name):
|
||||||
|
if checker is not None:
|
||||||
|
with contextlib.suppress(Exception):
|
||||||
|
if checker(cand):
|
||||||
|
return True
|
||||||
|
entry = litellm.model_cost.get(cand)
|
||||||
|
if entry and entry.get("supports_prompt_caching"):
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|||||||
@@ -35,11 +35,23 @@ class LlmSettings(BaseSettings):
|
|||||||
"OLLAMA_API_BASE",
|
"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")
|
reasoning_effort: ReasoningEffort = Field(default="high", alias="STRIX_REASONING_EFFORT")
|
||||||
force_required_tool_choice: bool = Field(
|
force_required_tool_choice: bool = Field(
|
||||||
default=False,
|
default=False,
|
||||||
alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE",
|
alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE",
|
||||||
)
|
)
|
||||||
|
prompt_cache: bool = Field(
|
||||||
|
default=True,
|
||||||
|
alias="STRIX_PROMPT_CACHE",
|
||||||
|
)
|
||||||
|
disable_streaming: bool = Field(
|
||||||
|
default=False,
|
||||||
|
alias="LLM_DISABLE_STREAMING",
|
||||||
|
)
|
||||||
timeout: int = Field(default=300, alias="LLM_TIMEOUT")
|
timeout: int = Field(default=300, alias="LLM_TIMEOUT")
|
||||||
|
|
||||||
|
|
||||||
@@ -53,6 +65,30 @@ class DedupeSettings(BaseSettings):
|
|||||||
)
|
)
|
||||||
api_key: str | None = Field(default=None, alias="DEDUPE_LLM_API_KEY")
|
api_key: str | None = Field(default=None, alias="DEDUPE_LLM_API_KEY")
|
||||||
api_base: str | None = Field(default=None, alias="DEDUPE_LLM_API_BASE")
|
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):
|
||||||
|
"""Context-window management: per-tool-output caps and history compaction."""
|
||||||
|
|
||||||
|
model_config = _BASE_CONFIG
|
||||||
|
|
||||||
|
auto_compact: bool = Field(default=True, alias="STRIX_CONTEXT_AUTO_COMPACT")
|
||||||
|
compact_buffer_tokens: int = Field(default=20_000, gt=0, alias="STRIX_CONTEXT_BUFFER_TOKENS")
|
||||||
|
keep_tokens: int = Field(default=8_000, gt=0, alias="STRIX_CONTEXT_KEEP_TOKENS")
|
||||||
|
fallback_context_tokens: int = Field(
|
||||||
|
default=200_000, gt=0, alias="STRIX_CONTEXT_FALLBACK_TOKENS"
|
||||||
|
)
|
||||||
|
summary_max_tokens: int = Field(default=4_096, gt=0, alias="STRIX_CONTEXT_SUMMARY_TOKENS")
|
||||||
|
tool_output_max_tokens: int = Field(default=8_000, gt=0, alias="STRIX_TOOL_OUTPUT_MAX_TOKENS")
|
||||||
|
tool_output_max_lines: int = Field(default=2_000, gt=0, alias="STRIX_TOOL_OUTPUT_MAX_LINES")
|
||||||
|
# Floor above the truncation-notice size so a preview always fits.
|
||||||
|
tool_output_max_bytes: int = Field(
|
||||||
|
default=50 * 1024, ge=1024, alias="STRIX_TOOL_OUTPUT_MAX_BYTES"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class RuntimeSettings(BaseSettings):
|
class RuntimeSettings(BaseSettings):
|
||||||
@@ -99,6 +135,7 @@ class Settings(BaseSettings):
|
|||||||
llm: LlmSettings = Field(default_factory=LlmSettings)
|
llm: LlmSettings = Field(default_factory=LlmSettings)
|
||||||
dedupe: DedupeSettings = Field(default_factory=DedupeSettings)
|
dedupe: DedupeSettings = Field(default_factory=DedupeSettings)
|
||||||
runtime: RuntimeSettings = Field(default_factory=RuntimeSettings)
|
runtime: RuntimeSettings = Field(default_factory=RuntimeSettings)
|
||||||
|
context: ContextSettings = Field(default_factory=ContextSettings)
|
||||||
telemetry: TelemetrySettings = Field(default_factory=TelemetrySettings)
|
telemetry: TelemetrySettings = Field(default_factory=TelemetrySettings)
|
||||||
integrations: IntegrationSettings = Field(default_factory=IntegrationSettings)
|
integrations: IntegrationSettings = Field(default_factory=IntegrationSettings)
|
||||||
viewer: ViewerSettings = Field(default_factory=ViewerSettings)
|
viewer: ViewerSettings = Field(default_factory=ViewerSettings)
|
||||||
|
|||||||
+86
-5
@@ -14,13 +14,15 @@ from strix.core.sessions import session_write_lock
|
|||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from collections.abc import Callable
|
||||||
|
|
||||||
from agents.items import TResponseInputItem
|
from agents.items import TResponseInputItem
|
||||||
from agents.memory import Session
|
from agents.memory import Session
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed"]
|
Status = Literal["running", "waiting", "completed", "stopped", "crashed", "failed", "budget_paused"]
|
||||||
|
|
||||||
|
|
||||||
@dataclass(slots=True)
|
@dataclass(slots=True)
|
||||||
@@ -47,6 +49,9 @@ class AgentCoordinator:
|
|||||||
self._snapshot_path: Path | None = None
|
self._snapshot_path: Path | None = None
|
||||||
self.is_shutting_down = False
|
self.is_shutting_down = False
|
||||||
self._budget_stopped = False
|
self._budget_stopped = False
|
||||||
|
self._reserve_stopped = False
|
||||||
|
self._budget_paused = False
|
||||||
|
self._extend_budget: Callable[[], None] | None = None
|
||||||
|
|
||||||
def set_snapshot_path(self, path: Path) -> None:
|
def set_snapshot_path(self, path: Path) -> None:
|
||||||
self._snapshot_path = path
|
self._snapshot_path = path
|
||||||
@@ -65,6 +70,71 @@ class AgentCoordinator:
|
|||||||
for runtime in self.runtimes.values():
|
for runtime in self.runtimes.values():
|
||||||
runtime.wake.set()
|
runtime.wake.set()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def reserve_stopped(self) -> bool:
|
||||||
|
return self._reserve_stopped
|
||||||
|
|
||||||
|
@property
|
||||||
|
def budget_paused(self) -> bool:
|
||||||
|
return self._budget_paused
|
||||||
|
|
||||||
|
def set_budget_extender(self, extend: Callable[[], None]) -> None:
|
||||||
|
self._extend_budget = extend
|
||||||
|
|
||||||
|
async def pause_for_budget(self, agent_id: str) -> None:
|
||||||
|
async with self._lock:
|
||||||
|
self._budget_paused = True
|
||||||
|
await self.set_status(agent_id, "budget_paused")
|
||||||
|
|
||||||
|
async def resume_from_budget_pause(self, *, exclude: str | None = None) -> None:
|
||||||
|
async with self._lock:
|
||||||
|
if not self._budget_paused:
|
||||||
|
return
|
||||||
|
self._budget_paused = False
|
||||||
|
paused = [aid for aid, status in self.statuses.items() if status == "budget_paused"]
|
||||||
|
if self._extend_budget is not None:
|
||||||
|
self._extend_budget()
|
||||||
|
for aid in paused:
|
||||||
|
await self.set_status(aid, "waiting")
|
||||||
|
if aid != exclude:
|
||||||
|
await self.send(
|
||||||
|
aid,
|
||||||
|
{
|
||||||
|
"from": "system",
|
||||||
|
"type": "budget_extended",
|
||||||
|
"content": (
|
||||||
|
"[Budget] The user extended the scan budget \u2014 continue your "
|
||||||
|
"current task."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def reset_budget_stops(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
budget_stopped: bool,
|
||||||
|
reserve_stopped: bool,
|
||||||
|
budget_paused: bool = False,
|
||||||
|
) -> None:
|
||||||
|
async with self._lock:
|
||||||
|
self._budget_stopped = budget_stopped
|
||||||
|
self._reserve_stopped = reserve_stopped
|
||||||
|
if not budget_paused:
|
||||||
|
self._budget_paused = False
|
||||||
|
for aid, status in self.statuses.items():
|
||||||
|
if status == "budget_paused":
|
||||||
|
self.statuses[aid] = "waiting"
|
||||||
|
await self._maybe_snapshot()
|
||||||
|
|
||||||
|
async def claim_reserve_notification(self) -> str | None:
|
||||||
|
async with self._lock:
|
||||||
|
if self._reserve_stopped:
|
||||||
|
return None
|
||||||
|
self._reserve_stopped = True
|
||||||
|
for runtime in self.runtimes.values():
|
||||||
|
runtime.wake.set()
|
||||||
|
return next((aid for aid, parent in self.parent_of.items() if parent is None), None)
|
||||||
|
|
||||||
async def register(
|
async def register(
|
||||||
self,
|
self,
|
||||||
agent_id: str,
|
agent_id: str,
|
||||||
@@ -130,8 +200,12 @@ class AgentCoordinator:
|
|||||||
logger.info("agent.status %s=%s", agent_id, status)
|
logger.info("agent.status %s=%s", agent_id, status)
|
||||||
await self._maybe_snapshot()
|
await self._maybe_snapshot()
|
||||||
|
|
||||||
async def send(self, target_agent_id: str, message: dict[str, Any]) -> bool:
|
async def send(
|
||||||
|
self, target_agent_id: str, message: dict[str, Any], *, interrupt: bool = True
|
||||||
|
) -> bool:
|
||||||
"""Deliver a user/peer message by appending it to the target SDK session."""
|
"""Deliver a user/peer message by appending it to the target SDK session."""
|
||||||
|
if message.get("from") == "user" and self._budget_paused:
|
||||||
|
await self.resume_from_budget_pause(exclude=target_agent_id)
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
if target_agent_id not in self.statuses:
|
if target_agent_id not in self.statuses:
|
||||||
logger.debug("agent.send dropped unknown target=%s", target_agent_id)
|
logger.debug("agent.send dropped unknown target=%s", target_agent_id)
|
||||||
@@ -139,7 +213,7 @@ class AgentCoordinator:
|
|||||||
runtime = self.runtimes.setdefault(target_agent_id, AgentRuntime())
|
runtime = self.runtimes.setdefault(target_agent_id, AgentRuntime())
|
||||||
session = runtime.session
|
session = runtime.session
|
||||||
stream = runtime.stream
|
stream = runtime.stream
|
||||||
interrupt = runtime.interrupt_on_message
|
interrupt_on_message = runtime.interrupt_on_message
|
||||||
if session is None:
|
if session is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"agent.send dropped target=%s because its SDK session is not attached",
|
"agent.send dropped target=%s because its SDK session is not attached",
|
||||||
@@ -158,7 +232,7 @@ class AgentCoordinator:
|
|||||||
async with self._lock:
|
async with self._lock:
|
||||||
self.pending_counts[target_agent_id] = self.pending_counts.get(target_agent_id, 0) + 1
|
self.pending_counts[target_agent_id] = self.pending_counts.get(target_agent_id, 0) + 1
|
||||||
self.runtimes.setdefault(target_agent_id, AgentRuntime()).wake.set()
|
self.runtimes.setdefault(target_agent_id, AgentRuntime()).wake.set()
|
||||||
if stream is not None and interrupt:
|
if stream is not None and interrupt and interrupt_on_message:
|
||||||
stream.cancel(mode="immediate")
|
stream.cancel(mode="immediate")
|
||||||
await self._maybe_snapshot()
|
await self._maybe_snapshot()
|
||||||
return True
|
return True
|
||||||
@@ -166,7 +240,8 @@ class AgentCoordinator:
|
|||||||
async def wait_for_message(self, agent_id: str) -> None:
|
async def wait_for_message(self, agent_id: str) -> None:
|
||||||
while True:
|
while True:
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
if self._budget_stopped or self.pending_counts.get(agent_id, 0) > 0:
|
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
|
return
|
||||||
wake = self.runtimes.setdefault(agent_id, AgentRuntime()).wake
|
wake = self.runtimes.setdefault(agent_id, AgentRuntime()).wake
|
||||||
wake.clear()
|
wake.clear()
|
||||||
@@ -300,6 +375,9 @@ class AgentCoordinator:
|
|||||||
"metadata": {aid: dict(md) for aid, md in self.metadata.items()},
|
"metadata": {aid: dict(md) for aid, md in self.metadata.items()},
|
||||||
"pending_counts": dict(self.pending_counts),
|
"pending_counts": dict(self.pending_counts),
|
||||||
"errors": dict(self.errors),
|
"errors": dict(self.errors),
|
||||||
|
"budget_stopped": self._budget_stopped,
|
||||||
|
"reserve_stopped": self._reserve_stopped,
|
||||||
|
"budget_paused": self._budget_paused,
|
||||||
}
|
}
|
||||||
|
|
||||||
async def restore(self, snap: dict[str, Any]) -> None:
|
async def restore(self, snap: dict[str, Any]) -> None:
|
||||||
@@ -310,6 +388,9 @@ class AgentCoordinator:
|
|||||||
self.metadata = {aid: dict(md) for aid, md in snap.get("metadata", {}).items()}
|
self.metadata = {aid: dict(md) for aid, md in snap.get("metadata", {}).items()}
|
||||||
self.pending_counts = dict(snap.get("pending_counts", {}))
|
self.pending_counts = dict(snap.get("pending_counts", {}))
|
||||||
self.errors = dict(snap.get("errors", {}))
|
self.errors = dict(snap.get("errors", {}))
|
||||||
|
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))
|
||||||
for aid in self.statuses:
|
for aid in self.statuses:
|
||||||
self.runtimes.setdefault(aid, AgentRuntime())
|
self.runtimes.setdefault(aid, AgentRuntime())
|
||||||
|
|
||||||
|
|||||||
+278
-41
@@ -13,15 +13,27 @@ from agents import RunConfig, Runner
|
|||||||
from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError
|
from agents.exceptions import AgentsException, MaxTurnsExceeded, UserError
|
||||||
from agents.sandbox.errors import ExecTransportError
|
from agents.sandbox.errors import ExecTransportError
|
||||||
from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore]
|
from docker import errors as docker_errors # type: ignore[import-untyped, unused-ignore]
|
||||||
from openai import APIError
|
from openai import (
|
||||||
|
APIConnectionError,
|
||||||
|
APIError,
|
||||||
|
APIStatusError,
|
||||||
|
APITimeoutError,
|
||||||
|
RateLimitError,
|
||||||
|
)
|
||||||
|
|
||||||
from strix.core.hooks import BudgetExceededError
|
from strix.config import codex
|
||||||
|
from strix.core.hooks import (
|
||||||
|
BudgetExceededError,
|
||||||
|
BudgetPausedError,
|
||||||
|
SubagentBudgetReservedError,
|
||||||
|
)
|
||||||
from strix.core.inputs import child_initial_input
|
from strix.core.inputs import child_initial_input
|
||||||
from strix.core.sessions import (
|
from strix.core.sessions import (
|
||||||
enforce_image_budget,
|
enforce_image_budget,
|
||||||
open_agent_session,
|
open_agent_session,
|
||||||
strip_all_images_from_session,
|
strip_all_images_from_session,
|
||||||
)
|
)
|
||||||
|
from strix.llm.compaction import is_context_overflow, maybe_compact
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -40,6 +52,91 @@ logger = logging.getLogger(__name__)
|
|||||||
StreamEventSink = Callable[[str, Any], None]
|
StreamEventSink = Callable[[str, Any], None]
|
||||||
|
|
||||||
_INPUT_REJECTION_CODES = frozenset({400, 404, 422})
|
_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
|
||||||
|
|
||||||
|
|
||||||
|
def _agent_instructions(agent: Any) -> str:
|
||||||
|
instructions = getattr(agent, "instructions", None)
|
||||||
|
return instructions if isinstance(instructions, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _agent_tools_text(agent: Any) -> str:
|
||||||
|
parts: list[str] = []
|
||||||
|
for tool in getattr(agent, "tools", []) or []:
|
||||||
|
name = getattr(tool, "name", "")
|
||||||
|
description = getattr(tool, "description", "") or ""
|
||||||
|
schema = getattr(tool, "params_json_schema", "") or ""
|
||||||
|
parts.append(f"{name} {description} {schema}")
|
||||||
|
return "\n".join(parts)
|
||||||
|
|
||||||
|
|
||||||
|
async def _compact_session(
|
||||||
|
agent: Any, session: Session, run_config: RunConfig, *, force: bool
|
||||||
|
) -> bool:
|
||||||
|
model = _run_config_model(run_config)
|
||||||
|
if session is None or model is None:
|
||||||
|
return False
|
||||||
|
return await maybe_compact(
|
||||||
|
session,
|
||||||
|
model=model,
|
||||||
|
instructions=_agent_instructions(agent),
|
||||||
|
tools_text=_agent_tools_text(agent),
|
||||||
|
force=force,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_GUARDRAIL_PARK_ERROR = (
|
||||||
|
"Blocked by the model's content guardrail (flagged as a possible cybersecurity risk). "
|
||||||
|
"Set STRIX_LLM to a model that isn't blocked and resume the scan to continue."
|
||||||
|
)
|
||||||
|
|
||||||
|
_TRANSIENT_MODEL_STATUS_CODES = frozenset({408, 500, 502, 503, 504})
|
||||||
|
_MAX_TRANSIENT_MODEL_RETRIES = 4
|
||||||
|
_TRANSIENT_MODEL_RETRY_BASE_DELAY_S = 2.0
|
||||||
|
_TRANSIENT_MODEL_RETRY_MAX_DELAY_S = 30.0
|
||||||
|
|
||||||
|
|
||||||
|
def _model_error_status_code(exc: BaseException) -> int | None:
|
||||||
|
code = getattr(exc, "status_code", None)
|
||||||
|
return code if isinstance(code, int) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _is_transient_model_error(exc: BaseException) -> bool:
|
||||||
|
if isinstance(exc, RateLimitError):
|
||||||
|
return False
|
||||||
|
if isinstance(exc, APITimeoutError | APIConnectionError):
|
||||||
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def _transient_model_retry_delay(attempt: int) -> float:
|
||||||
|
delay = _TRANSIENT_MODEL_RETRY_BASE_DELAY_S * float(2 ** (attempt - 1))
|
||||||
|
return min(delay, _TRANSIENT_MODEL_RETRY_MAX_DELAY_S)
|
||||||
|
|
||||||
|
|
||||||
async def run_agent_loop(
|
async def run_agent_loop(
|
||||||
@@ -64,21 +161,34 @@ async def run_agent_loop(
|
|||||||
)
|
)
|
||||||
result: RunResultBase | None = None
|
result: RunResultBase | None = None
|
||||||
|
|
||||||
|
budget_stopped = coordinator.budget_stopped
|
||||||
|
reserve_stopped = coordinator.reserve_stopped
|
||||||
|
if budget_stopped:
|
||||||
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
|
raise BudgetExceededError("scan budget reached")
|
||||||
|
if reserve_stopped and context.get("parent_id") is not None:
|
||||||
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
|
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
|
||||||
|
|
||||||
|
if reserve_stopped and start_parked and interactive and context.get("parent_id") is None:
|
||||||
|
await coordinator.send(agent_id, _reserve_notice())
|
||||||
|
|
||||||
if not (start_parked and interactive):
|
if not (start_parked and interactive):
|
||||||
if interactive:
|
if interactive:
|
||||||
result = await _run_cycle(
|
with contextlib.suppress(BudgetPausedError):
|
||||||
agent,
|
result = await _run_cycle(
|
||||||
coordinator,
|
agent,
|
||||||
agent_id,
|
coordinator,
|
||||||
input_data=initial_input,
|
agent_id,
|
||||||
run_config=run_config,
|
input_data=initial_input,
|
||||||
context=context,
|
run_config=run_config,
|
||||||
max_turns=max_turns,
|
context=context,
|
||||||
session=session,
|
max_turns=max_turns,
|
||||||
interactive=interactive,
|
session=session,
|
||||||
event_sink=event_sink,
|
interactive=interactive,
|
||||||
hooks=hooks,
|
event_sink=event_sink,
|
||||||
)
|
hooks=hooks,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
result = await _run_noninteractive_until_lifecycle(
|
result = await _run_noninteractive_until_lifecycle(
|
||||||
agent,
|
agent,
|
||||||
@@ -106,20 +216,25 @@ async def run_agent_loop(
|
|||||||
await coordinator.set_status(agent_id, "stopped")
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
raise BudgetExceededError("scan budget reached")
|
raise BudgetExceededError("scan budget reached")
|
||||||
|
|
||||||
|
if coordinator.reserve_stopped and context.get("parent_id") is not None:
|
||||||
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
|
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
|
||||||
|
|
||||||
await coordinator.consume_pending(agent_id)
|
await coordinator.consume_pending(agent_id)
|
||||||
result = await _run_cycle(
|
with contextlib.suppress(BudgetPausedError):
|
||||||
agent,
|
result = await _run_cycle(
|
||||||
coordinator,
|
agent,
|
||||||
agent_id,
|
coordinator,
|
||||||
input_data=[],
|
agent_id,
|
||||||
run_config=run_config,
|
input_data=[],
|
||||||
context=context,
|
run_config=run_config,
|
||||||
max_turns=max_turns,
|
context=context,
|
||||||
session=session,
|
max_turns=max_turns,
|
||||||
interactive=interactive,
|
session=session,
|
||||||
event_sink=event_sink,
|
interactive=interactive,
|
||||||
hooks=hooks,
|
event_sink=event_sink,
|
||||||
)
|
hooks=hooks,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def spawn_child_agent(
|
async def spawn_child_agent(
|
||||||
@@ -212,6 +327,7 @@ async def respawn_subagents(
|
|||||||
if coordinator.parent_of.get(aid) is None or aid == root_id:
|
if coordinator.parent_of.get(aid) is None or aid == root_id:
|
||||||
continue
|
continue
|
||||||
md["_restored_status"] = status
|
md["_restored_status"] = status
|
||||||
|
md["_restored_error"] = coordinator.errors.get(aid)
|
||||||
candidates.append(
|
candidates.append(
|
||||||
(
|
(
|
||||||
aid,
|
aid,
|
||||||
@@ -224,7 +340,8 @@ async def respawn_subagents(
|
|||||||
for child_id, name, parent_id, md in candidates:
|
for child_id, name, parent_id, md in candidates:
|
||||||
try:
|
try:
|
||||||
restored_status = str(md.get("_restored_status") or "running")
|
restored_status = str(md.get("_restored_status") or "running")
|
||||||
start_parked = interactive and restored_status != "running"
|
recoverable_park = restored_status == "waiting" and bool(md.get("_restored_error"))
|
||||||
|
start_parked = interactive and restored_status != "running" and not recoverable_park
|
||||||
|
|
||||||
if start_parked:
|
if start_parked:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -291,6 +408,10 @@ async def _run_noninteractive_until_lifecycle(
|
|||||||
await coordinator.set_status(agent_id, "stopped")
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
raise BudgetExceededError("scan budget reached")
|
raise BudgetExceededError("scan budget reached")
|
||||||
|
|
||||||
|
if coordinator.reserve_stopped and context.get("parent_id") is not None:
|
||||||
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
|
raise SubagentBudgetReservedError("scan reached the sub-agent budget reserve")
|
||||||
|
|
||||||
result = await _run_cycle(
|
result = await _run_cycle(
|
||||||
agent,
|
agent,
|
||||||
coordinator,
|
coordinator,
|
||||||
@@ -321,7 +442,7 @@ async def _run_noninteractive_until_lifecycle(
|
|||||||
|
|
||||||
if invalid_final_outputs >= invalid_final_output_limit:
|
if invalid_final_outputs >= invalid_final_output_limit:
|
||||||
await coordinator.set_status(agent_id, "crashed")
|
await coordinator.set_status(agent_id, "crashed")
|
||||||
await _notify_parent_on_crash(coordinator, agent_id, "crashed")
|
await _notify_parent_on_terminal(coordinator, agent_id, "crashed")
|
||||||
raise MaxTurnsExceeded(
|
raise MaxTurnsExceeded(
|
||||||
"Agent exhausted non-interactive recovery attempts without calling "
|
"Agent exhausted non-interactive recovery attempts without calling "
|
||||||
"finish_scan or agent_finish."
|
"finish_scan or agent_finish."
|
||||||
@@ -350,6 +471,8 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
hooks: RunHooks[dict[str, Any]] | None,
|
hooks: RunHooks[dict[str, Any]] | None,
|
||||||
) -> RunResultBase | None:
|
) -> RunResultBase | None:
|
||||||
image_strips = 0
|
image_strips = 0
|
||||||
|
compactions = 0
|
||||||
|
model_retries = 0
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
await coordinator.mark_running(agent_id)
|
await coordinator.mark_running(agent_id)
|
||||||
@@ -360,6 +483,10 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
await enforce_image_budget(session, max_images)
|
await enforce_image_budget(session, max_images)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("image-budget enforcement failed for %s", agent_id)
|
logger.exception("image-budget enforcement failed for %s", agent_id)
|
||||||
|
try:
|
||||||
|
await _compact_session(agent, session, run_config, force=False)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("proactive compaction failed for %s", agent_id)
|
||||||
stream = Runner.run_streamed(
|
stream = Runner.run_streamed(
|
||||||
agent,
|
agent,
|
||||||
input=input_data,
|
input=input_data,
|
||||||
@@ -380,9 +507,9 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
logger.exception("stream event sink failed for %s", agent_id)
|
logger.exception("stream event sink failed for %s", agent_id)
|
||||||
if stream.run_loop_exception is not None:
|
if stream.run_loop_exception is not None:
|
||||||
raise stream.run_loop_exception
|
raise stream.run_loop_exception
|
||||||
except BudgetExceededError:
|
if refusal := _structured_provider_refusal(stream):
|
||||||
# A RuntimeError subclass: re-raise explicitly so it is never
|
raise ProviderRefusalError(refusal)
|
||||||
# mistaken for the LiteLLM "after shutdown" race below.
|
except (BudgetExceededError, BudgetPausedError, SubagentBudgetReservedError):
|
||||||
raise
|
raise
|
||||||
except RuntimeError as stream_exc:
|
except RuntimeError as stream_exc:
|
||||||
if "after shutdown" not in str(stream_exc):
|
if "after shutdown" not in str(stream_exc):
|
||||||
@@ -401,6 +528,15 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
await coordinator.detach_stream(agent_id, stream)
|
await coordinator.detach_stream(agent_id, stream)
|
||||||
|
except BudgetPausedError as exc:
|
||||||
|
logger.info("agent %s paused at the scan budget limit: %s", agent_id, exc)
|
||||||
|
await coordinator.pause_for_budget(agent_id)
|
||||||
|
raise
|
||||||
|
except SubagentBudgetReservedError as exc:
|
||||||
|
logger.info("sub-agent %s stopped at the budget reserve: %s", agent_id, exc)
|
||||||
|
await coordinator.set_status(agent_id, "stopped")
|
||||||
|
await _notify_root_on_budget_reserve(coordinator)
|
||||||
|
raise
|
||||||
except BudgetExceededError as exc:
|
except BudgetExceededError as exc:
|
||||||
logger.info(
|
logger.info(
|
||||||
"agent %s reached the scan budget limit; stopping the scan: %s", agent_id, exc
|
"agent %s reached the scan budget limit; stopping the scan: %s", agent_id, exc
|
||||||
@@ -428,6 +564,50 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
)
|
)
|
||||||
input_data = []
|
input_data = []
|
||||||
continue
|
continue
|
||||||
|
if (
|
||||||
|
compactions < _MAX_COMPACTIONS_PER_CYCLE
|
||||||
|
and session is not None
|
||||||
|
and is_context_overflow(exc)
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
compacted = await _compact_session(agent, session, run_config, force=True)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("overflow compaction recovery failed for %s", agent_id)
|
||||||
|
compacted = False
|
||||||
|
if compacted:
|
||||||
|
compactions += 1
|
||||||
|
logger.info(
|
||||||
|
"Compacted %s session after context overflow; retrying (%d)",
|
||||||
|
agent_id,
|
||||||
|
compactions,
|
||||||
|
)
|
||||||
|
input_data = []
|
||||||
|
continue
|
||||||
|
if model_retries < _MAX_TRANSIENT_MODEL_RETRIES and _is_transient_model_error(exc):
|
||||||
|
model_retries += 1
|
||||||
|
delay = _transient_model_retry_delay(model_retries)
|
||||||
|
logger.warning(
|
||||||
|
"transient model/provider error for %s; replaying turn "
|
||||||
|
"(attempt %d/%d, backoff %.1fs): %r",
|
||||||
|
agent_id,
|
||||||
|
model_retries,
|
||||||
|
_MAX_TRANSIENT_MODEL_RETRIES,
|
||||||
|
delay,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
await asyncio.sleep(delay)
|
||||||
|
if session is not None:
|
||||||
|
input_data = []
|
||||||
|
continue
|
||||||
|
if codex.is_content_guardrail_error(exc):
|
||||||
|
return await _handle_content_guardrail(
|
||||||
|
coordinator, agent_id, exc, interactive=interactive
|
||||||
|
)
|
||||||
|
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:
|
if not interactive:
|
||||||
raise
|
raise
|
||||||
if isinstance(exc, MaxTurnsExceeded):
|
if isinstance(exc, MaxTurnsExceeded):
|
||||||
@@ -438,13 +618,29 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
|
|||||||
status = "crashed"
|
status = "crashed"
|
||||||
logger.exception("agent run failed for %s; parking as %s", agent_id, status)
|
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 coordinator.set_status(agent_id, status, error=str(exc) or type(exc).__name__)
|
||||||
await _notify_parent_on_crash(coordinator, agent_id, status)
|
await _notify_parent_on_terminal(coordinator, agent_id, status)
|
||||||
return None
|
return None
|
||||||
else:
|
else:
|
||||||
await _settle_run_result(coordinator, agent_id, interactive)
|
await _settle_run_result(coordinator, agent_id, interactive)
|
||||||
return stream
|
return stream
|
||||||
|
|
||||||
|
|
||||||
|
async def _handle_content_guardrail(
|
||||||
|
coordinator: AgentCoordinator,
|
||||||
|
agent_id: str,
|
||||||
|
exc: BaseException,
|
||||||
|
*,
|
||||||
|
interactive: bool,
|
||||||
|
) -> RunResultBase | None:
|
||||||
|
logger.warning("agent %s blocked by the model's content guardrail: %s", agent_id, exc)
|
||||||
|
if interactive:
|
||||||
|
await coordinator.set_status(agent_id, "waiting", error=_GUARDRAIL_PARK_ERROR)
|
||||||
|
return None
|
||||||
|
await coordinator.set_status(agent_id, "failed", error=_GUARDRAIL_PARK_ERROR)
|
||||||
|
await _notify_parent_on_terminal(coordinator, agent_id, "failed")
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
async def _settle_run_result(
|
async def _settle_run_result(
|
||||||
coordinator: AgentCoordinator,
|
coordinator: AgentCoordinator,
|
||||||
agent_id: str,
|
agent_id: str,
|
||||||
@@ -502,12 +698,31 @@ async def _append_noninteractive_tool_required_message(
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
async def _notify_parent_on_crash(
|
_TERMINAL_NOTICE = {
|
||||||
|
"crashed": (
|
||||||
|
"[Agent crash] {name} ({agent_id}) terminated unexpectedly. "
|
||||||
|
"Stop waiting on this child unless you want to message it again."
|
||||||
|
),
|
||||||
|
"failed": (
|
||||||
|
"[Agent failed] {name} ({agent_id}) stopped with an error and will not "
|
||||||
|
"send a completion report. Stop waiting on this child unless you want to "
|
||||||
|
"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."
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _notify_parent_on_terminal(
|
||||||
coordinator: AgentCoordinator,
|
coordinator: AgentCoordinator,
|
||||||
agent_id: str,
|
agent_id: str,
|
||||||
status: str,
|
status: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
if status != "crashed":
|
template = _TERMINAL_NOTICE.get(status)
|
||||||
|
if template is None:
|
||||||
return
|
return
|
||||||
async with coordinator._lock:
|
async with coordinator._lock:
|
||||||
parent = coordinator.parent_of.get(agent_id)
|
parent = coordinator.parent_of.get(agent_id)
|
||||||
@@ -518,16 +733,36 @@ async def _notify_parent_on_crash(
|
|||||||
parent,
|
parent,
|
||||||
{
|
{
|
||||||
"from": agent_id,
|
"from": agent_id,
|
||||||
"type": "crash",
|
"type": status,
|
||||||
"priority": "high",
|
"priority": "high",
|
||||||
"content": (
|
"content": template.format(name=name, agent_id=agent_id),
|
||||||
f"[Agent crash] {name} ({agent_id}) terminated unexpectedly. "
|
|
||||||
"Stop waiting on this child unless you want to message it again."
|
|
||||||
),
|
|
||||||
},
|
},
|
||||||
|
interrupt=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _reserve_notice() -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"from": "system",
|
||||||
|
"type": "budget_reserve_stop",
|
||||||
|
"priority": "high",
|
||||||
|
"content": (
|
||||||
|
"[Budget reserve] The scan has reached the sub-agent budget reserve: every "
|
||||||
|
"sub-agent is being force-stopped as soon as its in-flight turn completes, and "
|
||||||
|
"none will send a completion report. Their confirmed vulnerabilities are "
|
||||||
|
"already filed as they were found. Do not wait on any sub-agents and do not "
|
||||||
|
"spawn new ones — wrap up now and call finish_scan."
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _notify_root_on_budget_reserve(coordinator: AgentCoordinator) -> None:
|
||||||
|
root = await coordinator.claim_reserve_notification()
|
||||||
|
if root is None:
|
||||||
|
return
|
||||||
|
await coordinator.send(root, _reserve_notice())
|
||||||
|
|
||||||
|
|
||||||
async def _start_child_runner(
|
async def _start_child_runner(
|
||||||
*,
|
*,
|
||||||
parent_ctx: dict[str, Any],
|
parent_ctx: dict[str, Any],
|
||||||
@@ -579,6 +814,8 @@ async def _start_child_runner(
|
|||||||
)
|
)
|
||||||
except BudgetExceededError:
|
except BudgetExceededError:
|
||||||
logger.info("child %s stopped after reaching the scan budget limit", child_id)
|
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)
|
||||||
|
|
||||||
task_handle = asyncio.create_task(_child_loop(), name=f"agent-{name}-{child_id}")
|
task_handle = asyncio.create_task(_child_loop(), name=f"agent-{name}-{child_id}")
|
||||||
await coordinator.attach_runtime(child_id, task=task_handle)
|
await coordinator.attach_runtime(child_id, task=task_handle)
|
||||||
|
|||||||
+203
-4
@@ -14,26 +14,210 @@ from strix.report.state import get_global_report_state
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from agents import RunContextWrapper
|
from agents import RunContextWrapper
|
||||||
from agents.agent import Agent
|
from agents.agent import Agent
|
||||||
from agents.items import ModelResponse
|
from agents.items import ModelResponse, TResponseInputItem
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
_STAGE_LABELS: tuple[str, ...] = ("NOTICE", "URGENT", "CRITICAL")
|
||||||
|
_TURN_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
|
||||||
|
_ROOT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.70, 0.85, 0.95)
|
||||||
|
_SUBAGENT_BUDGET_WARN_BANDS: tuple[float, ...] = (0.75, 0.80, 0.85)
|
||||||
|
_SUBAGENT_BUDGET_RESERVE = 0.90
|
||||||
|
|
||||||
|
|
||||||
class BudgetExceededError(RuntimeError):
|
class BudgetExceededError(RuntimeError):
|
||||||
"""Raised when the accumulated LLM cost reaches the configured budget."""
|
"""Raised when the accumulated LLM cost reaches the configured budget."""
|
||||||
|
|
||||||
|
|
||||||
class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
class SubagentBudgetReservedError(RuntimeError):
|
||||||
"""Persist SDK-native usage after every model response."""
|
"""Raised to stop a single sub-agent once the reserve threshold is crossed."""
|
||||||
|
|
||||||
def __init__(self, *, model: str, max_budget_usd: float | None = None) -> None:
|
|
||||||
|
class BudgetPausedError(RuntimeError):
|
||||||
|
"""Raised to park one agent when an interactive scan reaches its budget."""
|
||||||
|
|
||||||
|
|
||||||
|
def recomputed_budget_flags(
|
||||||
|
cost: float,
|
||||||
|
max_budget_usd: float | None,
|
||||||
|
*,
|
||||||
|
interactive: bool,
|
||||||
|
) -> tuple[bool, bool]:
|
||||||
|
"""Return the (budget_stopped, reserve_stopped) flags a resumed scan should carry."""
|
||||||
|
if max_budget_usd is None:
|
||||||
|
return False, False
|
||||||
|
if interactive:
|
||||||
|
return False, False
|
||||||
|
budget_stopped = cost >= max_budget_usd
|
||||||
|
reserve_stopped = cost >= max_budget_usd * _SUBAGENT_BUDGET_RESERVE
|
||||||
|
return budget_stopped, reserve_stopped
|
||||||
|
|
||||||
|
|
||||||
|
def _crossed_stage(fraction: float, bands: tuple[float, ...]) -> int | None:
|
||||||
|
crossed: int | None = None
|
||||||
|
for index, band in enumerate(bands):
|
||||||
|
if fraction >= band:
|
||||||
|
crossed = index
|
||||||
|
return crossed
|
||||||
|
|
||||||
|
|
||||||
|
_ROOT_DIRECTIVES: tuple[str, ...] = (
|
||||||
|
(
|
||||||
|
"As the root agent, begin planning your wind-down of the whole scan: avoid "
|
||||||
|
"starting large new lines of investigation, and keep your required objectives on "
|
||||||
|
"track so you can call finish_scan comfortably before the limit."
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"As the root agent, prioritize wrapping up the whole scan now: stop opening new "
|
||||||
|
"lines of investigation, close out only what is essential, and move toward calling "
|
||||||
|
"finish_scan to compile and deliver the final report."
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"As the root agent, STOP all other work on the whole scan and finish immediately: "
|
||||||
|
"secure your findings and call finish_scan now — anything left unfinished when the "
|
||||||
|
"limit is hit is discarded."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
_SUBAGENT_DIRECTIVES: tuple[str, ...] = (
|
||||||
|
(
|
||||||
|
"As a sub-agent, begin planning your wind-down: avoid starting large new subtasks, "
|
||||||
|
"and if you are close to a confirmed, validated vulnerability, drive it to a result "
|
||||||
|
"you can report."
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"As a sub-agent, prioritize wrapping up your task now: report any confirmed, "
|
||||||
|
"validated vulnerability, finish work that is nearly done rather than starting "
|
||||||
|
"anything new, and prepare to call agent_finish."
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"As a sub-agent, STOP all other work and finish immediately: report any confirmed "
|
||||||
|
"vulnerability right now and call agent_finish to hand your results back to your "
|
||||||
|
"parent before you are cut off."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _wrapup_directive(context: RunContextWrapper[dict[str, Any]], stage: int) -> str:
|
||||||
|
is_root = context.context.get("parent_id") is None
|
||||||
|
directives = _ROOT_DIRECTIVES if is_root else _SUBAGENT_DIRECTIVES
|
||||||
|
return directives[stage]
|
||||||
|
|
||||||
|
|
||||||
|
def _urgency(stage: int) -> str:
|
||||||
|
return _STAGE_LABELS[stage]
|
||||||
|
|
||||||
|
|
||||||
|
class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
||||||
|
"""Persist SDK-native usage and warn/stop as turn and cost budgets are consumed."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
model: str,
|
||||||
|
max_budget_usd: float | None = None,
|
||||||
|
max_turns: int | None = None,
|
||||||
|
interactive: bool = False,
|
||||||
|
) -> None:
|
||||||
if max_budget_usd is not None and (
|
if max_budget_usd is not None and (
|
||||||
not math.isfinite(max_budget_usd) or max_budget_usd <= 0
|
not math.isfinite(max_budget_usd) or max_budget_usd <= 0
|
||||||
):
|
):
|
||||||
raise ValueError("max_budget_usd must be a finite number greater than 0")
|
raise ValueError("max_budget_usd must be a finite number greater than 0")
|
||||||
|
if max_turns is not None and max_turns <= 0:
|
||||||
|
raise ValueError("max_turns must be a positive integer")
|
||||||
self._model = model
|
self._model = model
|
||||||
self._max_budget_usd = max_budget_usd
|
self._max_budget_usd = max_budget_usd
|
||||||
|
self._budget_increment = max_budget_usd
|
||||||
|
self._max_turns = max_turns
|
||||||
|
self._interactive = interactive
|
||||||
|
|
||||||
|
def extend_budget(self) -> None:
|
||||||
|
if self._max_budget_usd is None or self._budget_increment is None:
|
||||||
|
return
|
||||||
|
self._max_budget_usd += self._budget_increment
|
||||||
|
|
||||||
|
async def on_llm_start(
|
||||||
|
self,
|
||||||
|
context: RunContextWrapper[dict[str, Any]],
|
||||||
|
agent: Agent[dict[str, Any]], # noqa: ARG002
|
||||||
|
system_prompt: str | None, # noqa: ARG002
|
||||||
|
input_items: list[TResponseInputItem],
|
||||||
|
) -> None:
|
||||||
|
try:
|
||||||
|
self._maybe_warn_turns(context, input_items)
|
||||||
|
self._maybe_warn_budget(context, input_items)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("budget/turn warning injection failed")
|
||||||
|
|
||||||
|
def _maybe_warn_turns(
|
||||||
|
self,
|
||||||
|
context: RunContextWrapper[dict[str, Any]],
|
||||||
|
input_items: list[TResponseInputItem],
|
||||||
|
) -> None:
|
||||||
|
if not self._max_turns:
|
||||||
|
return
|
||||||
|
usage = getattr(context, "usage", None)
|
||||||
|
requests = getattr(usage, "requests", None)
|
||||||
|
if not isinstance(requests, int):
|
||||||
|
return
|
||||||
|
turns_used = requests + 1
|
||||||
|
stage = _crossed_stage(turns_used / self._max_turns, _TURN_WARN_BANDS)
|
||||||
|
if stage is None:
|
||||||
|
return
|
||||||
|
remaining = max(self._max_turns - turns_used, 0)
|
||||||
|
pct = round(100 * turns_used / self._max_turns)
|
||||||
|
content = (
|
||||||
|
f"[{_urgency(stage)}] Turn budget: {turns_used}/{self._max_turns} used ({pct}%). "
|
||||||
|
f"About {remaining} turn(s) remain before this agent is force-stopped and any "
|
||||||
|
f"in-progress work is discarded. {_wrapup_directive(context, stage)}"
|
||||||
|
)
|
||||||
|
input_items.append({"role": "user", "content": content})
|
||||||
|
|
||||||
|
def _maybe_warn_budget(
|
||||||
|
self,
|
||||||
|
context: RunContextWrapper[dict[str, Any]],
|
||||||
|
input_items: list[TResponseInputItem],
|
||||||
|
) -> None:
|
||||||
|
if self._max_budget_usd is None:
|
||||||
|
return
|
||||||
|
report_state = get_global_report_state()
|
||||||
|
if report_state is None:
|
||||||
|
return
|
||||||
|
cost = report_state.get_total_llm_cost()
|
||||||
|
is_root = context.context.get("parent_id") is None
|
||||||
|
if self._interactive:
|
||||||
|
bands = _ROOT_BUDGET_WARN_BANDS
|
||||||
|
else:
|
||||||
|
bands = _ROOT_BUDGET_WARN_BANDS if is_root else _SUBAGENT_BUDGET_WARN_BANDS
|
||||||
|
stage = _crossed_stage(cost / self._max_budget_usd, bands)
|
||||||
|
if stage is None:
|
||||||
|
return
|
||||||
|
pct = round(100 * cost / self._max_budget_usd)
|
||||||
|
reserve_pct = round(_SUBAGENT_BUDGET_RESERVE * 100)
|
||||||
|
if self._interactive:
|
||||||
|
content = (
|
||||||
|
f"[{_urgency(stage)}] Scan cost budget: ${cost:.2f}/${self._max_budget_usd:.2f} "
|
||||||
|
f"spent ({pct}%). This budget is shared across every agent in the scan; when it "
|
||||||
|
"is reached all agents are paused until the user chooses to continue. "
|
||||||
|
f"{_wrapup_directive(context, stage)}"
|
||||||
|
)
|
||||||
|
elif is_root:
|
||||||
|
content = (
|
||||||
|
f"[{_urgency(stage)}] Scan cost budget: ${cost:.2f}/${self._max_budget_usd:.2f} "
|
||||||
|
f"spent ({pct}%). This budget is shared across every agent in the scan; when it "
|
||||||
|
"is reached the whole scan is stopped immediately, and sub-agents are stopped at "
|
||||||
|
f"{reserve_pct}% to reserve the remainder for your final report. "
|
||||||
|
f"{_wrapup_directive(context, stage)}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
content = (
|
||||||
|
f"[{_urgency(stage)}] Scan cost budget: ${cost:.2f}/${self._max_budget_usd:.2f} "
|
||||||
|
f"spent ({pct}%). This budget is shared across every agent in the scan; "
|
||||||
|
f"sub-agents are stopped at {reserve_pct}% to leave the remainder for the root "
|
||||||
|
f"agent's final report. {_wrapup_directive(context, stage)}"
|
||||||
|
)
|
||||||
|
input_items.append({"role": "user", "content": content})
|
||||||
|
|
||||||
async def on_llm_end(
|
async def on_llm_end(
|
||||||
self,
|
self,
|
||||||
@@ -66,6 +250,21 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
|
|||||||
if self._max_budget_usd is not None:
|
if self._max_budget_usd is not None:
|
||||||
cost = report_state.get_total_llm_cost()
|
cost = report_state.get_total_llm_cost()
|
||||||
if cost >= self._max_budget_usd:
|
if cost >= self._max_budget_usd:
|
||||||
|
if self._interactive:
|
||||||
|
raise BudgetPausedError(
|
||||||
|
f"Scan budget of ${self._max_budget_usd:.2f} reached "
|
||||||
|
f"(spent ${cost:.4f}); pausing until the user continues"
|
||||||
|
)
|
||||||
raise BudgetExceededError(
|
raise BudgetExceededError(
|
||||||
f"Token budget of ${self._max_budget_usd:.2f} exceeded (spent ${cost:.4f})"
|
f"Token budget of ${self._max_budget_usd:.2f} exceeded (spent ${cost:.4f})"
|
||||||
)
|
)
|
||||||
|
is_root = ctx.get("parent_id") is None
|
||||||
|
if not self._interactive and not is_root:
|
||||||
|
reserve_limit = self._max_budget_usd * _SUBAGENT_BUDGET_RESERVE
|
||||||
|
if cost >= reserve_limit:
|
||||||
|
raise SubagentBudgetReservedError(
|
||||||
|
f"Sub-agent budget reserve reached: spent ${cost:.4f} of "
|
||||||
|
f"${self._max_budget_usd:.2f} "
|
||||||
|
f"(>= {round(_SUBAGENT_BUDGET_RESERVE * 100)}% reserve); stopping this "
|
||||||
|
"sub-agent so the root agent can finish the scan."
|
||||||
|
)
|
||||||
|
|||||||
@@ -10,6 +10,9 @@ from openai.types.shared import Reasoning
|
|||||||
|
|
||||||
from strix.config.models import (
|
from strix.config.models import (
|
||||||
DEFAULT_MODEL_RETRY,
|
DEFAULT_MODEL_RETRY,
|
||||||
|
bedrock_route_supports_prompt_caching,
|
||||||
|
is_bedrock_route,
|
||||||
|
is_claude_model,
|
||||||
is_known_openai_bare_model,
|
is_known_openai_bare_model,
|
||||||
model_supports_reasoning,
|
model_supports_reasoning,
|
||||||
request_timeout_extra_args,
|
request_timeout_extra_args,
|
||||||
@@ -128,12 +131,15 @@ def make_model_settings(
|
|||||||
model_name: str,
|
model_name: str,
|
||||||
force_required_tool_choice: bool = False,
|
force_required_tool_choice: bool = False,
|
||||||
request_timeout: float | None = None,
|
request_timeout: float | None = None,
|
||||||
|
prompt_cache: bool = True,
|
||||||
|
extra_headers: dict[str, str] | None = None,
|
||||||
) -> ModelSettings:
|
) -> ModelSettings:
|
||||||
model_settings = ModelSettings(
|
model_settings = ModelSettings(
|
||||||
parallel_tool_calls=False,
|
parallel_tool_calls=False,
|
||||||
retry=DEFAULT_MODEL_RETRY,
|
retry=DEFAULT_MODEL_RETRY,
|
||||||
include_usage=True,
|
include_usage=True,
|
||||||
extra_args=request_timeout_extra_args(request_timeout),
|
extra_args=request_timeout_extra_args(request_timeout),
|
||||||
|
extra_headers=dict(extra_headers) if extra_headers else None,
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
reasoning_effort is not None
|
reasoning_effort is not None
|
||||||
@@ -145,9 +151,38 @@ def make_model_settings(
|
|||||||
)
|
)
|
||||||
if force_required_tool_choice and _accepts_required_tool_choice(model_name):
|
if force_required_tool_choice and _accepts_required_tool_choice(model_name):
|
||||||
model_settings = model_settings.resolve(ModelSettings(tool_choice="required"))
|
model_settings = model_settings.resolve(ModelSettings(tool_choice="required"))
|
||||||
|
|
||||||
|
cache_extra_args = _prompt_cache_extra_args(model_name) if prompt_cache else None
|
||||||
|
if cache_extra_args:
|
||||||
|
model_settings = model_settings.resolve(
|
||||||
|
ModelSettings(
|
||||||
|
extra_args={**(model_settings.extra_args or {}), **cache_extra_args},
|
||||||
|
),
|
||||||
|
)
|
||||||
return model_settings
|
return model_settings
|
||||||
|
|
||||||
|
|
||||||
|
def _prompt_cache_extra_args(model_name: str) -> dict[str, Any] | None:
|
||||||
|
"""LiteLLM ``cache_control_injection_points`` for Claude prompt caching.
|
||||||
|
|
||||||
|
System prompt + rolling last-message breakpoint everywhere; ``tool_config``
|
||||||
|
only on Bedrock Converse (the only route whose LiteLLM transform consumes
|
||||||
|
it — elsewhere it leaks onto the wire and native Anthropic 400s). Unmapped
|
||||||
|
Bedrock models get no points at all: Bedrock rejects the passed-through
|
||||||
|
field outright.
|
||||||
|
"""
|
||||||
|
if not is_claude_model(model_name):
|
||||||
|
return None
|
||||||
|
if is_bedrock_route(model_name) and not bedrock_route_supports_prompt_caching(model_name):
|
||||||
|
return None
|
||||||
|
|
||||||
|
points: list[dict[str, Any]] = [{"location": "message", "role": "system"}]
|
||||||
|
if is_bedrock_route(model_name):
|
||||||
|
points.append({"location": "tool_config"})
|
||||||
|
points.append({"location": "message", "index": -1})
|
||||||
|
return {"cache_control_injection_points": points}
|
||||||
|
|
||||||
|
|
||||||
def child_initial_input(
|
def child_initial_input(
|
||||||
*,
|
*,
|
||||||
name: str,
|
name: str,
|
||||||
|
|||||||
+52
-3
@@ -3,10 +3,12 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import contextlib
|
import contextlib
|
||||||
|
import io
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
from agents import RunConfig
|
from agents import RunConfig
|
||||||
@@ -29,7 +31,7 @@ from strix.core.execution import (
|
|||||||
from strix.core.execution import (
|
from strix.core.execution import (
|
||||||
spawn_child_agent as start_child_agent,
|
spawn_child_agent as start_child_agent,
|
||||||
)
|
)
|
||||||
from strix.core.hooks import BudgetExceededError, ReportUsageHooks
|
from strix.core.hooks import BudgetExceededError, ReportUsageHooks, recomputed_budget_flags
|
||||||
from strix.core.inputs import (
|
from strix.core.inputs import (
|
||||||
DEFAULT_MAX_TURNS,
|
DEFAULT_MAX_TURNS,
|
||||||
build_root_task,
|
build_root_task,
|
||||||
@@ -38,8 +40,13 @@ from strix.core.inputs import (
|
|||||||
)
|
)
|
||||||
from strix.core.paths import run_dir_for, runtime_state_dir
|
from strix.core.paths import run_dir_for, runtime_state_dir
|
||||||
from strix.core.sessions import open_agent_session
|
from strix.core.sessions import open_agent_session
|
||||||
|
from strix.report.state import get_global_report_state
|
||||||
from strix.runtime import session_manager
|
from strix.runtime import session_manager
|
||||||
from strix.telemetry.logging import set_scan_id, setup_scan_logging
|
from strix.telemetry.logging import set_scan_id, setup_scan_logging
|
||||||
|
from strix.tools.output_store import (
|
||||||
|
WORKSPACE_SPILL_DIR,
|
||||||
|
configure_spill_writer,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -179,6 +186,18 @@ async def run_strix_scan(
|
|||||||
f"Cannot resume scan {scan_id}: missing SDK session database at {agents_db}",
|
f"Cannot resume scan {scan_id}: missing SDK session database at {agents_db}",
|
||||||
)
|
)
|
||||||
await coordinator.restore(snap)
|
await coordinator.restore(snap)
|
||||||
|
report_state = get_global_report_state()
|
||||||
|
if report_state is not None:
|
||||||
|
budget_stopped, reserve_stopped = recomputed_budget_flags(
|
||||||
|
report_state.get_total_llm_cost(),
|
||||||
|
max_budget_usd,
|
||||||
|
interactive=interactive,
|
||||||
|
)
|
||||||
|
await coordinator.reset_budget_stops(
|
||||||
|
budget_stopped=budget_stopped,
|
||||||
|
reserve_stopped=reserve_stopped,
|
||||||
|
budget_paused=interactive and coordinator.budget_paused,
|
||||||
|
)
|
||||||
for aid, parent in coordinator.parent_of.items():
|
for aid, parent in coordinator.parent_of.items():
|
||||||
if parent is None:
|
if parent is None:
|
||||||
root_id = aid
|
root_id = aid
|
||||||
@@ -203,6 +222,20 @@ async def run_strix_scan(
|
|||||||
)
|
)
|
||||||
logger.info("Sandbox ready for scan %s", scan_id)
|
logger.info("Sandbox ready for scan %s", scan_id)
|
||||||
|
|
||||||
|
sandbox_session = bundle["session"]
|
||||||
|
|
||||||
|
async def _spill_to_workspace(output_id: str, text: str) -> str | None:
|
||||||
|
"""Write an oversized tool result into the sandbox; return its path or None."""
|
||||||
|
path = f"{WORKSPACE_SPILL_DIR}/{output_id}.txt"
|
||||||
|
try:
|
||||||
|
await sandbox_session.write(Path(path), io.BytesIO(text.encode("utf-8")))
|
||||||
|
except Exception:
|
||||||
|
logger.exception("failed to spill tool output to sandbox workspace")
|
||||||
|
return None
|
||||||
|
return path
|
||||||
|
|
||||||
|
configure_spill_writer(_spill_to_workspace)
|
||||||
|
|
||||||
sessions_to_close: list[SQLiteSession] = []
|
sessions_to_close: list[SQLiteSession] = []
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -216,6 +249,8 @@ async def run_strix_scan(
|
|||||||
model_name=resolved_model,
|
model_name=resolved_model,
|
||||||
force_required_tool_choice=settings.llm.force_required_tool_choice,
|
force_required_tool_choice=settings.llm.force_required_tool_choice,
|
||||||
request_timeout=settings.llm.timeout,
|
request_timeout=settings.llm.timeout,
|
||||||
|
prompt_cache=settings.llm.prompt_cache,
|
||||||
|
extra_headers=settings.llm.extra_headers,
|
||||||
)
|
)
|
||||||
run_config = RunConfig(
|
run_config = RunConfig(
|
||||||
model=resolved_model,
|
model=resolved_model,
|
||||||
@@ -224,7 +259,14 @@ async def run_strix_scan(
|
|||||||
sandbox=SandboxRunConfig(client=bundle["client"], session=bundle["session"]),
|
sandbox=SandboxRunConfig(client=bundle["client"], session=bundle["session"]),
|
||||||
trace_include_sensitive_data=False,
|
trace_include_sensitive_data=False,
|
||||||
)
|
)
|
||||||
hooks = ReportUsageHooks(model=resolved_model, max_budget_usd=max_budget_usd)
|
hooks = ReportUsageHooks(
|
||||||
|
model=resolved_model,
|
||||||
|
max_budget_usd=max_budget_usd,
|
||||||
|
max_turns=max_turns,
|
||||||
|
interactive=interactive,
|
||||||
|
)
|
||||||
|
if interactive:
|
||||||
|
coordinator.set_budget_extender(hooks.extend_budget)
|
||||||
|
|
||||||
scope_context = build_scope_context(scan_config)
|
scope_context = build_scope_context(scan_config)
|
||||||
root_context = _merge_root_prompt_context(scope_context, extra_system_prompt_context)
|
root_context = _merge_root_prompt_context(scope_context, extra_system_prompt_context)
|
||||||
@@ -335,6 +377,12 @@ async def run_strix_scan(
|
|||||||
|
|
||||||
async with coordinator._lock:
|
async with coordinator._lock:
|
||||||
root_status = coordinator.statuses.get(root_id)
|
root_status = coordinator.statuses.get(root_id)
|
||||||
|
root_error = coordinator.errors.get(root_id)
|
||||||
|
|
||||||
|
root_recoverable_park = root_status == "waiting" and bool(root_error)
|
||||||
|
root_start_parked = bool(
|
||||||
|
interactive and is_resume and root_status != "running" and not root_recoverable_park
|
||||||
|
)
|
||||||
|
|
||||||
result = await run_agent_loop(
|
result = await run_agent_loop(
|
||||||
agent=root_agent,
|
agent=root_agent,
|
||||||
@@ -346,7 +394,7 @@ async def run_strix_scan(
|
|||||||
agent_id=root_id,
|
agent_id=root_id,
|
||||||
interactive=interactive,
|
interactive=interactive,
|
||||||
session=root_session,
|
session=root_session,
|
||||||
start_parked=bool(interactive and is_resume and root_status != "running"),
|
start_parked=root_start_parked,
|
||||||
event_sink=event_sink,
|
event_sink=event_sink,
|
||||||
hooks=hooks,
|
hooks=hooks,
|
||||||
)
|
)
|
||||||
@@ -399,6 +447,7 @@ async def run_strix_scan(
|
|||||||
await coordinator.set_status(root_id, "failed")
|
await coordinator.set_status(root_id, "failed")
|
||||||
raise
|
raise
|
||||||
finally:
|
finally:
|
||||||
|
configure_spill_writer(None)
|
||||||
for s in sessions_to_close:
|
for s in sessions_to_close:
|
||||||
with contextlib.suppress(Exception):
|
with contextlib.suppress(Exception):
|
||||||
s.close()
|
s.close()
|
||||||
|
|||||||
@@ -92,6 +92,39 @@ async def _rewrite_session(
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
async def replace_session_items(
|
||||||
|
session: Session,
|
||||||
|
new_items: list[Any],
|
||||||
|
*,
|
||||||
|
expected_len: int | None = None,
|
||||||
|
) -> bool:
|
||||||
|
"""Overwrite the session's items, restoring the originals on failure.
|
||||||
|
|
||||||
|
When ``expected_len`` is given, the rewrite is skipped if the session no
|
||||||
|
longer has that many items (a concurrent writer changed it), so a slow
|
||||||
|
compaction summary can't clobber newer turns.
|
||||||
|
"""
|
||||||
|
async with session_write_lock(session):
|
||||||
|
original = list(await session.get_items())
|
||||||
|
if expected_len is not None and len(original) != expected_len:
|
||||||
|
logger.warning(
|
||||||
|
"skipping session rewrite: expected %d items, found %d",
|
||||||
|
expected_len,
|
||||||
|
len(original),
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
rebuilt = cast("list[TResponseInputItem]", new_items)
|
||||||
|
await session.clear_session()
|
||||||
|
try:
|
||||||
|
await session.add_items(rebuilt)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("session rewrite failed; restoring original items")
|
||||||
|
await session.clear_session()
|
||||||
|
await session.add_items(original)
|
||||||
|
raise
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
async def strip_all_images_from_session(session: Session) -> bool:
|
async def strip_all_images_from_session(session: Session) -> bool:
|
||||||
"""Replace every image tool output with a text placeholder (rejection recovery)."""
|
"""Replace every image tool output with a text placeholder (rejection recovery)."""
|
||||||
|
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from rich.panel import Panel
|
|||||||
from rich.text import Text
|
from rich.text import Text
|
||||||
|
|
||||||
from strix.config import load_settings
|
from strix.config import load_settings
|
||||||
|
from strix.core.inputs import DEFAULT_MAX_TURNS
|
||||||
from strix.core.runner import run_strix_scan
|
from strix.core.runner import run_strix_scan
|
||||||
from strix.report.state import ReportState, set_global_report_state
|
from strix.report.state import ReportState, set_global_report_state
|
||||||
from strix.runtime import session_manager
|
from strix.runtime import session_manager
|
||||||
@@ -184,6 +185,7 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915
|
|||||||
local_sources=getattr(args, "local_sources", None) or [],
|
local_sources=getattr(args, "local_sources", None) or [],
|
||||||
interactive=bool(getattr(args, "interactive", False)),
|
interactive=bool(getattr(args, "interactive", False)),
|
||||||
max_budget_usd=getattr(args, "max_budget_usd", None),
|
max_budget_usd=getattr(args, "max_budget_usd", None),
|
||||||
|
max_turns=getattr(args, "max_turns", DEFAULT_MAX_TURNS),
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
stop_updates.set()
|
stop_updates.set()
|
||||||
|
|||||||
+56
-10
@@ -31,6 +31,7 @@ from strix.config.models import (
|
|||||||
is_known_openai_bare_model,
|
is_known_openai_bare_model,
|
||||||
is_recommended_or_frontier_model,
|
is_recommended_or_frontier_model,
|
||||||
)
|
)
|
||||||
|
from strix.core.inputs import DEFAULT_MAX_TURNS, make_model_settings
|
||||||
from strix.core.paths import run_dir_for, runtime_state_dir
|
from strix.core.paths import run_dir_for, runtime_state_dir
|
||||||
from strix.interface.cli import run_cli
|
from strix.interface.cli import run_cli
|
||||||
from strix.interface.tui import run_tui
|
from strix.interface.tui import run_tui
|
||||||
@@ -210,7 +211,7 @@ def validate_environment() -> None:
|
|||||||
padding=(1, 2),
|
padding=(1, 2),
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.error("Missing required env vars: %s", missing_required_vars)
|
logger.debug("Missing required env vars: %s", missing_required_vars)
|
||||||
console.print("\n")
|
console.print("\n")
|
||||||
console.print(panel)
|
console.print(panel)
|
||||||
console.print()
|
console.print()
|
||||||
@@ -223,7 +224,7 @@ def validate_environment() -> None:
|
|||||||
|
|
||||||
def check_docker_installed() -> None:
|
def check_docker_installed() -> None:
|
||||||
if shutil.which("docker") is None:
|
if shutil.which("docker") is None:
|
||||||
logger.error("Docker CLI not found in PATH")
|
logger.debug("Docker CLI not found in PATH")
|
||||||
console = Console()
|
console = Console()
|
||||||
error_text = Text()
|
error_text = Text()
|
||||||
error_text.append("DOCKER NOT INSTALLED", style="bold red")
|
error_text.append("DOCKER NOT INSTALLED", style="bold red")
|
||||||
@@ -381,7 +382,13 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
|
|||||||
model.get_response(
|
model.get_response(
|
||||||
system_instructions="You are a helpful assistant.",
|
system_instructions="You are a helpful assistant.",
|
||||||
input="Reply with just 'OK'.",
|
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=[],
|
tools=[],
|
||||||
output_schema=None,
|
output_schema=None,
|
||||||
handoffs=[],
|
handoffs=[],
|
||||||
@@ -403,7 +410,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
|
# Match the runtime path: send the dedupe key/endpoint per call so a
|
||||||
# separate-provider dedupe model authenticates during warm-up too.
|
# separate-provider dedupe model authenticates during warm-up too.
|
||||||
deduper_extra = _dedupe_extra_args(settings.dedupe)
|
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(
|
await asyncio.wait_for(
|
||||||
deduper.get_response(
|
deduper.get_response(
|
||||||
system_instructions="You are a helpful assistant.",
|
system_instructions="You are a helpful assistant.",
|
||||||
@@ -422,7 +441,7 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
|
|||||||
logger.info("LLM warm-up succeeded for dedupe model %s", dedupe_model)
|
logger.info("LLM warm-up succeeded for dedupe model %s", dedupe_model)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception("LLM warm-up failed")
|
logger.debug("LLM warm-up failed", exc_info=True)
|
||||||
error_text = Text()
|
error_text = Text()
|
||||||
sub_hint = _subscription_error_hint(e)
|
sub_hint = _subscription_error_hint(e)
|
||||||
if sub_hint is not None:
|
if sub_hint is not None:
|
||||||
@@ -481,6 +500,16 @@ def _positive_budget(value: str) -> float:
|
|||||||
return budget
|
return budget
|
||||||
|
|
||||||
|
|
||||||
|
def _positive_int(value: str) -> int:
|
||||||
|
try:
|
||||||
|
parsed = int(value)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise argparse.ArgumentTypeError(f"invalid int value: {value!r}") from exc
|
||||||
|
if parsed <= 0:
|
||||||
|
raise argparse.ArgumentTypeError("must be an integer greater than 0")
|
||||||
|
return parsed
|
||||||
|
|
||||||
|
|
||||||
def parse_arguments() -> argparse.Namespace:
|
def parse_arguments() -> argparse.Namespace:
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Strix Multi-Agent Cybersecurity Penetration Testing Tool",
|
description="Strix Multi-Agent Cybersecurity Penetration Testing Tool",
|
||||||
@@ -636,10 +665,27 @@ Examples:
|
|||||||
)
|
)
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max-budget-usd",
|
"--max-budget",
|
||||||
|
dest="max_budget_usd",
|
||||||
|
metavar="USD",
|
||||||
type=_positive_budget,
|
type=_positive_budget,
|
||||||
default=None,
|
default=None,
|
||||||
help="Maximum LLM cost in USD (> 0). The scan stops cleanly when this limit is reached.",
|
help=(
|
||||||
|
"Maximum LLM cost in USD (> 0). The scan stops cleanly when this limit is reached. "
|
||||||
|
"Graduated wrap-up warnings are sent to all agents as it is approached."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--max-turns",
|
||||||
|
dest="max_turns",
|
||||||
|
metavar="N",
|
||||||
|
type=_positive_int,
|
||||||
|
default=DEFAULT_MAX_TURNS,
|
||||||
|
help=(
|
||||||
|
"Maximum turns per agent (> 0, default %(default)s). Each agent is force-stopped "
|
||||||
|
"when it reaches this limit, with graduated wrap-up warnings as it is approached."
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -856,7 +902,7 @@ def display_completion_message(args: argparse.Namespace, results_path: Path) ->
|
|||||||
view_text = Text()
|
view_text = Text()
|
||||||
view_text.append("\n")
|
view_text.append("\n")
|
||||||
view_text.append("View", style="dim")
|
view_text.append("View", style="dim")
|
||||||
view_text.append(" ")
|
view_text.append(" ")
|
||||||
view_text.append(f"strix view {args.run_name}", style="#22c55e")
|
view_text.append(f"strix view {args.run_name}", style="#22c55e")
|
||||||
panel_parts.extend(["\n", view_text])
|
panel_parts.extend(["\n", view_text])
|
||||||
|
|
||||||
@@ -918,7 +964,7 @@ def pull_docker_image() -> None:
|
|||||||
last_update = process_pull_line(line, layers_info, status, last_update)
|
last_update = process_pull_line(line, layers_info, status, last_update)
|
||||||
|
|
||||||
except DockerException as e:
|
except DockerException as e:
|
||||||
logger.exception("Failed to pull docker image %s", image)
|
logger.debug("Failed to pull docker image %s", image, exc_info=True)
|
||||||
console.print()
|
console.print()
|
||||||
error_text = Text()
|
error_text = Text()
|
||||||
error_text.append("FAILED TO PULL IMAGE", style="bold red")
|
error_text.append("FAILED TO PULL IMAGE", style="bold red")
|
||||||
@@ -952,7 +998,7 @@ def main() -> None:
|
|||||||
# `strix view [<run>]` is a viewer-only subcommand, dispatched before the
|
# `strix view [<run>]` is a viewer-only subcommand, dispatched before the
|
||||||
# scan argument parser (which requires a target) and before any scan setup.
|
# scan argument parser (which requires a target) and before any scan setup.
|
||||||
if len(sys.argv) > 1 and sys.argv[1] == "view":
|
if len(sys.argv) > 1 and sys.argv[1] == "view":
|
||||||
from strix.viewer.cli import run_view
|
from strix.interface.viewer.cli import run_view
|
||||||
|
|
||||||
run_view(sys.argv[2:])
|
run_view(sys.argv[2:])
|
||||||
return
|
return
|
||||||
|
|||||||
+38
-10
@@ -15,6 +15,7 @@ from typing import TYPE_CHECKING, Any, ClassVar
|
|||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from pygments.token import _TokenType
|
||||||
from textual.timer import Timer
|
from textual.timer import Timer
|
||||||
|
|
||||||
from rich.align import Align
|
from rich.align import Align
|
||||||
@@ -34,6 +35,7 @@ from textual.widgets.tree import TreeNode
|
|||||||
from strix.config import load_settings
|
from strix.config import load_settings
|
||||||
from strix.config.models import is_recommended_or_frontier_model
|
from strix.config.models import is_recommended_or_frontier_model
|
||||||
from strix.core.hooks import BudgetExceededError
|
from strix.core.hooks import BudgetExceededError
|
||||||
|
from strix.core.inputs import DEFAULT_MAX_TURNS
|
||||||
from strix.core.runner import run_strix_scan
|
from strix.core.runner import run_strix_scan
|
||||||
from strix.interface.tui.live_view import TuiLiveView
|
from strix.interface.tui.live_view import TuiLiveView
|
||||||
from strix.interface.tui.messages import send_user_message_to_agent
|
from strix.interface.tui.messages import send_user_message_to_agent
|
||||||
@@ -351,7 +353,7 @@ class VulnerabilityDetailScreen(ModalScreen): # type: ignore[misc]
|
|||||||
if not token_value:
|
if not token_value:
|
||||||
continue
|
continue
|
||||||
color = None
|
color = None
|
||||||
tt = token_type
|
tt: _TokenType | None = token_type
|
||||||
while tt:
|
while tt:
|
||||||
if tt in colors:
|
if tt in colors:
|
||||||
color = colors[tt]
|
color = colors[tt]
|
||||||
@@ -814,6 +816,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
self._scan_completed = threading.Event()
|
self._scan_completed = threading.Event()
|
||||||
self._scan_error: BaseException | None = None
|
self._scan_error: BaseException | None = None
|
||||||
self._error_noted_agents: set[str] = set()
|
self._error_noted_agents: set[str] = set()
|
||||||
|
self._budget_pause_notified = False
|
||||||
|
|
||||||
self._spinner_frame_index: int = 0
|
self._spinner_frame_index: int = 0
|
||||||
self._sweep_num_squares: int = 6
|
self._sweep_num_squares: int = 6
|
||||||
@@ -1046,6 +1049,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
self.live_view.record_agent_error(agent_id, error)
|
self.live_view.record_agent_error(agent_id, error)
|
||||||
else:
|
else:
|
||||||
self._error_noted_agents.discard(agent_id)
|
self._error_noted_agents.discard(agent_id)
|
||||||
|
self._notify_budget_pause(statuses)
|
||||||
|
|
||||||
if self._scan_loop is None or self._scan_loop.is_closed():
|
if self._scan_loop is None or self._scan_loop.is_closed():
|
||||||
return
|
return
|
||||||
@@ -1057,6 +1061,19 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
|
|
||||||
self._agent_graph_sync_future = asyncio.run_coroutine_threadsafe(collect(), self._scan_loop)
|
self._agent_graph_sync_future = asyncio.run_coroutine_threadsafe(collect(), self._scan_loop)
|
||||||
|
|
||||||
|
def _notify_budget_pause(self, statuses: dict[str, Any]) -> None:
|
||||||
|
paused = any(status == "budget_paused" for status in statuses.values())
|
||||||
|
if paused and not self._budget_pause_notified:
|
||||||
|
self._budget_pause_notified = True
|
||||||
|
self.notify(
|
||||||
|
"Budget limit reached \u2014 agents paused. Send a message to continue "
|
||||||
|
"(this extends the budget), or ctrl-q to quit.",
|
||||||
|
severity="warning",
|
||||||
|
timeout=15,
|
||||||
|
)
|
||||||
|
elif not paused:
|
||||||
|
self._budget_pause_notified = False
|
||||||
|
|
||||||
def _update_agent_node(self, agent_id: str, agent_data: dict[str, Any]) -> bool:
|
def _update_agent_node(self, agent_id: str, agent_data: dict[str, Any]) -> bool:
|
||||||
if agent_id not in self.agent_nodes:
|
if agent_id not in self.agent_nodes:
|
||||||
return False
|
return False
|
||||||
@@ -1069,6 +1086,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
status_indicators = {
|
status_indicators = {
|
||||||
"running": "⚪",
|
"running": "⚪",
|
||||||
"waiting": "⏸",
|
"waiting": "⏸",
|
||||||
|
"budget_paused": "⏸",
|
||||||
"completed": "🟢",
|
"completed": "🟢",
|
||||||
"failed": "🔴",
|
"failed": "🔴",
|
||||||
"crashed": "🔴",
|
"crashed": "🔴",
|
||||||
@@ -1266,10 +1284,17 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
self._stop_dot_animation()
|
self._stop_dot_animation()
|
||||||
return (text, Text(), False)
|
return (text, Text(), False)
|
||||||
|
|
||||||
if status == "waiting":
|
if status in {"waiting", "budget_paused"}:
|
||||||
text = Text()
|
text = Text()
|
||||||
text.append("Send message to resume", style="dim")
|
keymap = Text()
|
||||||
return (text, Text(), False)
|
if status == "budget_paused":
|
||||||
|
text.append("Budget limit reached", style="yellow")
|
||||||
|
text.append(" \u00b7 ", style="dim")
|
||||||
|
text.append("Send a message to continue", style="dim")
|
||||||
|
keymap = keymap_styled([("ctrl-q", "quit")])
|
||||||
|
else:
|
||||||
|
text.append("Send message to resume", style="dim")
|
||||||
|
return (text, keymap, False)
|
||||||
|
|
||||||
if status == "running":
|
if status == "running":
|
||||||
if self._agent_has_real_activity(agent_id):
|
if self._agent_has_real_activity(agent_id):
|
||||||
@@ -1494,6 +1519,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
coordinator=self.coordinator,
|
coordinator=self.coordinator,
|
||||||
interactive=True,
|
interactive=True,
|
||||||
max_budget_usd=getattr(self.args, "max_budget_usd", None),
|
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,
|
event_sink=self._capture_sdk_event,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -1501,10 +1527,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
except (KeyboardInterrupt, asyncio.CancelledError):
|
except (KeyboardInterrupt, asyncio.CancelledError):
|
||||||
logger.info("Scan interrupted by user")
|
logger.info("Scan interrupted by user")
|
||||||
except BudgetExceededError:
|
except BudgetExceededError:
|
||||||
# Defensive: the runner stops the scan cleanly on budget and
|
logger.info("Scan stopped: --max-budget limit reached")
|
||||||
# returns, so this normally never propagates. Treat it as a
|
|
||||||
# graceful stop, not a scan error, if it ever does.
|
|
||||||
logger.info("Scan stopped: --max-budget-usd limit reached")
|
|
||||||
except (ConnectionError, TimeoutError) as e:
|
except (ConnectionError, TimeoutError) as e:
|
||||||
logging.exception("Network error during scan")
|
logging.exception("Network error during scan")
|
||||||
self._scan_error = e
|
self._scan_error = e
|
||||||
@@ -1559,6 +1582,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
status_indicators = {
|
status_indicators = {
|
||||||
"running": "⚪",
|
"running": "⚪",
|
||||||
"waiting": "⏸",
|
"waiting": "⏸",
|
||||||
|
"budget_paused": "⏸",
|
||||||
"completed": "🟢",
|
"completed": "🟢",
|
||||||
"failed": "🔴",
|
"failed": "🔴",
|
||||||
"crashed": "🔴",
|
"crashed": "🔴",
|
||||||
@@ -1605,6 +1629,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
status_indicators = {
|
status_indicators = {
|
||||||
"running": "⚪",
|
"running": "⚪",
|
||||||
"waiting": "⏸",
|
"waiting": "⏸",
|
||||||
|
"budget_paused": "⏸",
|
||||||
"completed": "🟢",
|
"completed": "🟢",
|
||||||
"failed": "🔴",
|
"failed": "🔴",
|
||||||
"crashed": "🔴",
|
"crashed": "🔴",
|
||||||
@@ -1729,7 +1754,10 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
message=message,
|
message=message,
|
||||||
)
|
)
|
||||||
if not submitted:
|
if not submitted:
|
||||||
self.notify("Scan loop is not ready; message was not sent", severity="warning")
|
if self._scan_completed.is_set():
|
||||||
|
self.notify("The scan has ended; message was not sent", severity="warning")
|
||||||
|
else:
|
||||||
|
self.notify("Scan loop is not ready; message was not sent", severity="warning")
|
||||||
return
|
return
|
||||||
|
|
||||||
self._displayed_events.clear()
|
self._displayed_events.clear()
|
||||||
@@ -1862,7 +1890,7 @@ class StrixTUIApp(App): # type: ignore[misc]
|
|||||||
webbrowser.open(self._viewer_url)
|
webbrowser.open(self._viewer_url)
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
from strix.viewer.server import authorized_url, bundle_is_built, serve
|
from strix.interface.viewer.server import authorized_url, bundle_is_built, serve
|
||||||
|
|
||||||
if not bundle_is_built():
|
if not bundle_is_built():
|
||||||
self._set_viewer_cta("[#eab308]Viewer UI not built[/]")
|
self._set_viewer_cta("[#eab308]Viewer UI not built[/]")
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ class TuiLiveView:
|
|||||||
self.events: list[dict[str, Any]] = []
|
self.events: list[dict[str, Any]] = []
|
||||||
self._next_event_id = 1
|
self._next_event_id = 1
|
||||||
self._open_assistant_event_by_agent: dict[str, dict[str, Any]] = {}
|
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:
|
def hydrate_from_run_dir(self, run_dir: Path) -> None:
|
||||||
state_dir = runtime_state_dir(run_dir)
|
state_dir = runtime_state_dir(run_dir)
|
||||||
@@ -223,7 +223,8 @@ class TuiLiveView:
|
|||||||
timestamp: str | None = None,
|
timestamp: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
call_id = call["call_id"]
|
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_data = {
|
||||||
"tool_name": call["tool_name"],
|
"tool_name": call["tool_name"],
|
||||||
"args": call["args"],
|
"args": call["args"],
|
||||||
@@ -233,7 +234,7 @@ class TuiLiveView:
|
|||||||
}
|
}
|
||||||
if existing is None:
|
if existing is None:
|
||||||
event = self._append_event(agent_id, "tool", tool_data, timestamp=timestamp)
|
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:
|
else:
|
||||||
existing["data"].update(tool_data)
|
existing["data"].update(tool_data)
|
||||||
self._bump_event(existing, timestamp=timestamp)
|
self._bump_event(existing, timestamp=timestamp)
|
||||||
@@ -249,7 +250,8 @@ class TuiLiveView:
|
|||||||
timestamp: str | None = None,
|
timestamp: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
call_id = output["call_id"]
|
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:
|
if event is None:
|
||||||
event = self._append_event(
|
event = self._append_event(
|
||||||
agent_id,
|
agent_id,
|
||||||
@@ -263,7 +265,7 @@ class TuiLiveView:
|
|||||||
},
|
},
|
||||||
timestamp=timestamp,
|
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"])
|
result = _parse_json_value(output["output"])
|
||||||
event["data"]["result"] = result
|
event["data"]["result"] = result
|
||||||
|
|||||||
@@ -7,6 +7,13 @@ from .base_renderer import BaseToolRenderer
|
|||||||
from .registry import register_tool_renderer
|
from .registry import register_tool_renderer
|
||||||
|
|
||||||
|
|
||||||
|
def _author_label(note: dict[str, Any]) -> str:
|
||||||
|
if note.get("by_you"):
|
||||||
|
return "you"
|
||||||
|
agent_name = note.get("agent_name")
|
||||||
|
return str(agent_name).strip() if agent_name else ""
|
||||||
|
|
||||||
|
|
||||||
@register_tool_renderer
|
@register_tool_renderer
|
||||||
class CreateNoteRenderer(BaseToolRenderer):
|
class CreateNoteRenderer(BaseToolRenderer):
|
||||||
tool_name: ClassVar[str] = "create_note"
|
tool_name: ClassVar[str] = "create_note"
|
||||||
@@ -123,6 +130,9 @@ class ListNotesRenderer(BaseToolRenderer):
|
|||||||
text.append("\n - ")
|
text.append("\n - ")
|
||||||
text.append(title)
|
text.append(title)
|
||||||
text.append(f" ({category})", style="dim")
|
text.append(f" ({category})", style="dim")
|
||||||
|
author = _author_label(note)
|
||||||
|
if author:
|
||||||
|
text.append(f" by {author}", style="dim")
|
||||||
|
|
||||||
if note_content:
|
if note_content:
|
||||||
text.append("\n ")
|
text.append("\n ")
|
||||||
@@ -156,6 +166,9 @@ class GetNoteRenderer(BaseToolRenderer):
|
|||||||
text.append("\n ")
|
text.append("\n ")
|
||||||
text.append(title)
|
text.append(title)
|
||||||
text.append(f" ({category})", style="dim")
|
text.append(f" ({category})", style="dim")
|
||||||
|
author = _author_label(note)
|
||||||
|
if author:
|
||||||
|
text.append(f" by {author}", style="dim")
|
||||||
if content:
|
if content:
|
||||||
text.append("\n ")
|
text.append("\n ")
|
||||||
text.append(content, style="dim")
|
text.append(content, style="dim")
|
||||||
|
|||||||
@@ -431,3 +431,117 @@ class CreateDependencyReportRenderer(BaseToolRenderer):
|
|||||||
|
|
||||||
css_classes = cls.get_css_classes("completed")
|
css_classes = cls.get_css_classes("completed")
|
||||||
return Static(padded, classes=css_classes)
|
return Static(padded, classes=css_classes)
|
||||||
|
|
||||||
|
|
||||||
|
_LIST_SEVERITY_COLORS = {
|
||||||
|
"critical": "#dc2626",
|
||||||
|
"high": "#ea580c",
|
||||||
|
"medium": "#d97706",
|
||||||
|
"low": "#65a30d",
|
||||||
|
"info": "#0284c7",
|
||||||
|
"none": "#6b7280",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _severity_style(severity: Any) -> str:
|
||||||
|
return _LIST_SEVERITY_COLORS.get(str(severity or "").lower(), "#d97706")
|
||||||
|
|
||||||
|
|
||||||
|
def _author_label(report: dict[str, Any]) -> str:
|
||||||
|
if report.get("by_you"):
|
||||||
|
return "you"
|
||||||
|
agent_name = report.get("agent_name")
|
||||||
|
return str(agent_name).strip() if agent_name else ""
|
||||||
|
|
||||||
|
|
||||||
|
@register_tool_renderer
|
||||||
|
class ListReportsRenderer(BaseToolRenderer):
|
||||||
|
tool_name: ClassVar[str] = "list_reports"
|
||||||
|
css_classes: ClassVar[list[str]] = ["tool-call", "reporting-tool"]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def render(cls, tool_data: dict[str, Any]) -> Static:
|
||||||
|
result = _coerce_dict(tool_data.get("result"))
|
||||||
|
|
||||||
|
text = Text()
|
||||||
|
text.append("◆ ", style="#ef4444")
|
||||||
|
text.append("reports", style="dim")
|
||||||
|
|
||||||
|
if isinstance(tool_data.get("result"), str) and str(tool_data["result"]).strip():
|
||||||
|
text.append("\n ")
|
||||||
|
text.append(str(tool_data["result"]).strip(), style="dim")
|
||||||
|
elif result.get("success"):
|
||||||
|
total = result.get("total_count", 0)
|
||||||
|
reports = _coerce_list_of_dicts(result.get("reports"))
|
||||||
|
counts = _coerce_dict(result.get("severity_counts"))
|
||||||
|
|
||||||
|
text.append(f" ({total})", style="dim")
|
||||||
|
for sev, count in counts.items():
|
||||||
|
text.append(" ")
|
||||||
|
text.append(f"{sev} {count}", style=_severity_style(sev))
|
||||||
|
|
||||||
|
if not reports:
|
||||||
|
text.append("\n ")
|
||||||
|
text.append("No reports filed yet", style="dim")
|
||||||
|
else:
|
||||||
|
for report in reports:
|
||||||
|
rid = str(report.get("id", "")).strip()
|
||||||
|
title = str(report.get("title", "")).strip() or "(untitled)"
|
||||||
|
severity = str(report.get("severity", "")).strip()
|
||||||
|
text.append("\n - ")
|
||||||
|
if severity:
|
||||||
|
text.append(severity.upper(), style=f"bold {_severity_style(severity)}")
|
||||||
|
text.append(" ")
|
||||||
|
if rid:
|
||||||
|
text.append(f"{rid} ", style="dim")
|
||||||
|
text.append(title)
|
||||||
|
author = _author_label(report)
|
||||||
|
if author:
|
||||||
|
text.append(f" ({author})", style="dim")
|
||||||
|
else:
|
||||||
|
text.append("\n ")
|
||||||
|
text.append("Loading...", style="dim")
|
||||||
|
|
||||||
|
css_classes = cls.get_css_classes("completed")
|
||||||
|
return Static(text, classes=css_classes)
|
||||||
|
|
||||||
|
|
||||||
|
@register_tool_renderer
|
||||||
|
class GetReportRenderer(BaseToolRenderer):
|
||||||
|
tool_name: ClassVar[str] = "get_report"
|
||||||
|
css_classes: ClassVar[list[str]] = ["tool-call", "reporting-tool"]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def render(cls, tool_data: dict[str, Any]) -> Static:
|
||||||
|
result = _coerce_dict(tool_data.get("result"))
|
||||||
|
|
||||||
|
text = Text()
|
||||||
|
text.append("◆ ", style="#ef4444")
|
||||||
|
text.append("report read", style="dim")
|
||||||
|
|
||||||
|
report = _coerce_dict(result.get("report")) if result.get("success") else {}
|
||||||
|
if report:
|
||||||
|
rid = str(report.get("id", "")).strip()
|
||||||
|
title = str(report.get("title", "")).strip() or "(untitled)"
|
||||||
|
severity = str(report.get("severity", "")).strip()
|
||||||
|
text.append("\n ")
|
||||||
|
if severity:
|
||||||
|
text.append(severity.upper(), style=f"bold {_severity_style(severity)}")
|
||||||
|
text.append(" ")
|
||||||
|
if rid:
|
||||||
|
text.append(f"{rid} ", style="dim")
|
||||||
|
text.append(title)
|
||||||
|
author = _author_label(report)
|
||||||
|
if author:
|
||||||
|
text.append(f" ({author})", style="dim")
|
||||||
|
target = str(report.get("target", "")).strip()
|
||||||
|
if target:
|
||||||
|
text.append("\n ")
|
||||||
|
text.append(target, style="dim")
|
||||||
|
else:
|
||||||
|
text.append("\n ")
|
||||||
|
detail = result.get("error") if result.get("success") is False else None
|
||||||
|
text.append(str(detail) if detail else "Loading...", style="dim")
|
||||||
|
|
||||||
|
css_classes = cls.get_css_classes("completed")
|
||||||
|
return Static(text, classes=css_classes)
|
||||||
|
|||||||
@@ -271,7 +271,13 @@ def _release_target() -> str | None:
|
|||||||
if os_name is None:
|
if os_name is None:
|
||||||
return None
|
return None
|
||||||
target = f"{os_name}-{arch}"
|
target = f"{os_name}-{arch}"
|
||||||
supported = {"linux-x86_64", "macos-x86_64", "macos-arm64", "windows-x86_64"}
|
supported = {
|
||||||
|
"linux-x86_64",
|
||||||
|
"linux-arm64",
|
||||||
|
"macos-x86_64",
|
||||||
|
"macos-arm64",
|
||||||
|
"windows-x86_64",
|
||||||
|
}
|
||||||
return target if target in supported else None
|
return target if target in supported else None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -11,11 +11,10 @@ import tempfile
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from urllib.error import HTTPError, URLError
|
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
from urllib.request import Request, urlopen
|
|
||||||
|
|
||||||
import docker
|
import docker
|
||||||
|
import requests
|
||||||
from docker.errors import DockerException, ImageNotFound
|
from docker.errors import DockerException, ImageNotFound
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
from rich.panel import Panel
|
from rich.panel import Panel
|
||||||
@@ -1088,13 +1087,12 @@ def resolve_diff_scope_context(
|
|||||||
def _is_http_git_repo(url: str) -> bool:
|
def _is_http_git_repo(url: str) -> bool:
|
||||||
check_url = f"{url.rstrip('/')}/info/refs?service=git-upload-pack"
|
check_url = f"{url.rstrip('/')}/info/refs?service=git-upload-pack"
|
||||||
try:
|
try:
|
||||||
req = Request(check_url, headers={"User-Agent": "git/strix"}) # noqa: S310
|
resp = requests.get(check_url, headers={"User-Agent": "git/strix"}, timeout=10)
|
||||||
with urlopen(req, timeout=10) as resp: # noqa: S310 # nosec B310
|
except (requests.RequestException, ValueError):
|
||||||
return "x-git-upload-pack-advertisement" in resp.headers.get("Content-Type", "")
|
|
||||||
except HTTPError as e:
|
|
||||||
return e.code == 401
|
|
||||||
except (URLError, OSError, ValueError):
|
|
||||||
return False
|
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
|
def infer_target_type(target: str) -> tuple[str, dict[str, str]]: # noqa: PLR0911
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ directly from the run's on-disk files. No cloud dependency, no file picker.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from strix.viewer.server import serve
|
from strix.interface.viewer.server import serve
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["serve"]
|
__all__ = ["serve"]
|
||||||
@@ -15,12 +15,12 @@ import base64
|
|||||||
import contextlib
|
import contextlib
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import urllib.error
|
|
||||||
import urllib.request
|
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
from strix.config.loader import load_settings
|
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.
|
map, not raised.
|
||||||
"""
|
"""
|
||||||
url = f"{_app_url()}{path}"
|
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:
|
try:
|
||||||
with urllib.request.urlopen(request, timeout=timeout) as response: # noqa: S310
|
response = requests.post(
|
||||||
return response.status, _parse_body(response.read())
|
url,
|
||||||
except urllib.error.HTTPError as exc:
|
json=payload,
|
||||||
return exc.code, _parse_body(exc.read())
|
headers={"Accept": "application/json"},
|
||||||
except (urllib.error.URLError, TimeoutError, OSError) as exc:
|
timeout=timeout,
|
||||||
|
)
|
||||||
|
except requests.RequestException as exc:
|
||||||
logger.warning("relay request to %s failed: %s", path, exc)
|
logger.warning("relay request to %s failed: %s", path, exc)
|
||||||
raise RelayError("unavailable") from exc
|
raise RelayError("unavailable") from exc
|
||||||
|
return response.status_code, _parse_body(response.content)
|
||||||
|
|
||||||
|
|
||||||
def _parse_body(raw: bytes) -> dict[str, Any]:
|
def _parse_body(raw: bytes) -> dict[str, Any]:
|
||||||
@@ -16,8 +16,8 @@ from strix.core.paths import (
|
|||||||
run_record_path,
|
run_record_path,
|
||||||
runs_base_dir,
|
runs_base_dir,
|
||||||
)
|
)
|
||||||
from strix.viewer.server import authorized_url, bundle_is_built, serve
|
from strix.interface.viewer.server import authorized_url, bundle_is_built, serve
|
||||||
from strix.viewer.transcript import read_run_summary
|
from strix.interface.viewer.transcript import read_run_summary
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -58,7 +58,7 @@ def run_view(argv: list[str]) -> None:
|
|||||||
if not bundle_is_built():
|
if not bundle_is_built():
|
||||||
console.print(
|
console.print(
|
||||||
"[bold red]Viewer UI is not built.[/]\n"
|
"[bold red]Viewer UI is not built.[/]\n"
|
||||||
"Build it with: [cyan]cd strix/viewer/frontend && npm ci && npm run build[/]"
|
"Build it with: [cyan]cd strix/interface/viewer/frontend && npm ci && npm run build[/]"
|
||||||
)
|
)
|
||||||
raise SystemExit(1)
|
raise SystemExit(1)
|
||||||
|
|
||||||
Generated
|
Before Width: | Height: | Size: 3.7 KiB After Width: | Height: | Size: 3.7 KiB |
+6
@@ -49,6 +49,9 @@ export default function NotesRenderer({ toolName, args, result }: ToolRendererPr
|
|||||||
<div className="mt-1.5 text-[#999] text-[13px]">
|
<div className="mt-1.5 text-[#999] text-[13px]">
|
||||||
{note.title ?? "(untitled)"}
|
{note.title ?? "(untitled)"}
|
||||||
<span className="text-[#555] ml-1">({note.category ?? "general"})</span>
|
<span className="text-[#555] ml-1">({note.category ?? "general"})</span>
|
||||||
|
{(note.by_you || note.agent_name) && (
|
||||||
|
<span className="text-[#666] ml-1 text-xs">by {note.by_you ? "you" : note.agent_name}</span>
|
||||||
|
)}
|
||||||
</div>
|
</div>
|
||||||
{note.content && <div className="mt-1"><Markdown text={note.content} /></div>}
|
{note.content && <div className="mt-1"><Markdown text={note.content} /></div>}
|
||||||
</>
|
</>
|
||||||
@@ -74,6 +77,9 @@ export default function NotesRenderer({ toolName, args, result }: ToolRendererPr
|
|||||||
<span className="text-[#555] mr-1">-</span>
|
<span className="text-[#555] mr-1">-</span>
|
||||||
<span className="text-[#999]">{n.title ?? "(untitled)"}</span>
|
<span className="text-[#999]">{n.title ?? "(untitled)"}</span>
|
||||||
<span className="text-[#555] ml-1">({n.category ?? "general"})</span>
|
<span className="text-[#555] ml-1">({n.category ?? "general"})</span>
|
||||||
|
{(n.by_you || n.agent_name) && (
|
||||||
|
<span className="text-[#666] ml-1 text-xs">by {n.by_you ? "you" : n.agent_name}</span>
|
||||||
|
)}
|
||||||
{n.content && <div className="ml-3"><Markdown text={n.content} /></div>}
|
{n.content && <div className="ml-3"><Markdown text={n.content} /></div>}
|
||||||
</div>
|
</div>
|
||||||
))}
|
))}
|
||||||
+121
@@ -0,0 +1,121 @@
|
|||||||
|
"use client";
|
||||||
|
|
||||||
|
import type { ToolRendererProps } from "@/types/events";
|
||||||
|
import { TruncatedText } from "./ToolCard";
|
||||||
|
import Markdown from "./Markdown";
|
||||||
|
|
||||||
|
const SEVERITY_COLORS: Record<string, string> = {
|
||||||
|
critical: "text-red-400", high: "text-orange-400", medium: "text-yellow-400",
|
||||||
|
low: "text-blue-400", info: "text-cyan-400", none: "text-[#888]",
|
||||||
|
};
|
||||||
|
|
||||||
|
interface ReportEntry {
|
||||||
|
id?: string;
|
||||||
|
title?: string;
|
||||||
|
severity?: string;
|
||||||
|
cvss?: number;
|
||||||
|
cve?: string;
|
||||||
|
cwe?: string;
|
||||||
|
target?: string;
|
||||||
|
endpoint?: string;
|
||||||
|
method?: string;
|
||||||
|
description_preview?: string;
|
||||||
|
description?: string;
|
||||||
|
agent_name?: string;
|
||||||
|
by_you?: boolean;
|
||||||
|
}
|
||||||
|
|
||||||
|
function authorTag(r: ReportEntry) {
|
||||||
|
if (!r.agent_name && !r.by_you) return null;
|
||||||
|
const label = r.by_you ? "you" : r.agent_name;
|
||||||
|
return <span className="text-[#666] text-xs ml-1.5">({label})</span>;
|
||||||
|
}
|
||||||
|
|
||||||
|
function sevBadge(severity: string | undefined) {
|
||||||
|
const sev = String(severity ?? "").toLowerCase();
|
||||||
|
const color = SEVERITY_COLORS[sev] ?? "text-yellow-400";
|
||||||
|
return <span className={`font-semibold text-[13px] ${color}`}>{sev.toUpperCase() || "—"}</span>;
|
||||||
|
}
|
||||||
|
|
||||||
|
export default function ReportListRenderer({ toolName, result }: ToolRendererProps) {
|
||||||
|
const res = result as Record<string, unknown> | null;
|
||||||
|
const ok = res != null && typeof res === "object" && res.success === true;
|
||||||
|
|
||||||
|
if (toolName === "get_report") {
|
||||||
|
const report = ok ? (res.report as ReportEntry | undefined) : undefined;
|
||||||
|
return (
|
||||||
|
<div>
|
||||||
|
<span className="text-red-400/80 font-semibold text-sm">report</span>
|
||||||
|
{report ? (
|
||||||
|
<div className="mt-1.5 space-y-2">
|
||||||
|
<div className="flex items-center gap-2 flex-wrap">
|
||||||
|
{sevBadge(report.severity)}
|
||||||
|
{report.cvss != null && <span className="text-[#888] text-[13px]">CVSS {report.cvss}</span>}
|
||||||
|
{report.id && <span className="text-[#555] font-mono text-[13px]">{report.id}</span>}
|
||||||
|
{report.cve && <span className="text-[#888] font-mono text-[13px]">{report.cve}</span>}
|
||||||
|
{report.cwe && <span className="text-[#888] font-mono text-[13px]">{report.cwe}</span>}
|
||||||
|
{(report.agent_name || report.by_you) && (
|
||||||
|
<span className="text-[#666] text-[13px]">{report.by_you ? "you" : report.agent_name}</span>
|
||||||
|
)}
|
||||||
|
</div>
|
||||||
|
{report.title && <div className="text-[15px] text-white/80 font-semibold">{report.title}</div>}
|
||||||
|
{(report.target || report.endpoint) && (
|
||||||
|
<div className="text-[13px] text-[#888] font-mono">
|
||||||
|
{report.target}{report.endpoint ? ` ${report.method ?? ""} ${report.endpoint}` : ""}
|
||||||
|
</div>
|
||||||
|
)}
|
||||||
|
{report.description && <TruncatedText text={report.description} maxLines={20} />}
|
||||||
|
</div>
|
||||||
|
) : (
|
||||||
|
<div className="mt-1 text-[#555] text-xs">
|
||||||
|
{(res && typeof res === "object" && (res.error as string)) || "Report not found"}
|
||||||
|
</div>
|
||||||
|
)}
|
||||||
|
</div>
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
// list_reports
|
||||||
|
const rawReports = ok ? res.reports : null;
|
||||||
|
const reports: ReportEntry[] = Array.isArray(rawReports) ? (rawReports as ReportEntry[]) : [];
|
||||||
|
const total = ok && typeof res.total_count === "number" ? (res.total_count as number) : reports.length;
|
||||||
|
const counts = ok && res.severity_counts && typeof res.severity_counts === "object"
|
||||||
|
? (res.severity_counts as Record<string, number>)
|
||||||
|
: {};
|
||||||
|
const countEntries = Object.entries(counts);
|
||||||
|
|
||||||
|
return (
|
||||||
|
<div>
|
||||||
|
<div className="flex items-center gap-2 flex-wrap">
|
||||||
|
<span className="text-red-400/80 font-semibold text-sm">reports</span>
|
||||||
|
<span className="text-[#555] text-[13px]">({total})</span>
|
||||||
|
{countEntries.map(([sev, n]) => (
|
||||||
|
<span key={sev} className="text-[13px]">
|
||||||
|
{sevBadge(sev)}<span className="text-[#888] ml-0.5">{n}</span>
|
||||||
|
</span>
|
||||||
|
))}
|
||||||
|
</div>
|
||||||
|
{reports.length > 0 ? (
|
||||||
|
<div className="mt-1.5 space-y-1">
|
||||||
|
{reports.map((r, i) => (
|
||||||
|
<div key={r.id ?? i} className="text-[13px]">
|
||||||
|
<span className="text-[#555] mr-1">-</span>
|
||||||
|
{sevBadge(r.severity)}
|
||||||
|
{r.id && <span className="text-[#555] font-mono ml-1.5">{r.id}</span>}
|
||||||
|
<span className="text-[#999] ml-1.5">{r.title ?? "(untitled)"}</span>
|
||||||
|
{authorTag(r)}
|
||||||
|
{(r.target || r.endpoint) && (
|
||||||
|
<div className="ml-3 text-[#666] font-mono text-xs">
|
||||||
|
{r.target}{r.endpoint ? ` ${r.method ?? ""} ${r.endpoint}` : ""}
|
||||||
|
</div>
|
||||||
|
)}
|
||||||
|
{r.description_preview && (
|
||||||
|
<div className="ml-3"><Markdown text={r.description_preview} /></div>
|
||||||
|
)}
|
||||||
|
</div>
|
||||||
|
))}
|
||||||
|
</div>
|
||||||
|
) : <div className="mt-1 text-[#555] text-xs">No reports filed yet</div>}
|
||||||
|
</div>
|
||||||
|
);
|
||||||
|
}
|
||||||
+4
-1
@@ -12,6 +12,7 @@ import FileEditRenderer from "./FileEditRenderer";
|
|||||||
import ApplyPatchRenderer from "./ApplyPatchRenderer";
|
import ApplyPatchRenderer from "./ApplyPatchRenderer";
|
||||||
import ViewImageRenderer from "./ViewImageRenderer";
|
import ViewImageRenderer from "./ViewImageRenderer";
|
||||||
import VulnReportRenderer from "./VulnReportRenderer";
|
import VulnReportRenderer from "./VulnReportRenderer";
|
||||||
|
import ReportListRenderer from "./ReportListRenderer";
|
||||||
import ProxyRenderer from "./ProxyRenderer";
|
import ProxyRenderer from "./ProxyRenderer";
|
||||||
import ThinkRenderer from "./ThinkRenderer";
|
import ThinkRenderer from "./ThinkRenderer";
|
||||||
import AgentCommsRenderer from "./AgentCommsRenderer";
|
import AgentCommsRenderer from "./AgentCommsRenderer";
|
||||||
@@ -101,7 +102,7 @@ const CATEGORY_TOOLS: Record<ToolCategory, readonly string[]> = {
|
|||||||
filesystem: ["apply_patch", "view_image", "str_replace_editor", "list_files", "search_files"],
|
filesystem: ["apply_patch", "view_image", "str_replace_editor", "list_files", "search_files"],
|
||||||
// Caido proxy tools (legacy: send_request)
|
// Caido proxy tools (legacy: send_request)
|
||||||
proxy: ["list_requests", "view_request", "repeat_request", "list_sitemap", "view_sitemap_entry", "scope_rules", "send_request"],
|
proxy: ["list_requests", "view_request", "repeat_request", "list_sitemap", "view_sitemap_entry", "scope_rules", "send_request"],
|
||||||
reporting: ["create_vulnerability_report"],
|
reporting: ["create_vulnerability_report", "list_reports", "get_report"],
|
||||||
thinking: ["think"],
|
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_message", "view_agent_graph", "stop_agent"],
|
||||||
search: ["web_search"],
|
search: ["web_search"],
|
||||||
@@ -128,6 +129,8 @@ const RENDERER_OVERRIDES: Partial<Record<string, ComponentType<ToolRendererProps
|
|||||||
finish_scan: FinishRenderer,
|
finish_scan: FinishRenderer,
|
||||||
apply_patch: ApplyPatchRenderer,
|
apply_patch: ApplyPatchRenderer,
|
||||||
view_image: ViewImageRenderer,
|
view_image: ViewImageRenderer,
|
||||||
|
list_reports: ReportListRenderer,
|
||||||
|
get_report: ReportListRenderer,
|
||||||
};
|
};
|
||||||
|
|
||||||
/**
|
/**
|
||||||
+1
-1
@@ -5,7 +5,7 @@ import { fileURLToPath, URL } from "node:url";
|
|||||||
|
|
||||||
// The viewer is served as static files by a stdlib Python server on an
|
// The viewer is served as static files by a stdlib Python server on an
|
||||||
// arbitrary ephemeral port, so all asset URLs must be relative (base: "./").
|
// arbitrary ephemeral port, so all asset URLs must be relative (base: "./").
|
||||||
// The build output is committed at strix/viewer/static and shipped.
|
// The build output is committed at strix/interface/viewer/static and shipped.
|
||||||
export default defineConfig({
|
export default defineConfig({
|
||||||
base: "./",
|
base: "./",
|
||||||
plugins: [react(), tailwindcss()],
|
plugins: [react(), tailwindcss()],
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user