Compare commits

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

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

* test(cost): resolve OpenRouter handler via LiteLLM provider pipeline
2026-07-28 23:28:34 -04:00
b313d78f60 Scope viewer session cookie to the bound port (#922)
Co-authored-by: Jonathan Singer <jonathansinger@Mac-4078.lan>
2026-07-27 20:37:54 -04:00
e037d8d727 fix: recoverable guardrail blocks and decoupled crash-notify (#919)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-27 16:28:48 -07:00
alex sandGitHub fade37025d fix viewer tool call collisions across agents (#917) 2026-07-27 18:51:02 -04:00
Ahmed AllamandAhmed Allam f968f8e5a7 fix(cli): align View label spacing in final panel 2026-07-27 15:41:22 -07:00
Ahmed AllamandAhmed Allam ac0014fe65 chore: release v1.4.1 2026-07-27 12:57:39 -07:00
86282e83a8 fix(tls): replace raw urllib with requests for external HTTPS calls (frozen-build cert failures) (#903)
Co-authored-by: Jonathan Singer <jonathansinger@Mac-4051.lan>
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-27 12:34:26 -07:00
Ahmed AllamandAhmed Allam 37c7f5a6ba chore: release v1.4.0 2026-07-27 04:55:21 -07:00
Ahmed AllamandAhmed Allam 082d4ae62c fix(runtime): wake parent when child hits a terminal state (MaxTurnsExceeded) 2026-07-27 04:27:21 -07:00
c55a8fa4ba feat(runtime): graduated wrap-up warnings, budget reserve, and interactive budget pause/continue (#893)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-26 20:37:14 -07:00
devin-ai-integration[bot]andGitHub 47617969d3 fix(cli): don't dump raw warm-up traceback over the LLM error panel (#896) 2026-07-26 20:01:20 -07:00
27f9750cdc feat(llm): enable Bedrock/Anthropic prompt caching for Claude models (#772)
Co-authored-by: Sean Turner <sean.turner@zerohash.com>
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-26 17:12:57 -07:00
Matthew BrightandGitHub 427cdcd9d4 Add Linux ARM64 standalone release support (#886) 2026-07-26 16:27:02 -07:00
384338cf31 fix(runtime): retry transient mid-stream provider errors instead of crashing the scan (#891)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-26 16:23:31 -07:00
48 changed files with 3884 additions and 225 deletions
+16
View File
@@ -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
+14
View File
@@ -19,6 +19,14 @@ Configure Strix using environment variables or a config file.
Custom API base URL. Also accepts `OPENAI_API_BASE`, `LITELLM_BASE_URL`, or `OLLAMA_API_BASE`. 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.
+17
View File
@@ -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
View File
@@ -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
+2 -1
View File
@@ -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"
@@ -220,6 +220,7 @@ 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.interface.viewer.report_pdf. # circular dependency with strix.telemetry / strix.interface.viewer.report_pdf.
+1 -1
View File
@@ -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}"
+13 -20
View File
@@ -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
View File
@@ -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
+16
View File
@@ -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,10 @@ 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): class ContextSettings(BaseSettings):
+86 -5
View File
@@ -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())
+218 -41
View File
@@ -13,9 +13,20 @@ 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,
@@ -44,6 +55,23 @@ _INPUT_REJECTION_CODES = frozenset({400, 404, 422})
_MAX_COMPACTIONS_PER_CYCLE = 2 _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: def _run_config_model(run_config: RunConfig) -> str | None:
return run_config.model if isinstance(run_config.model, str) else None return run_config.model if isinstance(run_config.model, str) else None
@@ -78,6 +106,39 @@ async def _compact_session(
) )
_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(
*, *,
agent: Any, agent: Any,
@@ -100,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,
@@ -142,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(
@@ -248,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,
@@ -260,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(
@@ -327,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,
@@ -357,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."
@@ -387,6 +472,7 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
) -> RunResultBase | None: ) -> RunResultBase | None:
image_strips = 0 image_strips = 0
compactions = 0 compactions = 0
model_retries = 0
while True: while True:
try: try:
await coordinator.mark_running(agent_id) await coordinator.mark_running(agent_id)
@@ -421,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):
@@ -442,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
@@ -488,6 +583,31 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
) )
input_data = [] input_data = []
continue 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):
@@ -498,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,
@@ -562,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)
@@ -578,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],
@@ -639,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
View File
@@ -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."
)
+35
View File
@@ -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,
+31 -3
View File
@@ -31,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,
@@ -40,6 +40,7 @@ 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 ( from strix.tools.output_store import (
@@ -185,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
@@ -236,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,
@@ -244,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)
@@ -355,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,
@@ -366,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,
) )
+2
View File
@@ -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()
+55 -9
View File
@@ -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")
+37 -9
View File
@@ -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()
+7 -5
View File
@@ -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 -1
View File
@@ -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
+6 -8
View File
@@ -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
+10 -14
View File
@@ -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 # nosec B310 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]:
+11 -4
View File
@@ -107,8 +107,11 @@ def resolve_run_dir(base_dir: Path, run_param: str | None, default_run_dir: Path
return candidate return candidate
# Name of the cookie carrying the per-process session capability. # Prefix of the cookie carrying the per-process session capability. The bound
SESSION_COOKIE = "strix_viewer_session" # port is appended (``strix_viewer_session_<port>``) because browsers scope
# cookies by host only, never by port: concurrent viewers on 127.0.0.1 would
# otherwise share one cookie slot and clobber each other's session.
SESSION_COOKIE_PREFIX = "strix_viewer_session"
class _ViewerState: class _ViewerState:
@@ -135,6 +138,9 @@ class _ViewerState:
# enough to steer a live scan, trigger a report, or browse history -- # enough to steer a live scan, trigger a report, or browse history --
# the token is never handed to a caller who merely reaches ``/``. # the token is never handed to a caller who merely reaches ``/``.
self.session_token = secrets.token_urlsafe(32) self.session_token = secrets.token_urlsafe(32)
# Finalized in ``serve()`` once the port is known (the server binds
# after this state is constructed); see SESSION_COOKIE_PREFIX.
self.cookie_name = SESSION_COOKIE_PREFIX
def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]: def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]:
@@ -476,7 +482,7 @@ def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]:
the browser this process handed the page to can pass. A direct the browser this process handed the page to can pass. A direct
caller on an exposed port has no cookie and is rejected. caller on an exposed port has no cookie and is rejected.
""" """
supplied = self._cookies().get(SESSION_COOKIE, "") supplied = self._cookies().get(state.cookie_name, "")
return bool(supplied) and secrets.compare_digest(supplied, state.session_token) return bool(supplied) and secrets.compare_digest(supplied, state.session_token)
def _token_presented(self, query: dict[str, list[str]]) -> bool: def _token_presented(self, query: dict[str, list[str]]) -> bool:
@@ -512,7 +518,7 @@ def _make_handler(state: _ViewerState) -> type[BaseHTTPRequestHandler]:
# SameSite=Strict (never sent from a cross-site context). # SameSite=Strict (never sent from a cross-site context).
self.send_header( self.send_header(
"Set-Cookie", "Set-Cookie",
f"{SESSION_COOKIE}={state.session_token}; Path=/; HttpOnly; SameSite=Strict", f"{state.cookie_name}={state.session_token}; Path=/; HttpOnly; SameSite=Strict",
) )
self.end_headers() self.end_headers()
self.wfile.write(content) self.wfile.write(content)
@@ -586,6 +592,7 @@ def serve(
httpd.daemon_threads = True httpd.daemon_threads = True
bound_port = int(httpd.server_address[1]) bound_port = int(httpd.server_address[1])
state.cookie_name = f"{SESSION_COOKIE_PREFIX}_{bound_port}"
url = f"http://{host}:{bound_port}" url = f"http://{host}:{bound_port}"
thread = threading.Thread(target=httpd.serve_forever, name="strix-viewer", daemon=True) thread = threading.Thread(target=httpd.serve_forever, name="strix-viewer", daemon=True)
+44 -12
View File
@@ -12,15 +12,20 @@ from __future__ import annotations
import logging import logging
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import litellm from agents.model_settings import ModelSettings
from agents.models.interface import ModelTracing
from litellm.exceptions import BadRequestError, ContextWindowExceededError from litellm.exceptions import BadRequestError, ContextWindowExceededError
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from strix.config import load_settings from strix.config import load_settings
from strix.config.models import StrixProvider
from strix.core.inputs import make_model_settings
from strix.core.sessions import replace_session_items, session_write_lock from strix.core.sessions import replace_session_items, session_write_lock
from strix.llm.context_budget import context_window, count_tokens, output_limit from strix.llm.context_budget import context_window, count_tokens, output_limit
if TYPE_CHECKING: if TYPE_CHECKING:
from agents.items import ModelResponse
from agents.memory import Session from agents.memory import Session
@@ -268,26 +273,53 @@ def _checkpoint_item(summary: str) -> dict[str, Any]:
} }
def _extract_text(response: ModelResponse) -> str:
parts: list[str] = []
for item in response.output:
if not isinstance(item, ResponseOutputMessage):
continue
parts.extend(
chunk.text
for chunk in item.content
if isinstance(chunk, ResponseOutputText) and chunk.text
)
return "".join(parts)
async def _summarize(model: str, prompt: str, max_tokens: int) -> str | None: async def _summarize(model: str, prompt: str, max_tokens: int) -> str | None:
llm = load_settings().llm llm = load_settings().llm
model_settings = make_model_settings(
None,
model_name=model,
request_timeout=llm.timeout,
prompt_cache=False,
extra_headers=llm.extra_headers,
).resolve(ModelSettings(max_tokens=max_tokens))
try: try:
response = await litellm.acompletion( response = (
model=model, await StrixProvider()
messages=[{"role": "user", "content": prompt}], .get_model(model)
max_tokens=max_tokens, .get_response(
api_key=llm.api_key, system_instructions=None,
api_base=llm.api_base, input=prompt,
timeout=llm.timeout, model_settings=model_settings,
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
previous_response_id=None,
conversation_id=None,
prompt=None,
)
) )
except Exception: except Exception:
logger.exception("compaction summary call failed for model %s", model) logger.exception("compaction summary call failed for model %s", model)
return None return None
try: content = _extract_text(response).strip()
content = response.choices[0].message.content if not content:
except (AttributeError, IndexError, KeyError):
logger.warning("compaction summary returned no content") logger.warning("compaction summary returned no content")
return None return None
return content.strip() if isinstance(content, str) and content.strip() else None return content
async def maybe_compact( async def maybe_compact(
+13 -2
View File
@@ -17,7 +17,14 @@ logger = logging.getLogger(__name__)
# LiteLLM keys models without the routing prefix users type (``openai/``, # LiteLLM keys models without the routing prefix users type (``openai/``,
# ``litellm/``, ``ollama/`` ...). Strip a leading provider segment on lookup. # ``litellm/``, ``ollama/`` ...). Strip a leading provider segment on lookup.
_STRIPPABLE_PREFIXES = ("openai/", "litellm/", "any-llm/", "ollama/", "ollama_chat/") _STRIPPABLE_PREFIXES = (
"openai/",
"chatgpt/",
"litellm/",
"any-llm/",
"ollama/",
"ollama_chat/",
)
_DEFAULT_OUTPUT_TOKENS = 8_192 _DEFAULT_OUTPUT_TOKENS = 8_192
@@ -38,7 +45,11 @@ def _safe_get_model_info(model: str) -> dict[str, Any] | None:
@lru_cache(maxsize=128) @lru_cache(maxsize=128)
def _model_info(model: str) -> dict[str, int]: def _model_info(model: str) -> dict[str, int]:
for candidate in (model, _lookup_key(model)): lookup_key = _lookup_key(model)
# Provider-qualified ChatGPT lookups may start a synchronous device-login
# poll. LiteLLM keys the metadata by the underlying model slug.
candidates = (lookup_key,) if model.startswith("chatgpt/") else (model, lookup_key)
for candidate in candidates:
info = _safe_get_model_info(candidate) info = _safe_get_model_info(candidate)
if info is not None: if info is not None:
return { return {
+8 -3
View File
@@ -51,17 +51,24 @@ def _dedupe_extra_args(dedupe: DedupeSettings) -> dict[str, str]:
def _dedupe_model_settings( def _dedupe_model_settings(
dedupe: DedupeSettings, model_name: str, request_timeout: float | None dedupe: DedupeSettings, model_name: str, request_timeout: float | None
) -> ModelSettings: ) -> ModelSettings:
llm = load_settings().llm
settings = make_model_settings( settings = make_model_settings(
dedupe.reasoning_effort, dedupe.reasoning_effort,
model_name=model_name, model_name=model_name,
force_required_tool_choice=False, force_required_tool_choice=False,
request_timeout=request_timeout, request_timeout=request_timeout,
# The main model's headers apply only when dedupe falls back to the main
# model; a dedicated dedupe model may route to another provider, which
# must never receive the main endpoint's credentials. A dedicated model
# gets its own DEDUPE_LLM_EXTRA_HEADERS instead.
extra_headers=dedupe.extra_headers if dedupe.model else llm.extra_headers,
) )
extra = _dedupe_extra_args(dedupe) extra = _dedupe_extra_args(dedupe)
if extra: if extra:
settings = settings.resolve(ModelSettings(extra_args=extra)) settings = settings.resolve(ModelSettings(extra_args=extra))
return settings return settings
DEDUPE_SYSTEM_PROMPT = """You are an expert vulnerability report deduplication judge. DEDUPE_SYSTEM_PROMPT = """You are an expert vulnerability report deduplication judge.
Your task is to determine if a candidate vulnerability report describes the SAME vulnerability Your task is to determine if a candidate vulnerability report describes the SAME vulnerability
as any existing report. as any existing report.
@@ -347,9 +354,7 @@ async def check_duplicate(
response = await model.get_response( response = await model.get_response(
system_instructions=DEDUPE_SYSTEM_PROMPT, system_instructions=DEDUPE_SYSTEM_PROMPT,
input=user_msg, input=user_msg,
model_settings=_dedupe_model_settings( model_settings=_dedupe_model_settings(dedupe, resolved_model, settings.llm.timeout),
dedupe, resolved_model, settings.llm.timeout
),
tools=[], tools=[],
output_schema=None, output_schema=None,
handoffs=[], handoffs=[],
+74
View File
@@ -1,6 +1,7 @@
import json import json
import logging import logging
import subprocess import subprocess
import threading
from collections.abc import Callable from collections.abc import Callable
from datetime import UTC, datetime from datetime import UTC, datetime
from importlib.metadata import PackageNotFoundError, version from importlib.metadata import PackageNotFoundError, version
@@ -95,6 +96,8 @@ def get_global_report_state() -> Optional["ReportState"]:
def set_global_report_state(report_state: "ReportState") -> None: def set_global_report_state(report_state: "ReportState") -> None:
global _global_report_state # noqa: PLW0603 global _global_report_state # noqa: PLW0603
_global_report_state = report_state _global_report_state = report_state
# New run: drop any streamed-cost entries a prior run left unconsumed.
streamed_openrouter_costs.clear()
class ReportState: class ReportState:
@@ -507,6 +510,72 @@ class ReportState:
self._sync_llm_usage_record() self._sync_llm_usage_record()
def openrouter_stream_cost(usage: Any) -> float | None:
"""Total OpenRouter-reported cost from a raw stream ``usage`` block, or None.
Non-BYOK responses bill everything to ``usage.cost``. BYOK responses put the
OpenRouter fee in ``usage.cost`` (often 0) and the provider charge in
``usage.cost_details.upstream_inference_cost``, so BYOK totals sum the two.
"""
if not isinstance(usage, dict):
return None
total = 0.0
cost = usage.get("cost")
if isinstance(cost, int | float) and cost > 0:
total += float(cost)
if bool(usage.get("is_byok")):
details = usage.get("cost_details")
upstream = details.get("upstream_inference_cost") if isinstance(details, dict) else None
if isinstance(upstream, int | float) and upstream > 0:
total += float(upstream)
return total if total > 0 else None
def _response_id(completion_response: Any) -> str | None:
response_id = getattr(completion_response, "id", None)
if response_id is None and isinstance(completion_response, dict):
response_id = cast("dict[str, Any]", completion_response).get("id")
return response_id if isinstance(response_id, str) and response_id else None
class StreamedOpenRouterCosts:
"""Correlates OpenRouter's per-stream cost from the parser to the cost callback.
LiteLLM rebuilds streamed responses from token-only chunks and drops the
``usage.cost`` OpenRouter reports in its final stream chunk (its non-streamed
path preserves it; streaming snapshots hidden params at stream start). Every
scan streams, so the OpenRouter streaming handler (see strix.config.models)
records the cost here keyed by response id, and the callback takes it back out
for the matching rebuilt response. Entries are removed on read; ``clear()``
runs per scan so nothing accumulates across runs.
"""
def __init__(self) -> None:
self._costs: dict[str, float] = {}
self._lock = threading.Lock()
def remember(self, response_id: Any, usage: Any) -> None:
cost = openrouter_stream_cost(usage)
if cost is None or not (isinstance(response_id, str) and response_id):
return
with self._lock:
self._costs[response_id] = cost
def take(self, completion_response: Any) -> float | None:
response_id = _response_id(completion_response)
if response_id is None:
return None
with self._lock:
return self._costs.pop(response_id, None)
def clear(self) -> None:
with self._lock:
self._costs.clear()
streamed_openrouter_costs = StreamedOpenRouterCosts()
def litellm_cost_callback( def litellm_cost_callback(
kwargs: Any, kwargs: Any,
completion_response: Any, completion_response: Any,
@@ -541,6 +610,11 @@ def litellm_cost_callback(
if cost is None: if cost is None:
cost = _usage_reported_cost(completion_response) cost = _usage_reported_cost(completion_response)
# Recover the exact OpenRouter cost the streaming handler stashed for this
# response — LiteLLM drops it from streamed usage, so nothing above sees it.
if cost is None:
cost = streamed_openrouter_costs.take(completion_response)
if cost is None: if cost is None:
cost = _estimate_response_cost(kwargs, completion_response) cost = _estimate_response_cost(kwargs, completion_response)
+14
View File
@@ -110,6 +110,19 @@ def _apply_log_limits(create_kwargs: dict[str, Any]) -> None:
) )
def _apply_run_labels(create_kwargs: dict[str, Any]) -> None:
run_id = os.getenv("STRIX_RUN_ID")
if not run_id:
return
labels = create_kwargs.setdefault("labels", {})
if not isinstance(labels, dict):
return
labels["strix-run-id"] = run_id
run_type = os.getenv("STRIX_RUN_TYPE")
if run_type:
labels["strix-run-type"] = run_type
class StrixDockerSandboxSession(DockerSandboxSession): class StrixDockerSandboxSession(DockerSandboxSession):
sandbox_network: str = "" sandbox_network: str = ""
@@ -222,6 +235,7 @@ class StrixDockerSandboxClient(DockerSandboxClient):
_apply_sandbox_network(create_kwargs) _apply_sandbox_network(create_kwargs)
_apply_resource_limits(create_kwargs) _apply_resource_limits(create_kwargs)
_apply_log_limits(create_kwargs) _apply_log_limits(create_kwargs)
_apply_run_labels(create_kwargs)
# Strix injection: host bind mounts (e.g. large repos passed via --mount) # Strix injection: host bind mounts (e.g. large repos passed via --mount)
# that bypass the SDK's file-by-file LocalDir copy. # that bypass the SDK's file-by-file LocalDir copy.
+3 -9
View File
@@ -1,9 +1,9 @@
import json
import logging import logging
import urllib.request
from datetime import datetime from datetime import datetime
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import requests
from strix.config import load_settings from strix.config import load_settings
from strix.telemetry._common import ( from strix.telemetry._common import (
SESSION_ID, SESSION_ID,
@@ -37,13 +37,7 @@ def _send(event: str, properties: dict[str, Any]) -> bool:
"distinct_id": SESSION_ID, "distinct_id": SESSION_ID,
"properties": properties, "properties": properties,
} }
req = urllib.request.Request( # noqa: S310 requests.post(f"{_POSTHOG_HOST}/capture/", json=payload, timeout=10)
f"{_POSTHOG_HOST}/capture/",
data=json.dumps(payload).encode(),
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=10): # noqa: S310 # nosec B310
pass
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.debug("posthog send failed for event %s", event, exc_info=True) logger.debug("posthog send failed for event %s", event, exc_info=True)
return False return False
+3 -4
View File
@@ -2,10 +2,11 @@ from __future__ import annotations
import logging import logging
import urllib.parse import urllib.parse
import urllib.request
from datetime import datetime from datetime import datetime
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
import requests
from strix.config import load_settings from strix.config import load_settings
from strix.telemetry._common import ( from strix.telemetry._common import (
SESSION_ID, SESSION_ID,
@@ -42,9 +43,7 @@ def _send(event: str, properties: dict[str, Any]) -> bool:
url = f"{_SCARF_ENDPOINT}{path}" url = f"{_SCARF_ENDPOINT}{path}"
if query: if query:
url = f"{url}?{query}" url = f"{url}?{query}"
req = urllib.request.Request(url, method="POST") # noqa: S310 requests.post(url, timeout=10)
with urllib.request.urlopen(req, timeout=10): # noqa: S310 # nosec B310
pass
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.debug("scarf send failed for event %s", event, exc_info=True) logger.debug("scarf send failed for event %s", event, exc_info=True)
return False return False
+2
View File
@@ -253,6 +253,8 @@ async def finish_scan(
parent_id = inner.get("parent_id") parent_id = inner.get("parent_id")
if coordinator is not None and parent_id is None and me is not None: if coordinator is not None and parent_id is None and me is not None:
active_agents = await coordinator.active_agents_except(me) active_agents = await coordinator.active_agents_except(me)
if active_agents and coordinator.reserve_stopped:
active_agents = []
else: else:
active_agents = [] active_agents = []
+14
View File
@@ -7,8 +7,10 @@ import hashlib
import json import json
import time import time
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from unittest import mock
import pytest import pytest
import requests
from strix.config import codex from strix.config import codex
@@ -52,6 +54,18 @@ def test_authorize_url_carries_pkce_and_client() -> None:
assert "state=st8" in url assert "state=st8" in url
def test_post_form_returns_parsed_body() -> None:
resp = mock.MagicMock()
resp.status_code = 200
resp.content = b'{"access_token": "tok"}'
with mock.patch.object(requests, "post", return_value=resp) as post:
data = codex._post_form({"grant_type": "refresh_token"})
assert data == {"access_token": "tok"}
assert post.call_args.kwargs["timeout"] == codex._TOKEN_TIMEOUT
@pytest.mark.parametrize( @pytest.mark.parametrize(
("value", "expected"), ("value", "expected"),
[ [
+69 -42
View File
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Any
import pytest import pytest
from litellm.exceptions import BadRequestError, ContextWindowExceededError, RateLimitError from litellm.exceptions import BadRequestError, ContextWindowExceededError, RateLimitError
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from strix.config import ContextSettings from strix.config import ContextSettings
from strix.llm import compaction from strix.llm import compaction
@@ -146,17 +147,35 @@ def _patch_budget(monkeypatch: pytest.MonkeyPatch, *, keep_tokens: int, window:
context.auto_compact = True context.auto_compact = True
settings = SimpleNamespace( settings = SimpleNamespace(
context=context, context=context,
llm=SimpleNamespace(api_key=None, api_base=None, timeout=1), llm=SimpleNamespace(api_key=None, api_base=None, timeout=1, extra_headers=None),
) )
monkeypatch.setattr(compaction, "load_settings", lambda: settings) monkeypatch.setattr(compaction, "load_settings", lambda: settings)
def _patch_summary(monkeypatch: pytest.MonkeyPatch, text: str) -> None: def _model_response(text: str) -> Any:
async def fake_acompletion(**_kwargs: Any) -> Any: chunk = ResponseOutputText(annotations=[], text=text, type="output_text")
message = SimpleNamespace(content=text) message = ResponseOutputMessage(
return SimpleNamespace(choices=[SimpleNamespace(message=message)]) id="msg", content=[chunk], role="assistant", status="completed", type="message"
)
return SimpleNamespace(output=[message])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
def _patch_summary(
monkeypatch: pytest.MonkeyPatch, text: str, captured: dict[str, Any] | None = None
) -> None:
class FakeModel:
async def get_response(self, **kwargs: Any) -> Any:
if captured is not None:
captured.update(kwargs)
return _model_response(text)
class FakeProvider:
def get_model(self, model_name: str | None) -> Any:
if captured is not None:
captured["model"] = model_name
return FakeModel()
monkeypatch.setattr(compaction, "StrixProvider", FakeProvider)
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -189,19 +208,38 @@ async def test_maybe_compact_rewrites_and_keeps_pairs(monkeypatch: pytest.Monkey
async def test_maybe_compact_updates_previous_summary(monkeypatch: pytest.MonkeyPatch) -> None: async def test_maybe_compact_updates_previous_summary(monkeypatch: pytest.MonkeyPatch) -> None:
# Window large enough to leave real room for the summary instructions. # Window large enough to leave real room for the summary instructions.
_patch_budget(monkeypatch, keep_tokens=30, window=4_000) _patch_budget(monkeypatch, keep_tokens=30, window=4_000)
captured: dict[str, str] = {} captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "NEW", captured)
async def fake_acompletion(**kwargs: Any) -> Any:
captured["prompt"] = kwargs["messages"][0]["content"]
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="NEW"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
prior = compaction._checkpoint_item("OLD SUMMARY TEXT") prior = compaction._checkpoint_item("OLD SUMMARY TEXT")
session = FakeSession([prior, *_turns(12)]) session = FakeSession([prior, *_turns(12)])
assert await compaction.maybe_compact(session, model="m", force=True) is True assert await compaction.maybe_compact(session, model="m", force=True) is True
assert "OLD SUMMARY TEXT" in captured["prompt"] assert "OLD SUMMARY TEXT" in captured["input"]
@pytest.mark.asyncio
async def test_summarize_routes_through_provider_with_settings(
monkeypatch: pytest.MonkeyPatch,
) -> None:
_patch_budget(monkeypatch, keep_tokens=30, window=4_000)
monkeypatch.setattr(
compaction,
"load_settings",
lambda: SimpleNamespace(
llm=SimpleNamespace(
api_key=None, api_base=None, timeout=1, extra_headers={"X-Feature-Key": "svc"}
)
),
)
captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
assert await compaction._summarize("litellm/openai/some-model", "p", 64) == "S"
assert captured["model"] == "litellm/openai/some-model"
settings = captured["model_settings"]
assert settings.extra_headers == {"X-Feature-Key": "svc"}
assert settings.max_tokens == 64
def test_fit_to_tokens_truncates_oversized_text(monkeypatch: pytest.MonkeyPatch) -> None: def test_fit_to_tokens_truncates_oversized_text(monkeypatch: pytest.MonkeyPatch) -> None:
@@ -233,19 +271,14 @@ def test_summary_output_tokens_capped_at_model_limit(monkeypatch: pytest.MonkeyP
async def test_maybe_compact_bounds_summary_prompt(monkeypatch: pytest.MonkeyPatch) -> None: async def test_maybe_compact_bounds_summary_prompt(monkeypatch: pytest.MonkeyPatch) -> None:
# A tiny window with a huge head must not send an oversized summary request. # A tiny window with a huge head must not send an oversized summary request.
_patch_budget(monkeypatch, keep_tokens=30, window=4_000) _patch_budget(monkeypatch, keep_tokens=30, window=4_000)
captured: dict[str, str] = {} captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
async def fake_acompletion(**kwargs: Any) -> Any:
captured["prompt"] = kwargs["messages"][0]["content"]
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
big_turns = [{"role": "user", "content": "y" * 2_000} for _ in range(50)] big_turns = [{"role": "user", "content": "y" * 2_000} for _ in range(50)]
session = FakeSession(big_turns) session = FakeSession(big_turns)
assert await compaction.maybe_compact(session, model="m") is True assert await compaction.maybe_compact(session, model="m") is True
# count_tokens==len(chars); prompt must fit the model window. # count_tokens==len(chars); prompt must fit the model window.
assert len(captured["prompt"]) <= 4_000 assert len(captured["input"]) <= 4_000
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -256,27 +289,27 @@ async def test_summary_request_fits_when_room_is_below_old_floor(
instructions = len(compaction._SUMMARY_INSTRUCTIONS) instructions = len(compaction._SUMMARY_INSTRUCTIONS)
window = instructions + 64 + 256 + 300 # summary_max(64)+slack(256)+room(300) window = instructions + 64 + 256 + 300 # summary_max(64)+slack(256)+room(300)
_patch_budget(monkeypatch, keep_tokens=30, window=window) _patch_budget(monkeypatch, keep_tokens=30, window=window)
captured: dict[str, str] = {} captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
async def fake_acompletion(**kwargs: Any) -> Any:
captured["prompt"] = kwargs["messages"][0]["content"]
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
session = FakeSession([{"role": "user", "content": "y" * 5_000} for _ in range(20)]) session = FakeSession([{"role": "user", "content": "y" * 5_000} for _ in range(20)])
assert await compaction.maybe_compact(session, model="m") is True assert await compaction.maybe_compact(session, model="m") is True
assert len(captured["prompt"]) <= window assert len(captured["input"]) <= window
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_maybe_compact_skips_when_summary_fails(monkeypatch: pytest.MonkeyPatch) -> None: async def test_maybe_compact_skips_when_summary_fails(monkeypatch: pytest.MonkeyPatch) -> None:
_patch_budget(monkeypatch, keep_tokens=30, window=4_000) _patch_budget(monkeypatch, keep_tokens=30, window=4_000)
async def fake_acompletion(**_kwargs: Any) -> Any: class BoomModel:
raise RuntimeError("boom") async def get_response(self, **_kwargs: Any) -> Any:
raise RuntimeError("boom")
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion) class BoomProvider:
def get_model(self, _model_name: str | None) -> Any:
return BoomModel()
monkeypatch.setattr(compaction, "StrixProvider", BoomProvider)
session = FakeSession(_turns(12)) session = FakeSession(_turns(12))
before = await session.get_items() before = await session.get_items()
@@ -290,17 +323,11 @@ async def test_maybe_compact_skips_when_no_room_to_summarise(
) -> None: ) -> None:
# No room for any head -> no (doomed) summary is attempted. # No room for any head -> no (doomed) summary is attempted.
_patch_budget(monkeypatch, keep_tokens=30, window=200) _patch_budget(monkeypatch, keep_tokens=30, window=200)
called = False captured: dict[str, Any] = {}
_patch_summary(monkeypatch, "S", captured)
async def fake_acompletion(**_kwargs: Any) -> Any:
nonlocal called
called = True
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content="S"))])
monkeypatch.setattr("strix.llm.compaction.litellm.acompletion", fake_acompletion)
session = FakeSession(_turns(12)) session = FakeSession(_turns(12))
before = await session.get_items() before = await session.get_items()
assert await compaction.maybe_compact(session, model="m", force=True) is False assert await compaction.maybe_compact(session, model="m", force=True) is False
assert called is False assert not captured
assert await session.get_items() == before assert await session.get_items() == before
+18
View File
@@ -21,6 +21,24 @@ def test_context_window_strips_provider_prefix() -> None:
assert context_budget.context_window("openai/gpt-4o") == 128_000 assert context_budget.context_window("openai/gpt-4o") == 128_000
def test_context_window_chatgpt_prefix_skips_provider_auth(
monkeypatch: pytest.MonkeyPatch,
) -> None:
context_budget._model_info.cache_clear()
calls: list[str] = []
def _model_info(model: str) -> dict[str, int]:
calls.append(model)
return {"max_input_tokens": 1_050_000, "max_output_tokens": 128_000}
monkeypatch.setattr("strix.llm.context_budget.litellm.get_model_info", _model_info)
try:
assert context_budget.context_window("chatgpt/gpt-5.6-luna") == 1_050_000
assert calls == ["gpt-5.6-luna"]
finally:
context_budget._model_info.cache_clear()
def test_context_window_unmapped_uses_fallback(monkeypatch: pytest.MonkeyPatch) -> None: def test_context_window_unmapped_uses_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
context_budget._model_info.cache_clear() context_budget._model_info.cache_clear()
+105 -2
View File
@@ -7,9 +7,25 @@ from unittest.mock import MagicMock, patch
import litellm import litellm
import pytest import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
from strix.config.models import _configure_litellm_compatibility from strix.config.models import (
from strix.report.state import litellm_cost_callback _configure_litellm_compatibility,
_install_openrouter_stream_cost_capture,
)
from strix.report.state import (
ReportState,
litellm_cost_callback,
openrouter_stream_cost,
set_global_report_state,
streamed_openrouter_costs,
)
@pytest.fixture(autouse=True)
def _clear_streamed_costs() -> None:
streamed_openrouter_costs.clear()
def test_streaming_logging_stays_enabled_for_cost_callback() -> None: def test_streaming_logging_stays_enabled_for_cost_callback() -> None:
@@ -151,3 +167,90 @@ def test_cost_callback_records_nothing_when_no_cost_available() -> None:
litellm_cost_callback({"response_cost": None, "model": "x/y"}, response) litellm_cost_callback({"response_cost": None, "model": "x/y"}, response)
report_state.record_observed_llm_cost.assert_not_called() report_state.record_observed_llm_cost.assert_not_called()
def test_openrouter_stream_cost_extracts_plain_and_byok_totals() -> None:
assert openrouter_stream_cost({"cost": 0.003168}) == pytest.approx(0.003168)
assert openrouter_stream_cost(
{"cost": 0.01, "is_byok": True, "cost_details": {"upstream_inference_cost": 0.2}}
) == pytest.approx(0.21)
# Upstream cost is only added for BYOK responses.
assert openrouter_stream_cost(
{"cost": 0.05, "is_byok": False, "cost_details": {"upstream_inference_cost": 0.04}}
) == pytest.approx(0.05)
assert openrouter_stream_cost({"prompt_tokens": 10}) is None
assert openrouter_stream_cost(None) is None
def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> None:
report_state = MagicMock()
streamed_openrouter_costs.remember("gen-abc", {"cost": 0.42})
# LiteLLM strips cost from the rebuilt streamed usage; only the id survives.
response = SimpleNamespace(id="gen-abc", usage=SimpleNamespace(cost=None), _hidden_params={})
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=ValueError("unknown model")),
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.42)
# The entry is consumed so a later response cannot double-count it.
assert streamed_openrouter_costs.take(response) is None
def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None:
report_state = MagicMock()
streamed_openrouter_costs.remember("gen-xyz", {"cost": 0.9})
response = SimpleNamespace(
id="gen-xyz",
usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15),
_hidden_params={},
)
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", return_value=0.1) as estimate,
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.9)
estimate.assert_not_called()
def test_streamed_openrouter_costs_ignores_entries_without_cost() -> None:
streamed_openrouter_costs.remember("gen-none", {"prompt_tokens": 10})
streamed_openrouter_costs.remember("", {"cost": 0.5})
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-none")) is None
def test_streamed_openrouter_costs_cleared_on_new_run() -> None:
streamed_openrouter_costs.remember("gen-stale", {"cost": 0.7})
set_global_report_state(ReportState.__new__(ReportState))
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stale")) is None
def test_openrouter_stream_handler_records_cost() -> None:
_install_openrouter_stream_cost_capture()
# Resolve the config the way LiteLLM does in production so we prove the
# override is actually reachable through provider resolution, not just as a
# directly-constructed class.
config = ProviderConfigManager.get_provider_chat_config(
model="moonshotai/kimi-k3", provider=LlmProviders.OPENROUTER
)
assert config is not None
assert type(config).__name__ == "_StrixOpenrouterConfig"
handler = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True)
chunk = {
"id": "gen-stream",
"created": 1,
"model": "moonshotai/kimi-k3",
"choices": [{"index": 0, "delta": {"content": None}}],
"usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055},
}
handler.chunk_parser(chunk)
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx(
0.0035055
)
+32
View File
@@ -44,6 +44,38 @@ def test_dedupe_endpoint_sent_per_call() -> None:
assert (settings.extra_args or {})["api_key"] == "dedupe-key" assert (settings.extra_args or {})["api_key"] == "dedupe-key"
def test_dedicated_dedupe_model_uses_own_headers_not_main() -> None:
dedupe = DedupeSettings(
STRIX_DEDUPE_MODEL="deepseek/cheap",
DEDUPE_LLM_EXTRA_HEADERS={"X-Dedupe": "yes"},
)
settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300)
assert settings.extra_headers == {"X-Dedupe": "yes"}
def test_dedicated_dedupe_model_gets_no_main_headers_by_default(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Main": "secret"}))
loader._cached = None
try:
dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap")
settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300)
assert settings.extra_headers is None
finally:
loader._cached = None
def test_fallback_dedupe_inherits_main_headers(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Main": "svc"}))
loader._cached = None
try:
settings = _dedupe_model_settings(DedupeSettings(), "openai/main-model", 300)
assert settings.extra_headers == {"X-Main": "svc"}
finally:
loader._cached = None
def test_dedupe_defaults_are_empty() -> None: def test_dedupe_defaults_are_empty() -> None:
settings = DedupeSettings() settings = DedupeSettings()
assert settings.model is None assert settings.model is None
+326
View File
@@ -0,0 +1,326 @@
"""Tests for LLM_DISABLE_STREAMING: serve the streamed run loop without SSE.
A gateway that rejects ``stream:true`` (or delivers SSE unreliably) breaks the
SDK run loop, which only issues streamed requests. ``_NonStreamingModel`` wraps
the resolved model so each turn makes one non-streaming ``get_response`` and
replays the completed result as a single terminal stream event. A local server
that rejects streamed requests but answers non-streamed ones including a
structured tool call proves the wrapper works where the stock model fails.
"""
from __future__ import annotations
import json
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents import Agent, Runner, function_tool
from agents.model_settings import ModelSettings
from agents.models.interface import Model, ModelProvider, ModelTracing
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from agents.run import RunConfig
from openai import AsyncOpenAI, BadRequestError
from openai.types.responses import (
ResponseCompletedEvent,
ResponseFunctionToolCall,
ResponseOutputMessage,
ResponseOutputText,
)
from strix.config import codex, loader
from strix.config.loader import load_settings
from strix.config.models import StrixProvider, _NonStreamingModel
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterator
def _tool_call_completion() -> dict[str, Any]:
return {
"id": "chatcmpl-1",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "tool_calls",
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "do_thing", "arguments": '{"n": 1}'},
}
],
},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7},
}
def _text_completion() -> dict[str, Any]:
return {
"id": "chatcmpl-2",
"object": "chat.completion",
"created": 0,
"model": "gw-model",
"choices": [
{
"index": 0,
"finish_reason": "stop",
"message": {"role": "assistant", "content": "hello from gateway"},
}
],
"usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
}
_CAPTURED: dict[str, Any] = {}
_PAYLOAD: dict[str, dict[str, Any]] = {"value": _tool_call_completion()}
class _Handler(BaseHTTPRequestHandler):
"""A gateway that only speaks non-streaming Chat Completions."""
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
body = json.loads(self.rfile.read(length) or b"{}")
_CAPTURED.clear()
_CAPTURED.update(body)
if body.get("stream"):
payload = json.dumps(
{"error": {"message": "streaming is not supported by this endpoint"}}
).encode()
self.send_response(400)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
return
payload = json.dumps(_PAYLOAD["value"]).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
@pytest.fixture
def gateway_url() -> Iterator[str]:
_PAYLOAD["value"] = _tool_call_completion()
server = HTTPServer(("127.0.0.1", 0), _Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
server.shutdown()
server.server_close()
def _model(base_url: str) -> OpenAIChatCompletionsModel:
client = AsyncOpenAI(api_key="tok", base_url=base_url)
return OpenAIChatCompletionsModel(model="gw-model", openai_client=client)
def _call_kwargs() -> dict[str, Any]:
return {
"system_instructions": "s",
"input": "hi",
"model_settings": ModelSettings(),
"tools": [],
"output_schema": None,
"handoffs": [],
"tracing": ModelTracing.DISABLED,
"previous_response_id": None,
"conversation_id": None,
"prompt": None,
}
async def _drain(gen: AsyncIterator[Any]) -> list[Any]:
return [event async for event in gen]
@pytest.mark.asyncio
async def test_stock_model_streaming_fails_on_non_streaming_gateway(gateway_url: str) -> None:
# The stock model issues stream:true and the gateway rejects it.
model = _model(gateway_url)
with pytest.raises(BadRequestError, match="streaming is not supported"):
await _drain(model.stream_response(**_call_kwargs()))
assert _CAPTURED["stream"] is True
@pytest.mark.asyncio
async def test_wrapper_streams_tool_call_without_streaming_request(gateway_url: str) -> None:
# The wrapper turns the streamed run-loop call into one non-streaming
# request and replays the completed result as a terminal stream event.
model = _NonStreamingModel(_model(gateway_url))
events = await _drain(model.stream_response(**_call_kwargs()))
assert _CAPTURED.get("stream") is not True
assert len(events) == 1
completed = events[0]
assert isinstance(completed, ResponseCompletedEvent)
tool_call = completed.response.output[0]
assert isinstance(tool_call, ResponseFunctionToolCall)
assert tool_call.name == "do_thing"
assert json.loads(tool_call.arguments) == {"n": 1}
assert completed.response.usage is not None
assert completed.response.usage.total_tokens == 7
@pytest.mark.asyncio
async def test_wrapper_streams_plain_text(gateway_url: str) -> None:
_PAYLOAD["value"] = _text_completion()
model = _NonStreamingModel(_model(gateway_url))
events = await _drain(model.stream_response(**_call_kwargs()))
assert _CAPTURED.get("stream") is not True
message = events[0].response.output[0]
assert isinstance(message, ResponseOutputMessage)
text = message.content[0]
assert isinstance(text, ResponseOutputText)
assert text.text == "hello from gateway"
@pytest.mark.asyncio
async def test_wrapper_get_response_stays_non_streaming(gateway_url: str) -> None:
# The non-streaming path is a plain pass-through to the inner model.
model = _NonStreamingModel(_model(gateway_url))
response = await model.get_response(**_call_kwargs())
assert _CAPTURED.get("stream") is not True
tool_call = response.output[0]
assert isinstance(tool_call, ResponseFunctionToolCall)
assert tool_call.name == "do_thing"
_TURN_STREAM_FLAGS: list[bool] = []
class _MultiTurnHandler(BaseHTTPRequestHandler):
"""Non-streaming gateway: a tool call on turn 1, a final answer on turn 2."""
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
body = json.loads(self.rfile.read(length) or b"{}")
_TURN_STREAM_FLAGS.append(bool(body.get("stream")))
completion = _tool_call_completion() if len(_TURN_STREAM_FLAGS) == 1 else _text_completion()
if len(_TURN_STREAM_FLAGS) > 1:
completion["choices"][0]["message"]["content"] = "all done"
payload = json.dumps(completion).encode()
self.send_response(200)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
@pytest.fixture
def multiturn_url() -> Iterator[str]:
_TURN_STREAM_FLAGS.clear()
server = HTTPServer(("127.0.0.1", 0), _MultiTurnHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/v1"
finally:
server.shutdown()
server.server_close()
@pytest.mark.asyncio
async def test_run_loop_executes_tool_and_completes_without_streaming(multiturn_url: str) -> None:
# The whole streamed agent loop runs against a non-streaming gateway: the
# synthetic terminal event feeds the runner, which executes the tool and
# continues the turn until a final answer.
calls: list[int] = []
@function_tool
def do_thing(n: int) -> str:
calls.append(n)
return f"did {n}"
class _Provider(ModelProvider):
def get_model(self, model_name: str | None) -> Model: # noqa: ARG002
return _NonStreamingModel(_model(multiturn_url))
agent = Agent(name="t", instructions="use the tool", tools=[do_thing], model="gw-model")
result = Runner.run_streamed(
agent, input="please", run_config=RunConfig(model_provider=_Provider())
)
async for _ in result.stream_events():
pass
assert calls == [1] # tool executed with the streamed tool-call args
assert result.final_output == "all done"
assert len(_TURN_STREAM_FLAGS) == 2 # two turns, both...
assert not any(_TURN_STREAM_FLAGS) # ...issued as non-streaming requests
class _DummyModel(Model):
async def get_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
def stream_response(self, *args: Any, **kwargs: Any) -> Any:
raise NotImplementedError
@pytest.fixture
def _reset_settings(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in ("STRIX_LLM", "LLM_DISABLE_STREAMING"):
monkeypatch.delenv(key, raising=False)
monkeypatch.setattr(loader, "_cached", None)
monkeypatch.setattr(loader, "_override", None)
yield
def test_get_model_wraps_when_disabled(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
inner = _DummyModel()
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: inner)
monkeypatch.setenv("LLM_DISABLE_STREAMING", "true")
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert isinstance(model, _NonStreamingModel)
def test_get_model_unwrapped_by_default(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
inner = _DummyModel()
monkeypatch.setattr("strix.config.models.MultiProvider.get_model", lambda *_: inner)
load_settings()
model = StrixProvider().get_model("openai/gpt-4o-mini")
assert model is inner
def test_get_model_does_not_wrap_subscription_model(
monkeypatch: pytest.MonkeyPatch, _reset_settings: None
) -> None:
# Subscription (ChatGPT) models are always streamed and must not be wrapped.
monkeypatch.setattr(codex, "subscription_model", lambda *_: "gpt-5.5")
monkeypatch.setattr(codex, "get_subscription_client", lambda: AsyncOpenAI(api_key="x"))
monkeypatch.setenv("LLM_DISABLE_STREAMING", "true")
load_settings()
model = StrixProvider().get_model("gpt-5.5")
assert not isinstance(model, _NonStreamingModel)
+357
View File
@@ -0,0 +1,357 @@
from __future__ import annotations
import asyncio
from typing import TYPE_CHECKING, Any
from unittest.mock import MagicMock, patch
import pytest
from strix.core import execution
from strix.core.agents import AgentCoordinator
from strix.core.execution import _start_child_runner, run_agent_loop
from strix.core.hooks import BudgetExceededError, ReportUsageHooks
from strix.core.sessions import open_agent_session
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Callable
from pathlib import Path
MAX_BUDGET = 10.0
COST_PER_CALL = 1.0
class _FakeLedger:
def __init__(self) -> None:
self.cost = 0.0
self.calls: list[str] = []
def record_sdk_usage(self, **_kwargs: Any) -> None:
return
def get_total_llm_cost(self) -> float:
return self.cost
class _FakeStream:
def __init__(
self,
*,
ledger: _FakeLedger,
hooks: ReportUsageHooks,
context: dict[str, Any],
agent: Any,
) -> None:
self._ledger = ledger
self._hooks = hooks
self._context = context
self._agent = agent
self.run_loop_exception: BaseException | None = None
self.final_output = None
async def stream_events(self) -> AsyncIterator[Any]:
self._ledger.cost += COST_PER_CALL
self._ledger.calls.append(str(self._context.get("agent_id")))
ctx_wrapper = MagicMock()
ctx_wrapper.context = self._context
try:
await self._hooks.on_llm_end(ctx_wrapper, self._agent, MagicMock())
except Exception as exc: # noqa: BLE001
self.run_loop_exception = exc
items: tuple[Any, ...] = ()
for item in items:
yield item
def cancel(self, mode: str = "immediate") -> None: # noqa: ARG002
return
def _fake_runner(ledger: _FakeLedger) -> Any:
class _FakeRunner:
@staticmethod
def run_streamed(
agent: Any,
input: Any, # noqa: A002, ARG004
*,
run_config: Any, # noqa: ARG004
context: dict[str, Any],
max_turns: int, # noqa: ARG004
session: Any, # noqa: ARG004
hooks: ReportUsageHooks,
) -> _FakeStream:
return _FakeStream(ledger=ledger, hooks=hooks, context=context, agent=agent)
return _FakeRunner
async def _noop_compact(*_args: Any, **_kwargs: Any) -> bool:
return False
async def _wait_until(predicate: Callable[[], bool], *, timeout: float = 5.0) -> None:
async def _poll() -> None:
while not predicate():
await asyncio.sleep(0.01)
await asyncio.wait_for(_poll(), timeout=timeout)
@pytest.mark.asyncio
async def test_full_budget_lifecycle_reserve_then_cap( # noqa: PLR0915
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
ledger = _FakeLedger()
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
db_path = tmp_path / "agents.sqlite"
sessions: list[Any] = []
run_config = MagicMock()
await coordinator.register("root", "strix", parent_id=None)
root_session = open_agent_session("root", db_path)
sessions.append(root_session)
root_exc: list[BaseException] = []
async def _root_loop() -> None:
try:
await run_agent_loop(
agent=MagicMock(),
initial_input=[],
run_config=run_config,
context={"agent_id": "root", "parent_id": None},
max_turns=500,
coordinator=coordinator,
agent_id="root",
interactive=True,
session=root_session,
start_parked=True,
hooks=hooks,
)
except BaseException as exc:
root_exc.append(exc)
raise
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
root_task = asyncio.create_task(_root_loop())
await asyncio.sleep(0.05)
for child_id in ("child-a", "child-b"):
await coordinator.register(child_id, "recon", parent_id="root")
await _start_child_runner(
parent_ctx={"agent_id": "root", "parent_id": None},
coordinator=coordinator,
agents_db_path=db_path,
sessions_to_close=sessions,
run_config=run_config,
max_turns=500,
interactive=True,
child_agent=MagicMock(),
child_id=child_id,
name=f"recon-{child_id}",
parent_id="root",
task="probe things",
initial_input=[],
hooks=hooks,
)
await _wait_until(lambda: ledger.cost >= 2.0)
reserve_before = coordinator.reserve_stopped
assert reserve_before is False
async def _wait_spend_above(amount: float) -> None:
await _wait_until(lambda: ledger.cost > amount)
turn = 0
while ledger.cost < MAX_BUDGET * 0.90 - 1e-9:
target = ("child-a", "child-b")[turn % 2]
spent_before = ledger.cost
assert await coordinator.send(target, {"from": "user", "content": "keep going"})
await _wait_spend_above(spent_before)
turn += 1
await _wait_until(lambda: coordinator.reserve_stopped)
await _wait_until(
lambda: (
coordinator.statuses["child-a"] == "stopped"
and coordinator.statuses["child-b"] == "stopped"
)
)
assert coordinator.reserve_stopped is True
await _wait_until(lambda: coordinator.budget_stopped)
assert ledger.cost == pytest.approx(MAX_BUDGET)
assert len(ledger.calls) == 10
assert set(ledger.calls[:9]) == {"child-a", "child-b"}
assert ledger.calls[9] == "root"
root_items = await root_session.get_items()
notices = [item for item in root_items if "Budget reserve" in str(item)]
assert len(notices) == 1
with pytest.raises(BudgetExceededError):
await root_task
assert root_exc and isinstance(root_exc[0], BudgetExceededError)
assert {aid: str(status) for aid, status in coordinator.statuses.items()} == {
"root": "stopped",
"child-a": "stopped",
"child-b": "stopped",
}
assert coordinator.budget_stopped is True
assert coordinator.reserve_stopped is True
for session in sessions:
session.close()
@pytest.mark.asyncio
async def test_respawned_children_after_reserve_never_spend(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
ledger = _FakeLedger()
ledger.cost = 9.5
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child-a", "recon", parent_id="root")
snap = await coordinator.snapshot()
snap["reserve_stopped"] = True
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.reserve_stopped is True
sessions: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
await _start_child_runner(
parent_ctx={"agent_id": "root", "parent_id": None},
coordinator=restored,
agents_db_path=tmp_path / "agents.sqlite",
sessions_to_close=sessions,
run_config=MagicMock(),
max_turns=500,
interactive=True,
child_agent=MagicMock(),
child_id="child-a",
name="recon-child-a",
parent_id="root",
task="probe things",
initial_input=[],
hooks=hooks,
)
await _wait_until(lambda: restored.statuses["child-a"] == "stopped")
assert ledger.cost == pytest.approx(9.5)
assert ledger.calls == []
for session in sessions:
session.close()
@pytest.mark.asyncio
async def test_resumed_parked_root_after_reserve_is_renotified_and_finalizes(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
ledger = _FakeLedger()
ledger.cost = 9.0
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.set_status("root", "waiting")
snap = await coordinator.snapshot()
snap["reserve_stopped"] = True
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.reserve_stopped is True
root_session = open_agent_session("root", tmp_path / "agents.sqlite")
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
root_task = asyncio.create_task(
run_agent_loop(
agent=MagicMock(),
initial_input=[],
run_config=MagicMock(),
context={"agent_id": "root", "parent_id": None},
max_turns=500,
coordinator=restored,
agent_id="root",
interactive=True,
session=root_session,
start_parked=True,
hooks=hooks,
)
)
with pytest.raises(BudgetExceededError):
await asyncio.wait_for(root_task, timeout=5.0)
assert ledger.calls == ["root"]
assert ledger.cost == pytest.approx(MAX_BUDGET)
root_items = await root_session.get_items()
notices = [item for item in root_items if "Budget reserve" in str(item)]
assert len(notices) == 1
root_session.close()
@pytest.mark.asyncio
async def test_interactive_budget_pause_then_user_message_extends_and_resumes(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
ledger = _FakeLedger()
ledger.cost = 9.0
hooks = ReportUsageHooks(model="test-model", max_budget_usd=MAX_BUDGET, interactive=True)
monkeypatch.setattr(execution, "Runner", _fake_runner(ledger))
monkeypatch.setattr(execution, "_compact_session", _noop_compact)
coordinator = AgentCoordinator()
coordinator.set_budget_extender(hooks.extend_budget)
await coordinator.register("root", "strix", parent_id=None)
root_session = open_agent_session("root", tmp_path / "agents.sqlite")
with patch("strix.core.hooks.get_global_report_state", return_value=ledger):
root_task = asyncio.create_task(
run_agent_loop(
agent=MagicMock(),
initial_input=[],
run_config=MagicMock(),
context={"agent_id": "root", "parent_id": None},
max_turns=500,
coordinator=coordinator,
agent_id="root",
interactive=True,
session=root_session,
start_parked=True,
hooks=hooks,
)
)
await asyncio.sleep(0.05)
assert await coordinator.send("root", {"from": "user", "content": "go"})
await _wait_until(lambda: coordinator.budget_paused)
assert coordinator.statuses["root"] == "budget_paused"
assert ledger.cost == pytest.approx(MAX_BUDGET)
assert not root_task.done()
assert coordinator.budget_stopped is False
assert await coordinator.send("root", {"from": "user", "content": "keep going"})
await _wait_until(lambda: not coordinator.budget_paused)
await _wait_until(lambda: ledger.cost > MAX_BUDGET)
await _wait_until(lambda: coordinator.statuses["root"] == "waiting")
assert not root_task.done()
root_task.cancel()
await root_task
root_session.close()
+624
View File
@@ -3,10 +3,299 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import contextlib
import json
from typing import Any
from unittest.mock import MagicMock
import pytest import pytest
from agents.items import MessageOutputItem
from agents.memory import SQLiteSession
from agents.tool_context import ToolContext
from openai.types.responses import ResponseOutputMessage, ResponseOutputRefusal
from strix.config import codex
from strix.core import execution
from strix.core.agents import AgentCoordinator from strix.core.agents import AgentCoordinator
from strix.core.execution import (
_handle_content_guardrail,
_notify_parent_on_terminal,
_notify_root_on_budget_reserve,
respawn_subagents,
)
from strix.tools.finish.tool import finish_scan
class _StructuredRefusalStream:
def __init__(self, refusal: str) -> None:
self.run_loop_exception: BaseException | None = None
self.new_items = [
MessageOutputItem(
agent=MagicMock(),
raw_item=ResponseOutputMessage(
id="msg-refusal",
content=[ResponseOutputRefusal(type="refusal", refusal=refusal)],
role="assistant",
status="completed",
type="message",
),
)
]
async def stream_events(self) -> Any:
if False:
yield None
def cancel(self, mode: str = "immediate") -> None: # noqa: ARG002
return
async def _call_finish_scan(
coordinator: AgentCoordinator, agent_id: str, parent_id: str | None
) -> dict[str, Any]:
ctx = ToolContext(
context={"coordinator": coordinator, "agent_id": agent_id, "parent_id": parent_id},
tool_name="finish_scan",
tool_call_id="call-1",
tool_arguments="{}",
)
fields = ("executive_summary", "methodology", "technical_analysis", "recommendations")
result: str = await finish_scan.on_invoke_tool(ctx, json.dumps(dict.fromkeys(fields, "x")))
parsed: dict[str, Any] = json.loads(result)
return parsed
@pytest.mark.asyncio
async def test_reserve_stop_notifies_root_once(monkeypatch: pytest.MonkeyPatch) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child-a", "recon", parent_id="root")
await coordinator.register("child-b", "recon", parent_id="root")
sent: list[tuple[str, dict[str, Any]]] = []
async def _record(target_agent_id: str, message: dict[str, Any]) -> bool:
sent.append((target_agent_id, message))
return True
monkeypatch.setattr(coordinator, "send", _record)
await _notify_root_on_budget_reserve(coordinator)
await _notify_root_on_budget_reserve(coordinator)
assert len(sent) == 1
target, message = sent[0]
assert target == "root"
assert message["type"] == "budget_reserve_stop"
assert "finish_scan" in str(message["content"])
@pytest.mark.asyncio
async def test_concurrent_reserve_claims_yield_single_root() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
for i in range(12):
await coordinator.register(f"child-{i}", "recon", parent_id="root")
results = await asyncio.gather(*(coordinator.claim_reserve_notification() for _ in range(12)))
assert results.count("root") == 1
assert all(r is None for r in results if r != "root")
@pytest.mark.asyncio
async def test_claim_reserve_sets_flag_and_wakes_parked_agents() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
flag_before = coordinator.reserve_stopped
assert flag_before is False
waiter = asyncio.create_task(coordinator.wait_for_message("child"))
await asyncio.sleep(0)
assert not waiter.done()
await coordinator.claim_reserve_notification()
flag_after = coordinator.reserve_stopped
assert flag_after is True
await asyncio.wait_for(waiter, timeout=1.0)
@pytest.mark.asyncio
async def test_finish_scan_bypasses_active_agent_guard_after_reserve() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
await coordinator.set_status("child", "running")
blocked = await _call_finish_scan(coordinator, "root", None)
assert blocked["scan_completed"] is False
assert blocked["active_agents"]
await coordinator.claim_reserve_notification()
finished = await _call_finish_scan(coordinator, "root", None)
assert finished["scan_completed"] is True
assert coordinator.statuses["root"] == "completed"
@pytest.mark.asyncio
async def test_finish_scan_gate_ignores_sub_agent_caller() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
await coordinator.set_status("child", "running")
result = await _call_finish_scan(coordinator, "child", "root")
assert "active_agents" not in result
assert result["success"] is False
assert "root" in result["error"]
@pytest.mark.asyncio
async def test_reserve_stop_notify_noop_without_root(monkeypatch: pytest.MonkeyPatch) -> None:
coordinator = AgentCoordinator()
await coordinator.register("child", "recon", parent_id="missing")
sent: list[tuple[str, dict[str, Any]]] = []
async def _record(target_agent_id: str, message: dict[str, Any]) -> bool:
sent.append((target_agent_id, message))
return True
monkeypatch.setattr(coordinator, "send", _record)
await _notify_root_on_budget_reserve(coordinator)
assert sent == []
@pytest.mark.asyncio
async def test_snapshot_round_trip_preserves_stop_flags() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.trigger_budget_stop()
await coordinator.claim_reserve_notification()
snap = await coordinator.snapshot()
assert snap["budget_stopped"] is True
assert snap["reserve_stopped"] is True
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.budget_stopped is True
assert restored.reserve_stopped is True
@pytest.mark.asyncio
async def test_legacy_snapshot_without_stop_flags_defaults_to_false() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
snap = await coordinator.snapshot()
del snap["budget_stopped"]
del snap["reserve_stopped"]
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.budget_stopped is False
assert restored.reserve_stopped is False
@pytest.mark.asyncio
async def test_randomized_reserve_claim_race_many_interleavings() -> None:
for seed in range(25):
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
child_ids = [f"child-{i}" for i in range(8)]
for child_id in child_ids:
await coordinator.register(child_id, "recon", parent_id="root")
waiters = [asyncio.create_task(coordinator.wait_for_message(cid)) for cid in child_ids]
await asyncio.sleep(0)
async def _claim(delay: float, coord: AgentCoordinator = coordinator) -> str | None:
await asyncio.sleep(delay)
return await coord.claim_reserve_notification()
delays = [((seed * 31 + i * 17) % 50) / 10_000 for i in range(len(child_ids))]
results = await asyncio.gather(*(_claim(delay) for delay in delays))
assert results.count("root") == 1, f"seed {seed}: expected exactly one winner"
await asyncio.wait_for(asyncio.gather(*waiters), timeout=1.0)
assert coordinator.reserve_stopped is True
@pytest.mark.asyncio
async def test_reserve_claim_never_loses_root_wake() -> None:
for _ in range(10):
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
root_waiter = asyncio.create_task(coordinator.wait_for_message("root"))
await asyncio.sleep(0)
assert not root_waiter.done()
await coordinator.claim_reserve_notification()
await asyncio.sleep(0)
async with coordinator._lock:
coordinator.pending_counts["root"] = 1
coordinator.runtimes["root"].wake.set()
await asyncio.wait_for(root_waiter, timeout=1.0)
@pytest.mark.asyncio
async def test_budget_stop_takes_precedence_over_reserve_for_all_roles() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
await coordinator.claim_reserve_notification()
await coordinator.trigger_budget_stop()
await asyncio.wait_for(coordinator.wait_for_message("root"), timeout=1.0)
await asyncio.wait_for(coordinator.wait_for_message("child"), timeout=1.0)
assert coordinator.budget_stopped is True
assert coordinator.reserve_stopped is True
@pytest.mark.asyncio
async def test_root_not_released_by_reserve_alone() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
await coordinator.claim_reserve_notification()
root_waiter = asyncio.create_task(coordinator.wait_for_message("root"))
await asyncio.sleep(0.02)
assert not root_waiter.done()
root_waiter.cancel()
with contextlib.suppress(asyncio.CancelledError):
await root_waiter
@pytest.mark.asyncio
async def test_snapshot_during_concurrent_claims_is_consistent() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
for i in range(6):
await coordinator.register(f"child-{i}", "recon", parent_id="root")
claims = [asyncio.create_task(coordinator.claim_reserve_notification()) for _ in range(6)]
snap = await coordinator.snapshot()
await asyncio.gather(*claims)
assert isinstance(snap["reserve_stopped"], bool)
final_snap = await coordinator.snapshot()
assert final_snap["reserve_stopped"] is True
restored = AgentCoordinator()
await restored.restore(final_snap)
assert restored.reserve_stopped is True
assert await restored.claim_reserve_notification() is None
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -42,3 +331,338 @@ async def test_wait_for_message_returns_immediately_after_budget_stop() -> None:
# No pending messages, but the stop flag short-circuits the wait. # No pending messages, but the stop flag short-circuits the wait.
await asyncio.wait_for(coordinator.wait_for_message("agent"), timeout=1.0) await asyncio.wait_for(coordinator.wait_for_message("agent"), timeout=1.0)
@pytest.mark.asyncio
async def test_pause_for_budget_sets_flag_and_status() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.pause_for_budget("root")
assert coordinator.budget_paused is True
assert coordinator.statuses["root"] == "budget_paused"
@pytest.mark.asyncio
async def test_resume_from_budget_pause_extends_and_nudges(
monkeypatch: pytest.MonkeyPatch,
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child-a", "recon", parent_id="root")
await coordinator.register("child-b", "recon", parent_id="root")
await coordinator.pause_for_budget("root")
await coordinator.pause_for_budget("child-a")
await coordinator.pause_for_budget("child-b")
extensions: list[int] = []
coordinator.set_budget_extender(lambda: extensions.append(1))
sent: list[tuple[str, dict[str, Any]]] = []
async def _record(target_agent_id: str, message: dict[str, Any]) -> bool:
sent.append((target_agent_id, message))
return True
monkeypatch.setattr(coordinator, "send", _record)
await coordinator.resume_from_budget_pause(exclude="root")
assert coordinator.budget_paused is False
assert len(extensions) == 1
assert all(coordinator.statuses[aid] == "waiting" for aid in ("root", "child-a", "child-b"))
assert sorted(target for target, _ in sent) == ["child-a", "child-b"]
assert all(message["type"] == "budget_extended" for _, message in sent)
await coordinator.resume_from_budget_pause(exclude="root")
assert len(extensions) == 1
@pytest.mark.asyncio
async def test_user_send_resumes_budget_pause(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
await coordinator.pause_for_budget("root")
extensions: list[int] = []
coordinator.set_budget_extender(lambda: extensions.append(1))
delivered = await coordinator.send("root", {"from": "user", "content": "keep going"})
assert delivered is True
assert coordinator.budget_paused is False
assert len(extensions) == 1
assert coordinator.statuses["root"] == "waiting"
assert coordinator.pending_counts["root"] == 1
session.close()
@pytest.mark.asyncio
async def test_non_user_send_does_not_resume_budget_pause(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
await coordinator.pause_for_budget("root")
extensions: list[int] = []
coordinator.set_budget_extender(lambda: extensions.append(1))
await coordinator.send("root", {"from": "system", "content": "status"})
assert coordinator.budget_paused is True
assert extensions == []
assert coordinator.statuses["root"] == "budget_paused"
session.close()
@pytest.mark.asyncio
async def test_reset_budget_stops_clears_pause_and_normalizes_statuses() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.trigger_budget_stop()
await coordinator.claim_reserve_notification()
await coordinator.pause_for_budget("root")
await coordinator.reset_budget_stops(budget_stopped=False, reserve_stopped=False)
assert coordinator.budget_stopped is False
assert coordinator.reserve_stopped is False
assert coordinator.budget_paused is False
assert coordinator.statuses["root"] == "waiting"
@pytest.mark.asyncio
async def test_reset_budget_stops_can_preserve_pause() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.pause_for_budget("root")
await coordinator.reset_budget_stops(
budget_stopped=False, reserve_stopped=False, budget_paused=True
)
assert coordinator.budget_paused is True
assert coordinator.statuses["root"] == "budget_paused"
@pytest.mark.asyncio
async def test_snapshot_round_trip_preserves_budget_pause() -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.pause_for_budget("root")
snap = await coordinator.snapshot()
assert snap["budget_paused"] is True
restored = AgentCoordinator()
await restored.restore(snap)
assert restored.budget_paused is True
assert restored.statuses["root"] == "budget_paused"
@pytest.mark.asyncio
@pytest.mark.parametrize("status", ["stopped", "failed", "crashed"])
async def test_terminal_child_wakes_parked_parent(tmp_path: Any, status: str) -> None:
# Regression for #870: a child reaching a terminal state (e.g. MaxTurnsExceeded
# -> "stopped") must wake the parent parked in wait_for_message, so the root can
# finalize the scan instead of hanging for a completion report that never arrives.
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "SQL Injection", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
root_waiter = asyncio.create_task(coordinator.wait_for_message("root"))
await asyncio.sleep(0)
assert not root_waiter.done()
await coordinator.set_status("child", status, error="Max turns (500) exceeded")
await _notify_parent_on_terminal(coordinator, "child", status)
await asyncio.wait_for(root_waiter, timeout=1.0)
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_notify_parent_on_terminal_ignores_non_terminal_status(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
await _notify_parent_on_terminal(coordinator, "child", "waiting")
assert coordinator.pending_counts.get("root", 0) == 0
session.close()
class _RecordingStream:
def __init__(self) -> None:
self.cancelled = False
self.cancel_mode: str | None = None
def cancel(self, mode: str = "immediate") -> None:
self.cancelled = True
self.cancel_mode = mode
@pytest.mark.asyncio
async def test_terminal_notice_does_not_cancel_parent_stream(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
stream = _RecordingStream()
await coordinator.attach_runtime("root", session=session, interrupt_on_message=True)
await coordinator.attach_stream("root", stream)
await _notify_parent_on_terminal(coordinator, "child", "crashed")
assert stream.cancelled is False
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_guardrail_interactive_parks_agent_wakeable(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
exc = codex.CodexContentGuardrailError("chatgpt/gpt-5.6-sol")
result = await _handle_content_guardrail(coordinator, "child", exc, interactive=True)
assert result is None
assert coordinator.statuses["child"] == "waiting"
assert "STRIX_LLM" in coordinator.errors["child"]
waiter = asyncio.create_task(coordinator.wait_for_message("child"))
await asyncio.sleep(0)
assert not waiter.done()
session = SQLiteSession("child", tmp_path / "agents.db")
await coordinator.attach_runtime("child", session=session)
await coordinator.send("child", {"from": "user", "content": "switched model, resume"})
await asyncio.wait_for(waiter, timeout=1.0)
session.close()
@pytest.mark.asyncio
async def test_guardrail_noninteractive_fails_only_blocked_agent(tmp_path: Any) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
exc = codex.CodexContentGuardrailError("chatgpt/gpt-5.6-sol")
result = await _handle_content_guardrail(coordinator, "child", exc, interactive=False)
assert result is None
assert coordinator.statuses["child"] == "failed"
assert "STRIX_LLM" in coordinator.errors["child"]
assert coordinator.statuses["root"] == "running"
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_structured_provider_refusal_fails_interactive_agent(
monkeypatch: pytest.MonkeyPatch,
) -> None:
refusal = "This request was blocked under the provider's usage policy."
stream = _StructuredRefusalStream(refusal)
monkeypatch.setattr(execution.Runner, "run_streamed", lambda *_args, **_kwargs: stream)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
result = await execution._run_cycle(
MagicMock(),
coordinator,
"root",
input_data="task",
run_config=MagicMock(),
context={},
max_turns=5,
session=None,
interactive=True,
event_sink=None,
hooks=None,
)
assert result is None
assert coordinator.statuses["root"] == "failed"
assert coordinator.errors["root"] == refusal
@pytest.mark.asyncio
async def test_structured_provider_refusal_fails_noninteractive_child(
tmp_path: Any,
monkeypatch: pytest.MonkeyPatch,
) -> None:
refusal = "This request was blocked under the provider's usage policy."
stream = _StructuredRefusalStream(refusal)
monkeypatch.setattr(execution.Runner, "run_streamed", lambda *_args, **_kwargs: stream)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("child", "recon", parent_id="root")
session = SQLiteSession("root", tmp_path / "agents.db")
await coordinator.attach_runtime("root", session=session)
result = await execution._run_cycle(
MagicMock(),
coordinator,
"child",
input_data="task",
run_config=MagicMock(),
context={"parent_id": "root"},
max_turns=5,
session=None,
interactive=False,
event_sink=None,
hooks=None,
)
assert result is None
assert coordinator.statuses["child"] == "failed"
assert coordinator.errors["child"] == refusal
assert coordinator.pending_counts.get("root", 0) > 0
session.close()
@pytest.mark.asyncio
async def test_resume_revives_guardrail_parked_child_but_not_plain_waiting(
tmp_path: Any, monkeypatch: pytest.MonkeyPatch
) -> None:
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
await coordinator.register("blocked", "recon", parent_id="root")
await coordinator.register("peer_waiter", "recon", parent_id="root")
await coordinator.set_status("blocked", "waiting", error="STRIX_LLM guardrail")
await coordinator.set_status("peer_waiter", "waiting")
parked: dict[str, bool] = {}
async def _fake_start_child_runner(**kwargs: Any) -> None:
parked[kwargs["child_id"]] = bool(kwargs["start_parked"])
monkeypatch.setattr(execution, "_start_child_runner", _fake_start_child_runner)
await respawn_subagents(
coordinator=coordinator,
factory=lambda **_kwargs: object(),
agents_db_path=tmp_path / "agents.db",
sessions_to_close=[],
run_config=MagicMock(),
max_turns=10,
interactive=True,
parent_ctx={"agent_id": "root", "parent_id": None},
root_id="root",
)
assert parked["blocked"] is False
assert parked["peer_waiter"] is True
+157
View File
@@ -0,0 +1,157 @@
from __future__ import annotations
from typing import Any, cast
import httpx
import pytest
from agents import RunConfig, Runner
from openai import (
APIConnectionError,
APIError,
APIStatusError,
APITimeoutError,
BadRequestError,
InternalServerError,
RateLimitError,
)
from strix.core import execution
from strix.core.agents import AgentCoordinator
def _request() -> httpx.Request:
return httpx.Request("POST", "https://api.openai.com/v1/responses")
def _midstream_api_error() -> APIError:
return APIError("An error occurred while processing the request.", _request(), body=None)
def _status_error(status: int) -> APIStatusError:
return APIStatusError(
f"status {status}",
response=httpx.Response(status_code=status, request=_request()),
body=None,
)
def test_midstream_api_error_is_transient() -> None:
assert execution._is_transient_model_error(_midstream_api_error()) is True
def test_network_errors_are_transient() -> None:
assert execution._is_transient_model_error(APITimeoutError(_request())) is True
assert execution._is_transient_model_error(APIConnectionError(request=_request())) is True
def test_server_errors_are_transient() -> None:
assert (
execution._is_transient_model_error(
InternalServerError("boom", response=httpx.Response(500, request=_request()), body=None)
)
is True
)
for status in (502, 503, 504, 408):
assert execution._is_transient_model_error(_status_error(status)) is True
def test_rate_limit_is_not_retried_here() -> None:
rate_limited = RateLimitError(
"slow down", response=httpx.Response(429, request=_request()), body=None
)
assert execution._is_transient_model_error(rate_limited) is False
def test_client_errors_are_not_transient() -> None:
bad_request = BadRequestError(
"bad", response=httpx.Response(400, request=_request()), body=None
)
assert execution._is_transient_model_error(bad_request) is False
assert execution._is_transient_model_error(_status_error(404)) is False
assert execution._is_transient_model_error(ValueError("nope")) is False
class _FakeStream:
def __init__(self, exc: BaseException | None = None) -> None:
self._exc = exc
self._events: list[Any] = []
self.run_loop_exception: BaseException | None = None
async def stream_events(self) -> Any:
if self._exc is not None:
raise self._exc
for event in self._events:
yield event
def _patch_fast_backoff(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(execution, "_TRANSIENT_MODEL_RETRY_BASE_DELAY_S", 0.0)
monkeypatch.setattr(execution, "_TRANSIENT_MODEL_RETRY_MAX_DELAY_S", 0.0)
async def _run_once(
monkeypatch: pytest.MonkeyPatch,
streams: list[_FakeStream],
) -> Any:
_patch_fast_backoff(monkeypatch)
calls = {"n": 0}
def _fake_run_streamed(*_args: Any, **_kwargs: Any) -> _FakeStream:
stream = streams[calls["n"]]
calls["n"] += 1
return stream
monkeypatch.setattr(Runner, "run_streamed", _fake_run_streamed)
coordinator = AgentCoordinator()
await coordinator.register("root", "strix", parent_id=None)
result = await execution._run_cycle(
object(),
coordinator,
"root",
input_data="task",
run_config=cast("RunConfig", object()),
context={},
max_turns=5,
session=None,
interactive=False,
event_sink=None,
hooks=None,
)
return result, calls["n"], coordinator
@pytest.mark.asyncio
async def test_run_cycle_retries_transient_midstream_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
streams = [_FakeStream(exc=_midstream_api_error()), _FakeStream()]
result, attempts, _coordinator = await _run_once(monkeypatch, streams)
assert result is streams[1]
assert attempts == 2
@pytest.mark.asyncio
async def test_run_cycle_gives_up_after_max_retries(
monkeypatch: pytest.MonkeyPatch,
) -> None:
streams = [
_FakeStream(exc=_midstream_api_error())
for _ in range(execution._MAX_TRANSIENT_MODEL_RETRIES + 1)
]
with pytest.raises(APIError):
await _run_once(monkeypatch, streams)
@pytest.mark.asyncio
async def test_run_cycle_does_not_retry_permanent_error(
monkeypatch: pytest.MonkeyPatch,
) -> None:
bad_request = BadRequestError(
"bad", response=httpx.Response(400, request=_request()), body=None
)
streams = [_FakeStream(exc=bad_request), _FakeStream()]
with pytest.raises(BadRequestError):
await _run_once(monkeypatch, streams)
+388 -3
View File
@@ -2,11 +2,18 @@
from __future__ import annotations from __future__ import annotations
from typing import Any
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
from strix.core.hooks import BudgetExceededError, ReportUsageHooks from strix.core.hooks import (
BudgetExceededError,
BudgetPausedError,
ReportUsageHooks,
SubagentBudgetReservedError,
recomputed_budget_flags,
)
def _make_hooks(max_budget: float | None) -> ReportUsageHooks: def _make_hooks(max_budget: float | None) -> ReportUsageHooks:
@@ -20,9 +27,22 @@ def _make_report_state(cost: float) -> MagicMock:
return state return state
def _make_context(agent_id: str = "test-agent") -> MagicMock: def _make_context(agent_id: str = "test-agent", parent_id: str | None = None) -> MagicMock:
ctx: MagicMock = MagicMock() ctx: MagicMock = MagicMock()
ctx.context = {"agent_id": agent_id} ctx.context = {"agent_id": agent_id, "parent_id": parent_id}
return ctx
def _make_warn_context(
*,
requests: int,
parent_id: str | None = None,
agent_id: str = "test-agent",
) -> MagicMock:
ctx: MagicMock = MagicMock()
ctx.context = {"agent_id": agent_id, "parent_id": parent_id}
ctx.usage = MagicMock()
ctx.usage.requests = requests
return ctx return ctx
@@ -89,6 +109,127 @@ async def test_error_message_includes_amounts() -> None:
assert "7.1234" in str(exc_info.value) assert "7.1234" in str(exc_info.value)
@pytest.mark.asyncio
async def test_subagent_stops_at_reserve() -> None:
hooks = _make_hooks(10.0)
state = _make_report_state(9.0)
with (
patch("strix.core.hooks.get_global_report_state", return_value=state),
pytest.raises(SubagentBudgetReservedError),
):
await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_subagent_below_reserve_does_not_raise() -> None:
hooks = _make_hooks(10.0)
state = _make_report_state(8.99)
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_subagent_overshoot_to_full_budget_triggers_scan_wide_stop() -> None:
hooks = _make_hooks(10.0)
state = _make_report_state(10.5)
with (
patch("strix.core.hooks.get_global_report_state", return_value=state),
pytest.raises(BudgetExceededError),
):
await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_root_keeps_running_inside_reserve() -> None:
hooks = _make_hooks(10.0)
state = _make_report_state(9.5)
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_root_hard_stop_stays_at_full_budget() -> None:
hooks = _make_hooks(10.0)
state = _make_report_state(10.0)
with (
patch("strix.core.hooks.get_global_report_state", return_value=state),
pytest.raises(BudgetExceededError),
):
await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_budget_warning_mentions_reserve() -> None:
hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0)
state = _make_report_state(7.5)
root_items: list[Any] = []
sub_items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id=None), MagicMock(), None, root_items
)
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, sub_items
)
assert "stopped at 90%" in root_items[0]["content"]
assert "stopped at 90%" in sub_items[0]["content"]
assert "root agent's final report" in sub_items[0]["content"]
@pytest.mark.asyncio
async def test_subagent_critical_budget_warning_reachable_before_reserve() -> None:
hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0)
state = _make_report_state(8.6)
sub_items: list[Any] = []
root_items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, sub_items
)
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id=None), MagicMock(), None, root_items
)
assert "[CRITICAL]" in sub_items[0]["content"]
assert "[URGENT]" in root_items[0]["content"]
@pytest.mark.parametrize(
("parent_id", "cost", "expected"),
[
("root-1", 0.0, None),
("root-1", 8.9999, None),
("root-1", 9.0, SubagentBudgetReservedError),
("root-1", 9.0001, SubagentBudgetReservedError),
("root-1", 9.5, SubagentBudgetReservedError),
("root-1", 9.9999, SubagentBudgetReservedError),
("root-1", 10.0, BudgetExceededError),
("root-1", 10.0001, BudgetExceededError),
("root-1", 25.0, BudgetExceededError),
(None, 0.0, None),
(None, 8.9999, None),
(None, 9.0, None),
(None, 9.5, None),
(None, 9.9999, None),
(None, 10.0, BudgetExceededError),
(None, 10.0001, BudgetExceededError),
(None, 25.0, BudgetExceededError),
],
)
@pytest.mark.asyncio
async def test_budget_enforcement_decision_table(
parent_id: str | None, cost: float, expected: type[Exception] | None
) -> None:
hooks = _make_hooks(10.0)
state = _make_report_state(cost)
with patch("strix.core.hooks.get_global_report_state", return_value=state):
if expected is None:
await hooks.on_llm_end(_make_context(parent_id=parent_id), MagicMock(), MagicMock())
else:
with pytest.raises(expected):
await hooks.on_llm_end(_make_context(parent_id=parent_id), MagicMock(), MagicMock())
state.record_sdk_usage.assert_called_once()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_no_raise_when_report_state_none() -> None: async def test_no_raise_when_report_state_none() -> None:
hooks = _make_hooks(1.0) hooks = _make_hooks(1.0)
@@ -106,3 +247,247 @@ def test_non_positive_budget_rejected(bad_budget: float) -> None:
def test_budget_exceeded_error_is_runtime_error() -> None: def test_budget_exceeded_error_is_runtime_error() -> None:
err = BudgetExceededError("test") err = BudgetExceededError("test")
assert isinstance(err, RuntimeError) assert isinstance(err, RuntimeError)
def test_non_positive_max_turns_rejected() -> None:
with pytest.raises(ValueError, match="positive integer"):
ReportUsageHooks(model="test-model", max_turns=0)
@pytest.mark.asyncio
async def test_no_turn_warning_below_first_band() -> None:
hooks = ReportUsageHooks(model="test-model", max_turns=100)
items: list[Any] = []
await hooks.on_llm_start(_make_warn_context(requests=68), MagicMock(), None, items)
assert items == []
@pytest.mark.asyncio
async def test_turn_warning_notice_band() -> None:
hooks = ReportUsageHooks(model="test-model", max_turns=100)
items: list[Any] = []
await hooks.on_llm_start(_make_warn_context(requests=69), MagicMock(), None, items)
assert len(items) == 1
content = items[0]["content"]
assert "[NOTICE]" in content
assert "finish_scan" in content
@pytest.mark.asyncio
async def test_turn_warning_escalates_and_names_subagent_tool() -> None:
hooks = ReportUsageHooks(model="test-model", max_turns=100)
items: list[Any] = []
await hooks.on_llm_start(
_make_warn_context(requests=95, parent_id="root-1"), MagicMock(), None, items
)
assert len(items) == 1
content = items[0]["content"]
assert "[CRITICAL]" in content
assert "agent_finish" in content
@pytest.mark.asyncio
async def test_turn_warning_root_directive_distinct_from_subagent() -> None:
hooks = ReportUsageHooks(model="test-model", max_turns=100)
root_items: list[Any] = []
await hooks.on_llm_start(
_make_warn_context(requests=85, parent_id=None), MagicMock(), None, root_items
)
root = root_items[0]["content"]
sub_items: list[Any] = []
await hooks.on_llm_start(
_make_warn_context(requests=85, parent_id="root-1"), MagicMock(), None, sub_items
)
sub = sub_items[0]["content"]
assert root != sub
assert "root agent" in root
assert "finish_scan" in root
assert "agent_finish" not in root
assert "whole scan" in root
assert "sub-agent" in sub
assert "agent_finish" in sub
assert "finish_scan" not in sub
assert "confirmed" in sub
@pytest.mark.asyncio
async def test_budget_warning_root_directive_distinct_from_subagent() -> None:
hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0)
state = _make_report_state(8.6)
root_items: list[Any] = []
sub_items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id=None), MagicMock(), None, root_items
)
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, sub_items
)
root = root_items[0]["content"]
sub = sub_items[0]["content"]
assert "finish_scan" in root and "agent_finish" not in root
assert "agent_finish" in sub and "finish_scan" not in sub
assert "confirmed" in sub
@pytest.mark.parametrize("parent_id", [None, "root-1"])
@pytest.mark.asyncio
async def test_turn_warning_directive_escalates_per_stage(parent_id: str | None) -> None:
hooks = ReportUsageHooks(model="test-model", max_turns=100)
contents: dict[str, str] = {}
for label, requests in (("notice", 69), ("urgent", 85), ("critical", 95)):
items: list[Any] = []
await hooks.on_llm_start(
_make_warn_context(requests=requests, parent_id=parent_id), MagicMock(), None, items
)
contents[label] = items[0]["content"]
assert len({contents["notice"], contents["urgent"], contents["critical"]}) == 3
assert "[NOTICE]" in contents["notice"] and "begin planning" in contents["notice"]
assert "[URGENT]" in contents["urgent"] and "prioritize" in contents["urgent"]
assert "[CRITICAL]" in contents["critical"] and "STOP" in contents["critical"]
@pytest.mark.asyncio
async def test_no_turn_warning_when_max_turns_unset() -> None:
hooks = ReportUsageHooks(model="test-model")
items: list[Any] = []
await hooks.on_llm_start(_make_warn_context(requests=999), MagicMock(), None, items)
assert items == []
@pytest.mark.asyncio
async def test_no_budget_warning_below_first_band() -> None:
hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0)
state = _make_report_state(6.9)
items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(_make_warn_context(requests=0), MagicMock(), None, items)
assert items == []
@pytest.mark.asyncio
async def test_budget_warning_broadcast_content() -> None:
hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0)
state = _make_report_state(9.6)
items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(_make_warn_context(requests=0), MagicMock(), None, items)
assert len(items) == 1
content = items[0]["content"]
assert "[CRITICAL]" in content
assert "shared across every agent" in content
@pytest.mark.asyncio
async def test_turn_and_budget_warnings_stack() -> None:
hooks = ReportUsageHooks(model="test-model", max_budget_usd=10.0, max_turns=100)
state = _make_report_state(8.6)
items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(_make_warn_context(requests=89), MagicMock(), None, items)
assert len(items) == 2
joined = " ".join(i["content"] for i in items)
assert "Turn budget" in joined
assert "cost budget" in joined
def _make_interactive_hooks(max_budget: float | None) -> ReportUsageHooks:
return ReportUsageHooks(model="test-model", max_budget_usd=max_budget, interactive=True)
@pytest.mark.asyncio
async def test_interactive_at_budget_pauses_instead_of_stopping() -> None:
hooks = _make_interactive_hooks(10.0)
state = _make_report_state(10.0)
with (
patch("strix.core.hooks.get_global_report_state", return_value=state),
pytest.raises(BudgetPausedError),
):
await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_interactive_subagent_has_no_reserve() -> None:
hooks = _make_interactive_hooks(10.0)
state = _make_report_state(9.5)
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_interactive_subagent_pauses_at_full_budget() -> None:
hooks = _make_interactive_hooks(10.0)
state = _make_report_state(10.5)
with (
patch("strix.core.hooks.get_global_report_state", return_value=state),
pytest.raises(BudgetPausedError),
):
await hooks.on_llm_end(_make_context(parent_id="root-1"), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_extend_budget_lifts_the_pause() -> None:
hooks = _make_interactive_hooks(10.0)
state = _make_report_state(10.5)
hooks.extend_budget()
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_extend_budget_adds_original_amount_each_time() -> None:
hooks = _make_interactive_hooks(10.0)
hooks.extend_budget()
hooks.extend_budget()
state = _make_report_state(29.9)
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock())
state = _make_report_state(30.0)
with (
patch("strix.core.hooks.get_global_report_state", return_value=state),
pytest.raises(BudgetPausedError),
):
await hooks.on_llm_end(_make_context(parent_id=None), MagicMock(), MagicMock())
@pytest.mark.asyncio
async def test_interactive_subagent_uses_root_warning_bands() -> None:
hooks = _make_interactive_hooks(10.0)
state = _make_report_state(7.4)
items: list[Any] = []
with patch("strix.core.hooks.get_global_report_state", return_value=state):
await hooks.on_llm_start(
_make_warn_context(requests=0, parent_id="root-1"), MagicMock(), None, items
)
assert len(items) == 1
content = items[0]["content"]
assert "[NOTICE]" in content
assert "paused until the user chooses to continue" in content
assert "reserve" not in content.lower()
@pytest.mark.parametrize(
("cost", "max_budget", "interactive", "expected"),
[
(0.0, None, False, (False, False)),
(100.0, None, False, (False, False)),
(5.0, 10.0, False, (False, False)),
(9.0, 10.0, False, (False, True)),
(10.0, 10.0, False, (True, True)),
(10.0, 20.0, False, (False, False)),
(10.0, 10.0, True, (False, False)),
],
)
def test_recomputed_budget_flags(
cost: float,
max_budget: float | None,
interactive: bool,
expected: tuple[bool, bool],
) -> None:
assert recomputed_budget_flags(cost, max_budget, interactive=interactive) == expected
+122
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
from itertools import pairwise from itertools import pairwise
from typing import Any from typing import Any
import litellm
import pytest import pytest
from strix.core.inputs import build_root_task, child_initial_input, make_model_settings from strix.core.inputs import build_root_task, child_initial_input, make_model_settings
@@ -55,6 +56,103 @@ def test_child_initial_input_no_consecutive_same_role(parent_history: list[Any])
assert all(prev != nxt for prev, nxt in pairwise(roles)) assert all(prev != nxt for prev, nxt in pairwise(roles))
def _cache_points(model_name: str) -> Any:
extra = make_model_settings(None, model_name=model_name).extra_args or {}
return extra.get("cache_control_injection_points")
def test_make_model_settings_enables_prompt_cache_for_bedrock_claude() -> None:
assert _cache_points("bedrock/global.anthropic.claude-opus-4-8") == [
{"location": "message", "role": "system"},
{"location": "tool_config"},
{"location": "message", "index": -1},
]
@pytest.mark.parametrize(
"model_name",
[
"anthropic/claude-sonnet-4-5",
"openrouter/anthropic/claude-3.5-sonnet",
"vertex_ai/claude-sonnet-4-5",
],
)
def test_make_model_settings_enables_prompt_cache_for_non_bedrock_claude(model_name: str) -> None:
assert _cache_points(model_name) == [
{"location": "message", "role": "system"},
{"location": "message", "index": -1},
]
def test_tool_config_point_not_leaked_to_non_bedrock_claude() -> None:
# LiteLLM only consumes tool_config on Bedrock; elsewhere it leaks onto the
# wire and native Anthropic 400s.
for model in ("anthropic/claude-sonnet-4-5", "openrouter/anthropic/claude-3.5-sonnet"):
points = _cache_points(model) or []
assert all(p.get("location") != "tool_config" for p in points)
def test_prompt_cache_can_be_disabled() -> None:
assert (
make_model_settings(
None, model_name="anthropic/claude-sonnet-4-5", prompt_cache=False
).extra_args
is None
)
@pytest.mark.parametrize("model_name", ["gpt-5", "vertex_ai/gemini-2.5-pro", "openai/o3"])
def test_make_model_settings_no_prompt_cache_for_non_claude(model_name: str) -> None:
assert make_model_settings(None, model_name=model_name).extra_args is None
def test_no_prompt_cache_for_unmapped_bedrock_claude_model(monkeypatch: Any) -> None:
# A Bedrock Claude model LiteLLM hasn't mapped must run uncached, not crash.
unmapped = "bedrock/global.anthropic.claude-brand-new-9"
monkeypatch.setattr(litellm, "model_cost", {}, raising=False)
if getattr(getattr(litellm, "utils", None), "supports_prompt_caching", None):
monkeypatch.setattr(litellm.utils, "supports_prompt_caching", lambda *_a, **_k: False)
assert make_model_settings(None, model_name=unmapped).extra_args is None
def test_prompt_cache_kept_for_non_bedrock_claude_even_if_unmapped(monkeypatch: Any) -> None:
# Only Bedrock hard-rejects unknown cache fields, so only Bedrock is guarded.
monkeypatch.setattr(litellm, "model_cost", {}, raising=False)
if getattr(getattr(litellm, "utils", None), "supports_prompt_caching", None):
monkeypatch.setattr(litellm.utils, "supports_prompt_caching", lambda *_a, **_k: False)
for model in ("anthropic/claude-brand-new-9", "openrouter/anthropic/claude-brand-new"):
assert _cache_points(model) == [
{"location": "message", "role": "system"},
{"location": "message", "index": -1},
]
def test_conversation_tail_breakpoint_moves_with_appended_transcript() -> None:
# LiteLLM must place the index=-1 cache_control on the last message however
# long the transcript grows.
hook_mod = pytest.importorskip("litellm.integrations.anthropic_cache_control_hook")
apply = hook_mod.AnthropicCacheControlHook._apply_message_injections
points = _cache_points("bedrock/global.anthropic.claude-opus-4-8")
msg_points = [p for p in points if p.get("location") == "message"]
def last_msg_cache_control(n_turns: int) -> Any:
messages: list[dict[str, Any]] = [{"role": "system", "content": "stable prompt"}]
for i in range(n_turns):
messages.append({"role": "assistant", "content": f"turn {i} action"})
messages.append({"role": "user", "content": f"turn {i} tool result"})
processed = apply(msg_points, messages, 4)
last = processed[-1]
content = last.get("content")
if isinstance(content, list):
return content[-1].get("cache_control")
return last.get("cache_control")
assert last_msg_cache_control(2) == {"type": "ephemeral"}
assert last_msg_cache_control(20) == {"type": "ephemeral"}
def test_build_root_task_empty_config() -> None: def test_build_root_task_empty_config() -> None:
assert build_root_task({}) == "" assert build_root_task({}) == ""
@@ -174,6 +272,30 @@ def test_make_model_settings_omits_timeout_when_unset() -> None:
assert settings.extra_args is None assert settings.extra_args is None
def test_make_model_settings_sets_extra_headers() -> None:
settings = make_model_settings(
"none",
model_name="openai/some-model",
extra_headers={"X-Feature-Key": "svc", "X-Tenant": "acme"},
)
assert settings.extra_headers == {"X-Feature-Key": "svc", "X-Tenant": "acme"}
def test_make_model_settings_omits_extra_headers_when_unset() -> None:
assert make_model_settings("none", model_name="gpt-4o").extra_headers is None
def test_make_model_settings_extra_headers_survive_reasoning_resolve() -> None:
settings = make_model_settings(
"high",
model_name="openai/o3",
extra_headers={"X-Feature-Key": "svc"},
)
assert settings.extra_headers == {"X-Feature-Key": "svc"}
def test_make_model_settings_timeout_survives_reasoning_resolve() -> None: def test_make_model_settings_timeout_survives_reasoning_resolve() -> None:
# Reasoning is resolved via ModelSettings.resolve(); the timeout in extra_args # Reasoning is resolved via ModelSettings.resolve(); the timeout in extra_args
# must not be dropped when a reasoning override is merged in. # must not be dropped when a reasoning override is merged in.
+154
View File
@@ -0,0 +1,154 @@
from __future__ import annotations
import stat
import subprocess
import sys
import tarfile
from pathlib import Path
import pytest
RELEASE_VERSION = "9.9.9"
RELEASE_TARGET = "linux-arm64"
pytestmark = pytest.mark.skipif(
sys.platform == "win32",
reason="scripts/install.sh is a POSIX shell installer",
)
def _write_executable(path: Path, content: str) -> None:
path.write_text(content, encoding="utf-8")
path.chmod(path.stat().st_mode | stat.S_IXUSR)
def _create_release_archive(tmp_path: Path) -> Path:
binary_name = f"strix-{RELEASE_VERSION}-{RELEASE_TARGET}"
binary_path = tmp_path / binary_name
_write_executable(binary_path, f"#!/bin/sh\nprintf 'strix {RELEASE_VERSION}\\n'\n")
archive_path = tmp_path / f"{binary_name}.tar.gz"
with tarfile.open(archive_path, "w:gz") as archive:
archive.add(binary_path, arcname=binary_name)
return archive_path
def _create_mock_commands(tmp_path: Path, machine: str) -> Path:
mock_bin = tmp_path / "mock-bin"
mock_bin.mkdir()
_write_executable(
mock_bin / "uname",
f"""#!/bin/sh
case "$1" in
-s) echo Linux ;;
-m) echo {machine} ;;
*) echo "unexpected uname argument: $*" >&2; exit 1 ;;
esac
""",
)
_write_executable(mock_bin / "docker", "#!/bin/sh\nexit 0\n")
_write_executable(
mock_bin / "curl",
"""#!/bin/sh
output=""
while [ "$#" -gt 0 ]; do
if [ "$1" = "-o" ]; then
output="$2"
shift 2
continue
fi
printf '%s\\n' "$1" >> "$STRIX_TEST_CURL_LOG"
shift
done
cp "$STRIX_TEST_ARCHIVE" "$output"
""",
)
return mock_bin
def _create_installer_environment(
tmp_path: Path,
archive_path: Path,
mock_bin: Path,
) -> tuple[dict[str, str], Path, Path]:
"""Build the installer environment explicitly.
Every variable the installer reads is listed here, so no inherited value
(`XDG_CONFIG_HOME`, `GITHUB_ACTIONS`, `TMPDIR`, ...) can send a write
outside the sandbox or change the code path under test.
"""
home_path = tmp_path / "home"
home_path.mkdir()
download_path = tmp_path / "downloads"
download_path.mkdir()
curl_log_path = tmp_path / "curl.log"
environment = {
"HOME": str(home_path),
"XDG_CONFIG_HOME": str(home_path / ".config"),
"PATH": f"{mock_bin}:/usr/bin:/bin",
"SHELL": "/bin/bash",
"TMPDIR": str(download_path),
"STRIX_TEST_ARCHIVE": str(archive_path),
"STRIX_TEST_CURL_LOG": str(curl_log_path),
"VERSION": RELEASE_VERSION,
}
return environment, home_path, curl_log_path
def _run_installer(
repository_root: Path,
environment: dict[str, str],
) -> subprocess.CompletedProcess[str]:
return subprocess.run( # noqa: S603
["/bin/bash", str(repository_root / "scripts/install.sh")],
cwd=repository_root,
env=environment,
capture_output=True,
text=True,
check=False,
)
def test_installer_downloads_and_runs_linux_arm64_release(tmp_path: Path) -> None:
repository_root = Path(__file__).resolve().parents[1]
archive_path = _create_release_archive(tmp_path)
mock_bin = _create_mock_commands(tmp_path, machine="aarch64")
environment, home_path, curl_log_path = _create_installer_environment(
tmp_path,
archive_path,
mock_bin,
)
result = _run_installer(repository_root, environment)
assert result.returncode == 0, result.stderr
expected_filename = f"strix-{RELEASE_VERSION}-{RELEASE_TARGET}.tar.gz"
assert expected_filename in curl_log_path.read_text(encoding="utf-8")
installed_binary = home_path / ".strix/bin/strix"
installed_result = subprocess.run( # noqa: S603
[str(installed_binary), "--version"],
capture_output=True,
text=True,
check=True,
)
assert installed_result.stdout.strip() == f"strix {RELEASE_VERSION}"
def test_installer_rejects_unsupported_architecture(tmp_path: Path) -> None:
repository_root = Path(__file__).resolve().parents[1]
archive_path = _create_release_archive(tmp_path)
mock_bin = _create_mock_commands(tmp_path, machine="riscv64")
environment, home_path, curl_log_path = _create_installer_environment(
tmp_path,
archive_path,
mock_bin,
)
result = _run_installer(repository_root, environment)
assert result.returncode != 0
assert "Unsupported OS/Arch: linux/riscv64" in result.stdout
assert not curl_log_path.exists()
assert not (home_path / ".strix").exists()
+97
View File
@@ -0,0 +1,97 @@
"""Tests for LLM_EXTRA_HEADERS: custom default headers on OpenAI-compatible endpoints."""
from __future__ import annotations
import json
from typing import TYPE_CHECKING
import litellm
import pytest
from agents.models import _openai_shared
from strix.config import loader
from strix.config.loader import load_settings
from strix.config.models import configure_sdk_model_defaults
if TYPE_CHECKING:
from collections.abc import Iterator
_ENV_KEYS = ["STRIX_LLM", "LLM_API_KEY", "LLM_API_BASE", "LLM_EXTRA_HEADERS"]
@pytest.fixture(autouse=True)
def _reset(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
for key in _ENV_KEYS:
monkeypatch.delenv(key, raising=False)
monkeypatch.setattr(loader, "_cached", None)
monkeypatch.setattr(loader, "_override", None)
saved_headers = litellm.headers
saved_client = _openai_shared.get_default_openai_client()
litellm.headers = None
try:
yield
finally:
litellm.headers = saved_headers
_openai_shared.set_default_openai_client(saved_client) # type: ignore[arg-type]
def test_extra_headers_parsed_from_json_env(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-A": "1", "X-B": "2"}))
settings = load_settings()
assert settings.llm.extra_headers == {"X-A": "1", "X-B": "2"}
def test_extra_headers_merged_into_litellm_headers(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("STRIX_LLM", "litellm/openai/some-model")
monkeypatch.setenv("LLM_API_BASE", "https://gateway.example/v1")
monkeypatch.setenv("LLM_API_KEY", "token")
headers = {"X-Feature-Key": "svc", "X-Tenant": "acme"}
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps(headers))
configure_sdk_model_defaults(load_settings())
current: object = litellm.headers
assert isinstance(current, dict)
assert current["X-Feature-Key"] == "svc"
assert current["X-Tenant"] == "acme"
def test_extra_headers_applied_to_native_openai_client(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("STRIX_LLM", "openai/some-model")
monkeypatch.setenv("LLM_API_BASE", "https://gateway.example/v1")
monkeypatch.setenv("LLM_API_KEY", "token")
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Feature-Key": "svc"}))
configure_sdk_model_defaults(load_settings())
client = _openai_shared.get_default_openai_client()
assert client is not None
assert client.default_headers.get("X-Feature-Key") == "svc"
assert str(client.base_url).rstrip("/") == "https://gateway.example/v1"
def test_extra_headers_applied_to_native_openai_without_custom_base(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("STRIX_LLM", "openai/gpt-5")
monkeypatch.setenv("LLM_API_KEY", "token")
monkeypatch.setenv("LLM_EXTRA_HEADERS", json.dumps({"X-Feature-Key": "svc"}))
configure_sdk_model_defaults(load_settings())
client = _openai_shared.get_default_openai_client()
assert client is not None
assert client.default_headers.get("X-Feature-Key") == "svc"
def test_no_extra_headers_leaves_litellm_headers_untouched(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("STRIX_LLM", "openai/some-model")
monkeypatch.setenv("LLM_API_BASE", "https://gateway.example/v1")
monkeypatch.setenv("LLM_API_KEY", "token")
configure_sdk_model_defaults(load_settings())
assert litellm.headers is None
+5 -2
View File
@@ -14,6 +14,7 @@ import strix.tools.notes.tools as notes_tools
import strix.tools.todo.tools as todo_tools import strix.tools.todo.tools as todo_tools
from strix.core import runner from strix.core import runner
from strix.core.agents import AgentCoordinator from strix.core.agents import AgentCoordinator
from strix.runtime import session_manager
def _make_rate_limit_error() -> RateLimitError: def _make_rate_limit_error() -> RateLimitError:
@@ -38,6 +39,8 @@ async def test_persistent_rate_limit_stops_gracefully(
reasoning_effort="high", reasoning_effort="high",
force_required_tool_choice=False, force_required_tool_choice=False,
timeout=300, timeout=300,
prompt_cache=True,
extra_headers=None,
), ),
runtime=types.SimpleNamespace(max_context_images=3), runtime=types.SimpleNamespace(max_context_images=3),
) )
@@ -56,8 +59,8 @@ async def test_persistent_rate_limit_stops_gracefully(
async def _cleanup(*_args: Any, **_kwargs: Any) -> None: async def _cleanup(*_args: Any, **_kwargs: Any) -> None:
return None return None
monkeypatch.setattr(runner.session_manager, "create_or_reuse", _create_or_reuse) # type: ignore[attr-defined] monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse)
monkeypatch.setattr(runner.session_manager, "cleanup", _cleanup) # type: ignore[attr-defined] monkeypatch.setattr(session_manager, "cleanup", _cleanup)
monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task") monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task")
monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: "") monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: "")
+5 -2
View File
@@ -17,6 +17,7 @@ import strix.tools.notes.tools as notes_tools
import strix.tools.todo.tools as todo_tools import strix.tools.todo.tools as todo_tools
from strix.core import runner from strix.core import runner
from strix.core.agents import AgentCoordinator from strix.core.agents import AgentCoordinator
from strix.runtime import session_manager
def _make_rate_limit_error() -> RateLimitError: def _make_rate_limit_error() -> RateLimitError:
@@ -46,6 +47,8 @@ def _patch_engine_scaffold(
reasoning_effort="high", reasoning_effort="high",
force_required_tool_choice=False, force_required_tool_choice=False,
timeout=300, timeout=300,
prompt_cache=True,
extra_headers=None,
), ),
runtime=types.SimpleNamespace(max_context_images=3), runtime=types.SimpleNamespace(max_context_images=3),
) )
@@ -66,8 +69,8 @@ def _patch_engine_scaffold(
async def _cleanup(*_args: Any, **_kwargs: Any) -> None: async def _cleanup(*_args: Any, **_kwargs: Any) -> None:
return None return None
monkeypatch.setattr(runner.session_manager, "create_or_reuse", _create_or_reuse) monkeypatch.setattr(session_manager, "create_or_reuse", _create_or_reuse)
monkeypatch.setattr(runner.session_manager, "cleanup", _cleanup) monkeypatch.setattr(session_manager, "cleanup", _cleanup)
monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task") monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task")
monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: scope_context) monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: scope_context)
+35 -9
View File
@@ -159,14 +159,40 @@ def test_sha256_file(tmp_path: Path) -> None:
assert update_check._sha256_file(path) == hashlib.sha256(b"strix").hexdigest() assert update_check._sha256_file(path) == hashlib.sha256(b"strix").hexdigest()
def test_release_target(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.parametrize(
("system", "machine", "expected"),
[
("Linux", "x86_64", "linux-x86_64"),
("Linux", "aarch64", "linux-arm64"),
("Linux", "arm64", "linux-arm64"),
("Darwin", "arm64", "macos-arm64"),
("Darwin", "riscv64", None),
],
)
def test_release_target(
monkeypatch: pytest.MonkeyPatch,
system: str,
machine: str,
expected: str | None,
) -> None:
monkeypatch.setattr(platform, "system", lambda: system)
monkeypatch.setattr(platform, "machine", lambda: machine)
assert update_check._release_target() == expected
def test_self_update_uses_linux_arm64_release(monkeypatch: pytest.MonkeyPatch) -> None:
requested_update: list[tuple[str, str]] = []
def record_download(version: str, target: str, _console: Console) -> bool:
requested_update.append((version, target))
return True
monkeypatch.setattr(update_check, "is_binary_install", lambda: True)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
monkeypatch.setattr(platform, "system", lambda: "Linux") monkeypatch.setattr(platform, "system", lambda: "Linux")
monkeypatch.setattr(platform, "machine", lambda: "x86_64") monkeypatch.setattr(platform, "machine", lambda: "aarch64")
assert update_check._release_target() == "linux-x86_64" monkeypatch.setattr(update_check, "_download_and_replace", record_download)
monkeypatch.setattr(platform, "system", lambda: "Darwin") assert update_check.self_update(Console(file=io.StringIO()), version="1.1.0") is True
monkeypatch.setattr(platform, "machine", lambda: "arm64") assert requested_update == [("1.1.0", "linux-arm64")]
assert update_check._release_target() == "macos-arm64"
monkeypatch.setattr(platform, "machine", lambda: "riscv64")
assert update_check._release_target() is None
+125 -2
View File
@@ -4,9 +4,11 @@ from __future__ import annotations
import json import json
import os import os
import sqlite3
import urllib.error import urllib.error
import urllib.request import urllib.request
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from urllib.parse import urlsplit
from strix.core.paths import latest_run_dir, runs_base_dir from strix.core.paths import latest_run_dir, runs_base_dir
from strix.interface.viewer.server import serve from strix.interface.viewer.server import serve
@@ -86,6 +88,75 @@ def test_build_run_state_from_agents_json(tmp_path: Path) -> None:
assert state["events"] == [] assert state["events"] == []
def test_build_run_state_keeps_same_call_id_separate_per_agent(tmp_path: Path) -> None:
run_dir = _make_run(tmp_path, "tools", status="completed", end_time=None)
agents_db = run_dir / ".state" / "agents.db"
rows = [
(
"root",
{
"type": "function_call",
"call_id": "exec_command_0",
"name": "exec_command",
"arguments": json.dumps({"cmd": "echo root"}),
},
),
(
"root",
{
"type": "function_call_output",
"call_id": "exec_command_0",
"output": json.dumps({"success": True, "output": "root"}),
},
),
(
"child",
{
"type": "function_call",
"call_id": "exec_command_0",
"name": "exec_command",
"arguments": json.dumps({"cmd": "echo child"}),
},
),
(
"child",
{
"type": "function_call_output",
"call_id": "exec_command_0",
"output": json.dumps({"success": True, "output": "child"}),
},
),
]
with sqlite3.connect(agents_db) as conn:
conn.execute(
"""
create table agent_messages (
id integer primary key,
session_id text not null,
message_data text not null,
created_at text not null
)
"""
)
conn.executemany(
"""
insert into agent_messages (session_id, message_data, created_at)
values (?, ?, '2026-01-01T00:00:00+00:00')
""",
[(agent_id, json.dumps(message)) for agent_id, message in rows],
)
state = build_run_state(run_dir)
tools = [event for event in state["events"] if event["type"] == "tool"]
assert len(tools) == 2
by_agent = {event["agent_id"]: event for event in tools}
assert by_agent["root"]["data"]["args"] == {"cmd": "echo root"}
assert by_agent["root"]["data"]["result"]["output"] == "root"
assert by_agent["child"]["data"]["args"] == {"cmd": "echo child"}
assert by_agent["child"]["data"]["result"]["output"] == "child"
def _get(url: str, *, cookie: str | None = None) -> tuple[int, str, bytes]: def _get(url: str, *, cookie: str | None = None) -> tuple[int, str, bytes]:
headers = {"Cookie": cookie} if cookie else {} headers = {"Cookie": cookie} if cookie else {}
req = urllib.request.Request(url, headers=headers) # noqa: S310 - localhost test server req = urllib.request.Request(url, headers=headers) # noqa: S310 - localhost test server
@@ -271,6 +342,11 @@ def _session_cookie(url: str, token: str) -> str:
return raw.split(";", 1)[0] return raw.split(";", 1)[0]
def _cookie_name(url: str) -> str:
"""The per-server session cookie name, derived from the bound port."""
return f"strix_viewer_session_{urlsplit(url).port}"
def _get_status(url: str, *, cookie: str | None = None) -> int: def _get_status(url: str, *, cookie: str | None = None) -> int:
headers = {"Cookie": cookie} if cookie else {} headers = {"Cookie": cookie} if cookie else {}
req = urllib.request.Request(url, headers=headers) # noqa: S310 - localhost test server req = urllib.request.Request(url, headers=headers) # noqa: S310 - localhost test server
@@ -311,7 +387,7 @@ def test_capability_issued_only_for_tokened_bootstrap(
# Only the correct bootstrap token mints the session cookie. # Only the correct bootstrap token mints the session cookie.
with urllib.request.urlopen(f"{url}/?token={token}") as resp: # noqa: S310 # nosec B310 with urllib.request.urlopen(f"{url}/?token={token}") as resp: # noqa: S310 # nosec B310
cookie = str(resp.headers.get("Set-Cookie", "")) cookie = str(resp.headers.get("Set-Cookie", ""))
assert "strix_viewer_session=" in cookie assert f"{_cookie_name(url)}=" in cookie
assert "HttpOnly" in cookie and "SameSite=Strict" in cookie assert "HttpOnly" in cookie and "SameSite=Strict" in cookie
# Static assets never carry it. # Static assets never carry it.
@@ -344,7 +420,7 @@ def test_unauthorized_client_cannot_acquire_capability(
url, url,
"/api/agents/steer", "/api/agents/steer",
{"agent_id": "root", "message": "pwn"}, {"agent_id": "root", "message": "pwn"},
cookie="strix_viewer_session=", cookie=f"{_cookie_name(url)}=",
) )
assert status == 403 assert status == 403
assert delivered == [] assert delivered == []
@@ -541,6 +617,53 @@ def test_runs_list_requires_session_and_verification(
httpd.server_close() httpd.server_close()
def test_concurrent_servers_use_distinct_cookies(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Cookies are host-scoped, not port-scoped: two viewers on 127.0.0.1 must
not share a cookie slot, and one server's cookie must not pass the other's
session gate."""
run_a = _make_run(tmp_path / "a", "run-a", status="running", end_time=None)
run_b = _make_run(tmp_path / "b", "run-b", status="running", end_time=None)
_bundle(tmp_path, monkeypatch)
monkeypatch.setattr(
"strix.interface.viewer.auth.read_auth", lambda: {"email": "a@b.com", "token": "t"}
)
monkeypatch.setattr("strix.interface.viewer.auth.is_verified", lambda: True)
httpd_a, url_a, token_a = serve(run_a, open_browser=False)
httpd_b, url_b, token_b = serve(run_b, open_browser=False)
try:
cookie_a = _session_cookie(url_a, token_a)
cookie_b = _session_cookie(url_b, token_b)
# The two servers mint differently named cookies, so a browser stores both.
assert cookie_a.split("=", 1)[0] == _cookie_name(url_a)
assert cookie_b.split("=", 1)[0] == _cookie_name(url_b)
assert cookie_a.split("=", 1)[0] != cookie_b.split("=", 1)[0]
def _status(url: str, cookie: str) -> dict[str, object]:
_, _, body = _get(f"{url}/api/auth/status", cookie=cookie)
return dict(json.loads(body))
# Each server honors its own cookie...
assert _status(url_a, cookie_a)["verified"] is True
assert _status(url_b, cookie_b)["verified"] is True
# ...but treats the other server's cookie as session-less.
assert _status(url_a, cookie_b)["verified"] is False
assert _status(url_b, cookie_a)["verified"] is False
# Even both cookies together (what a real browser would send) only
# match the token minted by the receiving server.
both = f"{cookie_a}; {cookie_b}"
assert _status(url_a, both)["verified"] is True
assert _status(url_b, both)["verified"] is True
finally:
httpd_a.shutdown()
httpd_a.server_close()
httpd_b.shutdown()
httpd_b.server_close()
def test_server_rejects_path_traversal(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: def test_server_rejects_path_traversal(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
run_dir = _make_run(tmp_path, "guard", status="completed", end_time="2026-01-01T00:00:00Z") run_dir = _make_run(tmp_path, "guard", status="completed", end_time="2026-01-01T00:00:00Z")
secret = tmp_path / "secret.txt" secret = tmp_path / "secret.txt"
Generated
+1 -1
View File
@@ -2411,7 +2411,7 @@ wheels = [
[[package]] [[package]]
name = "strix-agent" name = "strix-agent"
version = "1.3.1" version = "1.4.1"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "caido-sdk-client" }, { name = "caido-sdk-client" },