Compare commits

...
Author SHA1 Message Date
Alex Schapiro aa7a097037 fix(update): never show update prompt/notice in non-interactive runs 2026-07-19 18:08:11 +00:00
Alex Schapiro df80fad960 feat(update): 3-way pre-scan prompt (update now / not now / skip this version) + package-manager upgrade 2026-07-19 17:45:56 +00:00
Alex Schapiro a0bddbc13b fix(update): verify release checksum, clean up staged binary, roll back Windows rename on failure 2026-07-19 17:07:10 +00:00
Alex Schapiro b5e82edb83 feat(cli): update notifications + self-update (strix --update) 2026-07-19 16:56:50 +00:00
Ahmed Allam 7d5a67d234 chore(llm): shorten timeout helper docstring; update tests 2026-07-17 19:45:32 -07:00
Ahmed Allam 88ad3e4472 fix(llm): use a JSON-serializable per-turn model timeout
An httpx.Timeout in ModelSettings.extra_args crashes
ModelSettings.to_json_dict() (PydanticSerializationError) on the Chat
Completions and LiteLLM model paths, which serialize settings for their
tracing generation span — failing every model turn on those paths. Pass
the timeout as a plain float, which httpx-based clients apply as the
read (inactivity) timeout.
2026-07-17 19:45:32 -07:00
Ahmed Allam cf7689e927 fix(llm): use httpx.Timeout read-inactivity for per-turn model timeout 2026-07-17 18:40:23 -07:00
Ahmed Allam 3bb95ab43d fix(llm): add per-turn model request timeout so stalled streams fail fast and retry 2026-07-17 18:40:23 -07:00
Ahmed Allam 9aa151c687 fix(llm): retry statusless mid-stream provider errors (quota/billing)
The SDK's http_status retry policy only retries errors carrying a known
HTTP status code, but quota/billing (and other provider-side) failures
often surface inside a streamed response as a bare error with no status
code, so they were failing on the first attempt. Add a statusless retry
policy to DEFAULT_MODEL_RETRY (retry count and backoff unchanged) so they
are retried before a genuine exhaustion fails the run; user aborts are
never retried.
2026-07-17 16:47:14 -07:00
Ahmed Allam b9c2592b53 fix(llm): retry statusless mid-stream provider errors (quota/billing)
The SDK's http_status retry policy only retries errors carrying a known
HTTP status code, but quota/billing (and other provider-side) failures
often surface inside a streamed response as a bare error with no status
code, so they were failing on the first attempt. Add a statusless retry
policy to DEFAULT_MODEL_RETRY so they are retried (before any content is
streamed; user aborts are never retried), restoring the pre-SDK engine's
resilience. If the provider is genuinely exhausted, the error still
propagates and fails the scan after retries.
2026-07-17 16:47:14 -07:00
devin-ai-integration[bot] f54ecb74f9 fix(report): restore cost tracking for OpenRouter and other LiteLLM-routed models (#801) 2026-07-17 13:38:23 -07:00
devin-ai-integration[bot]andAhmed Allam 96ca7e544d revert(proxy): drop overfit Caido reconnect/HTTPQL band-aids, keep serialization lock (#799)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-07-17 13:18:40 -07:00
21 changed files with 1048 additions and 448 deletions
+1
View File
@@ -434,6 +434,7 @@ SPECIALIZED TOOLS:
PROXY & INTERCEPTION:
- Caido CLI - Modern web proxy (already running). Use the proxy tools
directly, or import `caido_api` from sandbox Python scripts.
- HTTPQL filters (for `list_requests`): quote string values, leave integers unquoted (`resp.code.eq:200`, not `"200"`); combine terms with `AND`/`OR` (there is no `NOT` — use the negated operator `ne`/`ncont`/`nregex`). Numeric fields (`resp.code`, `req.port`) use `eq`/`ne`/`gt`/`gte`/`lt`/`lte`; text fields (`req.host`, `req.path`, `req.method`, `req.raw`) use `cont`/`ncont`/`eq`/`regex`. Example: `resp.code.gte:200 AND resp.code.lt:300 AND req.host.cont:"api"`.
- NOTE: If you are seeing proxy errors when sending requests, it usually means you are not sending requests to a correct url/host/port.
- Ignore Caido proxy-generated 50x HTML error pages; these are proxy issues (might happen when requesting a wrong host or SSL/TLS issues, etc).
+17
View File
@@ -10,6 +10,7 @@ from agents.models.multi_provider import MultiProvider
from agents.retry import (
ModelRetryBackoffSettings,
ModelRetrySettings,
RetryPolicyContext,
retry_policies,
)
@@ -20,6 +21,21 @@ if TYPE_CHECKING:
from strix.config.settings import Settings
def request_timeout_extra_args(timeout_s: float | None) -> dict[str, float] | None:
"""Per-request model timeout; a plain float so ``ModelSettings.to_json_dict()`` stays serializable.""" # noqa: E501
if not timeout_s or timeout_s <= 0:
return None
return {"timeout": timeout_s}
def _retry_statusless_provider_errors(context: RetryPolicyContext) -> bool:
"""Retry statusless provider errors (e.g. mid-stream quota/billing), but not aborts."""
normalized = context.normalized
if normalized.is_abort:
return False
return normalized.status_code is None
class StrixProvider(MultiProvider):
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
so users type ``deepseek/deepseek-chat`` rather than
@@ -56,6 +72,7 @@ DEFAULT_MODEL_RETRY = ModelRetrySettings(
retry_policies.provider_suggested(),
retry_policies.network_error(),
retry_policies.http_status((429, 500, 502, 503, 504)),
_retry_statusless_provider_errors,
),
)
+1 -2
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import logging
import math
from typing import TYPE_CHECKING, Any
from agents.lifecycle import RunHooks
@@ -27,8 +28,6 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]):
"""Persist SDK-native usage after every model response."""
def __init__(self, *, model: str, max_budget_usd: float | None = None) -> None:
import math
if max_budget_usd is not None and (
not math.isfinite(max_budget_usd) or max_budget_usd <= 0
):
+3
View File
@@ -12,6 +12,7 @@ from strix.config.models import (
DEFAULT_MODEL_RETRY,
is_known_openai_bare_model,
model_supports_reasoning,
request_timeout_extra_args,
)
from strix.core.sessions import scrub_images_from_items
@@ -126,11 +127,13 @@ def make_model_settings(
*,
model_name: str,
force_required_tool_choice: bool = False,
request_timeout: float | None = None,
) -> ModelSettings:
model_settings = ModelSettings(
parallel_tool_calls=False,
retry=DEFAULT_MODEL_RETRY,
include_usage=True,
extra_args=request_timeout_extra_args(request_timeout),
)
if (
reasoning_effort is not None
+1 -4
View File
@@ -215,6 +215,7 @@ async def run_strix_scan(
settings.llm.reasoning_effort,
model_name=resolved_model,
force_required_tool_choice=settings.llm.force_required_tool_choice,
request_timeout=settings.llm.timeout,
)
run_config = RunConfig(
model=resolved_model,
@@ -282,10 +283,6 @@ async def run_strix_scan(
context: dict[str, Any] = {
"coordinator": coordinator,
"sandbox_session": bundle["session"],
# One ``SharedCaidoClient`` is reused by every agent in the scan
# (child contexts are shallow copies via ``dict(parent_ctx)``). It
# serializes access to the non-concurrency-safe GraphQL transport
# and rebuilds it if it dies mid-scan.
"caido_client": bundle["caido_client"],
"agent_id": root_id,
"parent_id": None,
+27
View File
@@ -5,6 +5,7 @@ Strix Agent Interface
import argparse
import asyncio
import os
import shutil
import sys
from datetime import UTC, datetime
@@ -32,6 +33,13 @@ from strix.config.models import (
from strix.core.paths import run_dir_for, runtime_state_dir
from strix.interface.cli import run_cli
from strix.interface.tui import run_tui
from strix.interface.update_check import (
is_binary_install,
notify_update,
prompt_update_if_available,
self_update,
start_background_check,
)
from strix.interface.utils import (
assign_workspace_subdirs,
build_final_stats_text,
@@ -447,6 +455,14 @@ Examples:
version=f"strix {get_version()}",
)
parser.add_argument(
"--update",
action="store_true",
help="Update strix to the latest version and exit. Self-updates the "
"standalone binary install; for pip/pipx/uv installs, prints the "
"matching upgrade command instead.",
)
parser.add_argument(
"-t",
"--target",
@@ -565,6 +581,9 @@ Examples:
args = parser.parse_args()
if args.update:
sys.exit(0 if self_update() else 1)
if args.instruction and args.instruction_file:
parser.error(
"Cannot specify both --instruction and --instruction-file. Use one or the other."
@@ -788,6 +807,8 @@ def display_completion_message(args: argparse.Namespace, results_path: Path) ->
"[#60a5fa]discord.gg/strix-ai[/]"
)
console.print()
if not args.non_interactive:
notify_update(console)
def pull_docker_image() -> None:
@@ -851,6 +872,12 @@ def main() -> None:
if args.config:
apply_config_override(validate_config_file(args.config))
start_background_check()
if not args.non_interactive and prompt_update_if_available(Console()):
if is_binary_install() and sys.platform != "win32":
os.execv(sys.executable, sys.argv) # noqa: S606 # nosec B606
sys.exit(0)
check_docker_installed()
pull_docker_image()
+389
View File
@@ -0,0 +1,389 @@
"""Update notifications and self-update for the strix CLI.
Follows the pattern used by tools like gh, uv, and pip: a background,
rate-limited (once per 24h) check against the release source, a cached
result in ``~/.strix``, a non-intrusive notice with the upgrade command
for the detected install method, and a ``strix --update`` self-update
path for the standalone binary install.
"""
from __future__ import annotations
import hashlib
import json
import logging
import os
import platform
import shutil
import stat
import subprocess
import sys
import tarfile
import tempfile
import threading
import time
import zipfile
from pathlib import Path
from typing import cast
import requests
from rich.console import Console
from rich.prompt import Prompt
from strix.telemetry._common import get_version
logger = logging.getLogger(__name__)
GITHUB_REPO = "usestrix/strix"
PYPI_PACKAGE = "strix-agent"
CHECK_INTERVAL_SECONDS = 24 * 60 * 60
REQUEST_TIMEOUT_SECONDS = 5
_CACHE_PATH = Path.home() / ".strix" / "update-check.json"
_background_thread: threading.Thread | None = None
def _is_disabled() -> bool:
return bool(os.environ.get("STRIX_NO_UPDATE_CHECK")) or any(
os.environ.get(key)
for key in ("CI", "GITHUB_ACTIONS", "GITLAB_CI", "JENKINS_URL", "BUILDKITE", "CIRCLECI")
)
def is_binary_install() -> bool:
return bool(getattr(sys, "frozen", False))
def get_install_method() -> str:
if is_binary_install():
return "binary"
prefix = str(Path(sys.prefix)).replace("\\", "/")
if "/pipx/" in prefix or prefix.endswith("/pipx"):
return "pipx"
if "/uv/tools/" in prefix:
return "uv"
return "pip"
def get_upgrade_command(method: str | None = None) -> str:
method = method or get_install_method()
commands = {
"binary": "strix --update",
"pipx": "pipx upgrade strix-agent",
"uv": "uv tool upgrade strix-agent",
"pip": "pip install --upgrade strix-agent",
}
return commands[method]
def _parse_version(value: str) -> tuple[int, ...] | None:
parts = value.strip().lstrip("v").split(".")
try:
return tuple(int(part) for part in parts)
except ValueError:
return None
def _is_newer(latest: str, current: str) -> bool:
latest_parts = _parse_version(latest)
current_parts = _parse_version(current)
if latest_parts is None or current_parts is None:
return False
return latest_parts > current_parts
def _fetch_latest_version() -> str | None:
try:
if is_binary_install():
response = requests.get(
f"https://api.github.com/repos/{GITHUB_REPO}/releases/latest",
timeout=REQUEST_TIMEOUT_SECONDS,
)
response.raise_for_status()
tag = response.json().get("tag_name", "")
return tag.lstrip("v") or None
response = requests.get(
f"https://pypi.org/pypi/{PYPI_PACKAGE}/json",
timeout=REQUEST_TIMEOUT_SECONDS,
)
response.raise_for_status()
version = response.json().get("info", {}).get("version")
return str(version) if version else None
except Exception: # noqa: BLE001
logger.debug("update check failed", exc_info=True)
return None
def _fetch_asset_digest(version: str, filename: str) -> str | None:
"""Return the expected sha256 (hex) for a release asset, if the API provides one."""
try:
response = requests.get(
f"https://api.github.com/repos/{GITHUB_REPO}/releases/tags/v{version}",
timeout=REQUEST_TIMEOUT_SECONDS,
)
response.raise_for_status()
for asset in response.json().get("assets", []):
if asset.get("name") == filename:
digest = asset.get("digest") or ""
if digest.startswith("sha256:"):
return digest.removeprefix("sha256:")
except Exception: # noqa: BLE001
logger.debug("release asset digest lookup failed", exc_info=True)
return None
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as f:
for chunk in iter(lambda: f.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def _read_cache() -> dict[str, object]:
try:
with _CACHE_PATH.open(encoding="utf-8") as f:
data = json.load(f)
if isinstance(data, dict):
return cast("dict[str, object]", data)
except Exception: # noqa: BLE001, S110
pass # nosec B110
return {}
def _write_cache(**fields: object) -> None:
try:
cache = _read_cache()
cache.update(fields)
_CACHE_PATH.parent.mkdir(parents=True, exist_ok=True)
_CACHE_PATH.write_text(json.dumps(cache), encoding="utf-8")
except Exception: # noqa: BLE001, S110
pass # nosec B110
def skip_version(version: str) -> None:
"""Remember not to prompt again for this version (newer releases still notify)."""
_write_cache(skipped_version=version)
def _refresh_cache() -> None:
latest = _fetch_latest_version()
if latest:
_write_cache(latest_version=latest, checked_at=time.time())
def start_background_check() -> None:
"""Refresh the cached latest-version info in a daemon thread (at most once per 24h)."""
global _background_thread # noqa: PLW0603
if _is_disabled():
return
cache = _read_cache()
checked_at = cache.get("checked_at")
if isinstance(checked_at, int | float) and time.time() - checked_at < CHECK_INTERVAL_SECONDS:
return
_background_thread = threading.Thread(target=_refresh_cache, daemon=True)
_background_thread.start()
def get_available_update(*, respect_skip: bool = True) -> str | None:
"""Return the newer version from the cache, or None if up to date / unknown."""
if _is_disabled():
return None
if _background_thread is not None:
_background_thread.join(timeout=0.2)
cache = _read_cache()
latest = cache.get("latest_version")
current = get_version()
if not isinstance(latest, str) or current == "unknown" or not _is_newer(latest, current):
return None
if respect_skip and cache.get("skipped_version") == latest:
return None
return latest
def notify_update(console: Console) -> None:
latest = get_available_update()
if not latest:
return
console.print(
f"[#eab308]A new version of strix is available:[/] "
f"[dim]{get_version()}[/] [dim]→[/] [bold #22c55e]{latest}[/]"
f" [dim]·[/] [#60a5fa]{get_upgrade_command()}[/]"
)
console.print()
def run_package_upgrade(console: Console, method: str) -> bool:
"""Upgrade a package-manager install by running its upgrade command."""
command = get_upgrade_command(method).split()
console.print(f"[dim]Running[/] [#60a5fa]{' '.join(command)}[/]")
try:
result = subprocess.run(command, check=False) # noqa: S603
except OSError as e:
console.print(f"[bold red]Update failed:[/] {e}")
return False
if result.returncode != 0:
console.print(
f"[bold red]Update failed[/] [dim](exit code {result.returncode}).[/] "
f"Run it manually: [#60a5fa]{get_upgrade_command(method)}[/]"
)
return False
console.print("[#22c55e]✓ strix updated — restart the scan to use the new version[/]")
return True
def prompt_update_if_available(console: Console) -> bool:
"""Offer an interactive update before a scan starts.
Returns True if strix was updated (caller should re-exec / exit).
"""
latest = get_available_update()
if not latest or not sys.stdin.isatty() or not sys.stdout.isatty():
return False
console.print()
console.print(
f"[#eab308]A new version of strix is available:[/] "
f"[dim]{get_version()}[/] [dim]→[/] [bold #22c55e]{latest}[/]"
)
console.print(
"[dim] y — update now n — not now (ask again next run) s — skip this version[/]"
)
choice = Prompt.ask("Update strix?", choices=["y", "n", "s"], default="n")
console.print()
if choice == "s":
skip_version(latest)
return False
if choice != "y":
return False
method = get_install_method()
if method == "binary":
return self_update(console, version=latest)
return run_package_upgrade(console, method)
def _release_target() -> str | None:
raw_os = platform.system().lower()
os_name = {"darwin": "macos", "linux": "linux", "windows": "windows"}.get(raw_os)
arch = platform.machine().lower()
arch = {"aarch64": "arm64", "amd64": "x86_64"}.get(arch, arch)
if os_name is None:
return None
target = f"{os_name}-{arch}"
supported = {"linux-x86_64", "macos-x86_64", "macos-arm64", "windows-x86_64"}
return target if target in supported else None
def _download_and_replace(version: str, target: str, console: Console) -> bool:
is_windows = target.startswith("windows")
archive_ext = ".zip" if is_windows else ".tar.gz"
filename = f"strix-{version}-{target}{archive_ext}"
url = f"https://github.com/{GITHUB_REPO}/releases/download/v{version}/{filename}"
binary_name = f"strix-{version}-{target}" + (".exe" if is_windows else "")
current_exe = Path(sys.executable).resolve()
with tempfile.TemporaryDirectory() as tmp:
tmp_dir = Path(tmp)
archive_path = tmp_dir / filename
console.print(f"[dim]Downloading[/] {url}")
with requests.get( # nosec B113
url,
stream=True,
timeout=REQUEST_TIMEOUT_SECONDS * 12,
) as response:
response.raise_for_status()
with archive_path.open("wb") as f:
for chunk in response.iter_content(chunk_size=1 << 20):
f.write(chunk)
expected_digest = _fetch_asset_digest(version, filename)
if expected_digest:
actual_digest = _sha256_file(archive_path)
if actual_digest != expected_digest:
raise RuntimeError(
f"checksum mismatch for {filename}: "
f"expected sha256 {expected_digest}, got {actual_digest}"
)
else:
console.print("[dim yellow]No published checksum available; skipping verification[/]")
if is_windows:
with zipfile.ZipFile(archive_path) as zf:
zf.extract(binary_name, tmp_dir)
else:
with tarfile.open(archive_path, "r:gz") as tf:
tf.extract(binary_name, tmp_dir, filter="data")
new_binary = tmp_dir / binary_name
new_binary.chmod(new_binary.stat().st_mode | stat.S_IXUSR | stat.S_IXGRP | stat.S_IXOTH)
staged = current_exe.with_name(current_exe.name + ".new")
try:
shutil.copy2(new_binary, staged)
if is_windows:
# Windows can't replace a running executable in place; move it aside first.
old = current_exe.with_name(current_exe.name + ".old")
old.unlink(missing_ok=True)
current_exe.rename(old)
try:
staged.replace(current_exe)
except Exception:
old.rename(current_exe)
raise
else:
staged.replace(current_exe)
except Exception:
staged.unlink(missing_ok=True)
raise
return True
def self_update(console: Console | None = None, version: str | None = None) -> bool:
"""Replace the running standalone binary with the latest release.
Returns True on success. For package-manager installs this only
prints the right upgrade command and returns False.
"""
console = console or Console()
if not is_binary_install():
method = get_install_method()
console.print(
f"[#eab308]This strix was installed via {method};[/] "
f"upgrade it with: [#60a5fa]{get_upgrade_command(method)}[/]"
)
return False
latest = version or _fetch_latest_version()
if not latest:
console.print("[bold red]Could not determine the latest strix version.[/]")
return False
current = get_version()
if current != "unknown" and not _is_newer(latest, current):
console.print(f"[#22c55e]strix {current} is already the latest version.[/]")
return True
target = _release_target()
if not target:
console.print(
f"[bold red]No prebuilt binary for this platform "
f"({platform.system()}/{platform.machine()}).[/]"
)
return False
try:
_download_and_replace(latest, target, console)
except Exception as e: # noqa: BLE001
logger.debug("self-update failed", exc_info=True)
console.print(f"[bold red]Update failed:[/] {e}")
console.print(
"[dim]You can reinstall manually with:[/] "
"[#60a5fa]curl -sSL https://strix.ai/install | bash[/]"
)
return False
_write_cache(latest_version=latest, checked_at=time.time())
console.print(f"[#22c55e]✓ Updated strix to {latest}[/]")
return True
+6 -1
View File
@@ -16,6 +16,7 @@ from strix.config.models import (
DEFAULT_MODEL_RETRY,
StrixProvider,
configure_sdk_model_defaults,
request_timeout_extra_args,
)
from strix.report.state import get_global_report_state
@@ -310,7 +311,11 @@ async def check_duplicate(
response = await model.get_response(
system_instructions=DEDUPE_SYSTEM_PROMPT,
input=user_msg,
model_settings=ModelSettings(retry=DEFAULT_MODEL_RETRY, include_usage=True),
model_settings=ModelSettings(
retry=DEFAULT_MODEL_RETRY,
include_usage=True,
extra_args=request_timeout_extra_args(settings.llm.timeout),
),
tools=[],
output_schema=None,
handoffs=[],
+102 -10
View File
@@ -534,16 +534,10 @@ def litellm_cost_callback(
cost = value
if cost is None:
usage: Any = getattr(completion_response, "usage", None)
if usage is None and isinstance(completion_response, dict):
usage = cast("dict[str, Any]", completion_response).get("usage")
usage_cost: Any
if isinstance(usage, dict):
usage_cost = cast("dict[str, Any]", usage).get("cost")
else:
usage_cost = getattr(usage, "cost", None)
if isinstance(usage_cost, int | float) and usage_cost > 0:
cost = float(usage_cost)
cost = _usage_reported_cost(completion_response)
if cost is None:
cost = _estimate_response_cost(kwargs, completion_response)
if cost is None or cost <= 0:
return
@@ -554,3 +548,101 @@ def litellm_cost_callback(
report_state.record_observed_llm_cost(cost)
except Exception:
logger.exception("Failed to record observed LiteLLM cost")
def _usage_reported_cost(completion_response: Any) -> float | None:
"""Provider-reported cost from the ``usage`` block (e.g. OpenRouter).
Non-BYOK responses charge everything to ``usage.cost``. BYOK responses
charge only the OpenRouter fee to ``usage.cost`` (often 0) and report the
provider charge in ``usage.cost_details.upstream_inference_cost``, so the
true BYOK total is the sum of the two.
"""
usage: Any = getattr(completion_response, "usage", None)
if usage is None and isinstance(completion_response, dict):
usage = cast("dict[str, Any]", completion_response).get("usage")
if usage is None:
return None
def _field(container: Any, name: str) -> Any:
if isinstance(container, dict):
return cast("dict[str, Any]", container).get(name)
return getattr(container, name, None)
total = 0.0
usage_cost = _field(usage, "cost")
if isinstance(usage_cost, int | float) and usage_cost > 0:
total += float(usage_cost)
if bool(_field(usage, "is_byok")):
upstream = _field(_field(usage, "cost_details"), "upstream_inference_cost")
if isinstance(upstream, int | float) and upstream > 0:
total += float(upstream)
return total if total > 0 else None
def _estimate_response_cost(kwargs: Any, completion_response: Any) -> float | None:
"""Best-effort LiteLLM cost-map estimate when no provider-reported cost exists.
LiteLLM strips provider cost fields when rebuilding streamed responses and
returns no ``response_cost`` for models missing from its cost map, so try
the provider-prefixed name, the raw name, and the bare model name.
"""
from litellm import completion_cost
model = kwargs.get("model") if isinstance(kwargs, dict) else None
if not isinstance(model, str) or not model:
if isinstance(completion_response, dict):
model = cast("dict[str, Any]", completion_response).get("model")
else:
model = getattr(completion_response, "model", None)
if not isinstance(model, str) or not model:
return None
provider = None
litellm_params = kwargs.get("litellm_params") if isinstance(kwargs, dict) else None
if isinstance(litellm_params, dict):
provider = litellm_params.get("custom_llm_provider")
usage_payload = _usage_payload(completion_response)
if usage_payload is None:
return None
candidates: list[str] = []
if isinstance(provider, str) and provider and not model.startswith(f"{provider}/"):
candidates.append(f"{provider}/{model}")
candidates.append(model)
if "/" in model:
candidates.append(model.rsplit("/", 1)[-1])
for candidate in candidates:
try:
value = completion_cost(
completion_response={"model": candidate, "usage": usage_payload},
model=candidate,
)
except Exception: # nosec B112 # noqa: BLE001, S112
continue
if isinstance(value, int | float) and value > 0:
return float(value)
return None
def _usage_payload(completion_response: Any) -> dict[str, Any] | None:
"""Token counts as a plain dict, detached from the response's provider metadata."""
usage: Any = getattr(completion_response, "usage", None)
if usage is None and isinstance(completion_response, dict):
usage = cast("dict[str, Any]", completion_response).get("usage")
if usage is None:
return None
if hasattr(usage, "model_dump"):
usage = usage.model_dump()
if not isinstance(usage, dict):
return None
payload = cast("dict[str, Any]", usage)
if not payload.get("total_tokens") and not (
payload.get("prompt_tokens") or payload.get("completion_tokens")
):
return None
return payload
+11 -53
View File
@@ -80,72 +80,30 @@ async def _login_as_guest(
raise RuntimeError(f"loginAsGuest failed after {attempts} attempts: {last_err}")
async def _aclose_quietly(client: Client) -> None:
"""Best-effort close of a client whose setup failed; never raises."""
with contextlib.suppress(Exception):
await client.aclose()
async def _connect_client(
session: BaseSandboxSession,
*,
host_url: str,
container_url: str,
) -> Client:
access_token = await _login_as_guest(session, container_url=container_url)
client = Client(host_url, auth=TokenAuthOptions(token=access_token))
await client.connect()
return client
async def bootstrap_caido(
session: BaseSandboxSession,
*,
host_url: str,
container_url: str,
) -> tuple[Client, str]:
"""Connect to the in-container Caido sidecar and select a fresh project.
Returns the connected client and the id of the temporary project it
selected. The project id lets :func:`reconnect_caido` rebuild a dead
transport while staying on the *same* project (and its captured traffic)
instead of creating a new empty one.
"""
) -> Client:
"""Connect to the in-container Caido sidecar and select a fresh project."""
logger.info("Bootstrapping Caido client (host=%s, container=%s)", host_url, container_url)
client = await _connect_client(session, host_url=host_url, container_url=container_url)
access_token = await _login_as_guest(session, container_url=container_url)
client = Client(host_url, auth=TokenAuthOptions(token=access_token))
await client.connect()
try:
project = await client.project.create(
CreateProjectOptions(name="sandbox", temporary=True),
)
await client.project.select(project.id)
except BaseException:
# Don't leak the connected transport if project setup fails.
await _aclose_quietly(client)
# The connected client never reaches the session bundle if project
# setup fails, so close it here to avoid leaking the transport.
with contextlib.suppress(Exception):
await client.aclose()
raise
logger.info("Caido project selected: %s", project.id)
return client, str(project.id)
async def reconnect_caido(
session: BaseSandboxSession,
*,
host_url: str,
container_url: str,
project_id: str,
) -> Client:
"""Rebuild a Caido client after its transport died, keeping the project.
Re-authenticates, reconnects, and re-selects the existing project so the
caller keeps access to the traffic captured before the disconnect.
"""
logger.info("Reconnecting Caido client (host=%s, project=%s)", host_url, project_id)
client = await _connect_client(session, host_url=host_url, container_url=container_url)
try:
await client.project.select(project_id)
except BaseException:
# A missing/unavailable project must not leave the freshly-connected
# transport dangling — otherwise every retry leaks another one.
await _aclose_quietly(client)
raise
return client
+4 -17
View File
@@ -5,20 +5,15 @@ from __future__ import annotations
import logging
import shutil
from pathlib import Path
from typing import TYPE_CHECKING, Any
from typing import Any
from agents.sandbox.entries import BaseEntry, LocalDir
from agents.sandbox.manifest import Environment, Manifest
from strix.config import load_settings
from strix.runtime.backends import get_backend
from strix.runtime.caido_bootstrap import bootstrap_caido, reconnect_caido
from strix.runtime.caido_bootstrap import bootstrap_caido
from strix.runtime.local_dir_staging import stage_symlink_safe_dir
from strix.tools.proxy.caido_api import SharedCaidoClient
if TYPE_CHECKING:
from caido_sdk_client import Client as CaidoClient
logger = logging.getLogger(__name__)
@@ -136,24 +131,16 @@ async def create_or_reuse(
host_caido_url = f"{scheme}://{caido_endpoint.host}:{caido_endpoint.port}"
logger.debug("Caido host endpoint resolved: %s", host_caido_url)
caido_client, caido_project_id = await bootstrap_caido(
caido_client = await bootstrap_caido(
session,
host_url=host_caido_url,
container_url=container_caido_url,
)
async def _reconnect_caido() -> CaidoClient:
return await reconnect_caido(
session,
host_url=host_caido_url,
container_url=container_caido_url,
project_id=caido_project_id,
)
bundle = {
"client": client,
"session": session,
"caido_client": SharedCaidoClient(caido_client, _reconnect_caido),
"caido_client": caido_client,
}
_SESSION_CACHE[scan_id] = bundle
logger.info("Sandbox session for scan %s ready and cached", scan_id)
+7 -99
View File
@@ -3,9 +3,7 @@
from __future__ import annotations
import asyncio
import contextlib
import json
import logging
import os
import time
import urllib.request
@@ -28,9 +26,6 @@ if TYPE_CHECKING:
from caido_sdk_client import Client as CaidoClient
logger = logging.getLogger(__name__)
RequestPart = Literal["request", "response"]
SortBy = Literal[
"timestamp",
@@ -50,18 +45,6 @@ _SITEMAP_PAGE_SIZE = 30
_DEFAULT_CAIDO_URL = "http://127.0.0.1:48080"
_CLIENT_CACHE: dict[str, Client] = {}
_CLIENT_LOCK = asyncio.Lock()
# Substrings that mean the shared client's transport has died or is being used
# concurrently — recoverable by rebuilding the client and retrying once.
_CONNECTION_ERROR_MARKERS = (
"transport is already connected",
"connector is closed",
"server disconnected",
"session is closed",
"cannot write to closing transport",
"connection reset",
"connection closed",
)
_REQ_FIELD_MAP: dict[SortBy, tuple[str, str]] = {
"timestamp": ("req", "created_at"),
"host": ("req", "host"),
@@ -108,22 +91,6 @@ async def _new_client() -> Client:
return client
async def _safe_aclose(client: Client | None) -> None:
"""Close a (possibly dead) client without letting teardown errors escape."""
if client is None:
return
with contextlib.suppress(Exception):
await client.aclose()
def _is_connection_error(exc: BaseException) -> bool:
message = str(exc).lower()
if any(marker in message for marker in _CONNECTION_ERROR_MARKERS):
return True
cause = exc.__cause__ or exc.__context__
return cause is not None and cause is not exc and _is_connection_error(cause)
async def get_client() -> Client:
"""Return the shared Caido client, creating it under a lock if needed.
@@ -139,73 +106,19 @@ async def get_client() -> Client:
return client
async def call_with_client[T](
fn: Callable[[Client], Awaitable[T]], *, idempotent: bool = True
) -> T:
"""Run ``fn`` against the shared client, serialized and reconnect-safe.
async def call_with_client[T](fn: Callable[[Client], Awaitable[T]]) -> T:
"""Run ``fn`` against the shared client, serialized through ``_CLIENT_LOCK``.
The Caido GraphQL transport is not safe for concurrent use: two in-flight
requests race and raise "Transport is already connected". All proxy calls
are therefore serialized through ``_CLIENT_LOCK``. If the cached client's
transport has since died ("Connector is closed" / "Server disconnected"),
the stale client is closed and rebuilt so subsequent calls stop failing
against a dead client.
``fn`` is only re-run automatically when ``idempotent`` is true. For
mutations (replay, scope create/update/delete) a connection error may
arrive *after* Caido applied the change, so we heal the client for future
calls but re-raise instead of risking a double-apply.
requests race and raise "Transport is already connected". Serializing every
proxy call through the lock prevents that.
"""
async with _CLIENT_LOCK:
client = _CLIENT_CACHE.get("default")
if client is None:
client = await _new_client()
_CLIENT_CACHE["default"] = client
try:
return await fn(client)
except Exception as exc:
if not _is_connection_error(exc):
raise
new_client = await _new_client()
_CLIENT_CACHE["default"] = new_client
await _safe_aclose(client)
if not idempotent:
raise
return await fn(new_client)
class SharedCaidoClient:
"""Serialized, reconnect-safe wrapper around one host-side Caido client.
Every agent in a scan shares a single instance (propagated through the
shallow-copied run context). ``call`` serializes access — the SDK transport
is not concurrency-safe — and, when the transport dies, rebuilds the client
via ``reconnect`` (which preserves the Caido project) and closes the dead
one, so a transient Caido restart no longer disables proxy tools for the
rest of the scan.
"""
def __init__(self, client: Client, reconnect: Callable[[], Awaitable[Client]]) -> None:
self._client = client
self._reconnect = reconnect
self._lock = asyncio.Lock()
async def call[T](self, fn: Callable[[Client], Awaitable[T]], *, idempotent: bool = True) -> T:
async with self._lock:
try:
return await fn(self._client)
except Exception as exc:
if not _is_connection_error(exc):
raise
dead, self._client = self._client, await self._reconnect()
await _safe_aclose(dead)
if not idempotent:
raise
return await fn(self._client)
async def aclose(self) -> None:
async with self._lock:
await _safe_aclose(self._client)
return await fn(client)
async def close_client() -> None:
@@ -546,9 +459,7 @@ async def repeat_request(
)
return await replay_send_raw(client, raw=raw, connection=connection)
# A replay mutates server state; don't auto-retry if the transport dies
# mid-send (the request may already have been sent).
return await call_with_client(_run, idempotent=False)
return await call_with_client(_run)
async def scope_rules(
@@ -569,8 +480,7 @@ async def scope_rules(
scope_name=scope_name,
)
# get/list are read-only and safe to retry; create/update/delete mutate.
return await call_with_client(_run, idempotent=action in {"get", "list"})
return await call_with_client(_run)
async def _scope_rules_with_client(
@@ -819,11 +729,9 @@ async def view_sitemap_entry(entry_id: str) -> dict[str, Any]:
__all__ = [
"RequestPart",
"ScopeAction",
"SharedCaidoClient",
"SitemapDepth",
"SortBy",
"SortOrder",
"call_with_client",
"close_client",
"get_client",
"list_requests",
+48 -77
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import asyncio
import dataclasses
import json
import logging
@@ -13,13 +14,14 @@ from typing import TYPE_CHECKING, Any, Literal
from agents import RunContextWrapper, function_tool
from strix.tools.proxy import caido_api
from strix.tools.proxy.caido_api import SharedCaidoClient
logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from caido_sdk_client import Client
from strix.tools.proxy.caido_api import (
@@ -29,7 +31,7 @@ if TYPE_CHECKING:
SortOrder,
)
else:
from strix.tools.proxy.caido_api import (
from strix.tools.proxy.caido_api import ( # noqa: TC001
RequestPart,
SitemapDepth,
SortBy,
@@ -39,19 +41,21 @@ else:
ScopeAction = Literal["get", "list", "create", "update", "delete"]
# All agents in a scan share one host-side Caido client whose GraphQL transport
# is not concurrency-safe (parallel calls raise "Transport is already
# connected"). Serialize every host-side proxy call through this lock.
_CAIDO_CALL_LOCK = asyncio.Lock()
def _ctx_proxy(ctx: RunContextWrapper) -> SharedCaidoClient | None:
"""Return the scan-wide serialized, reconnect-safe Caido client holder.
All agents in a scan share one :class:`SharedCaidoClient` whose GraphQL
transport is not concurrency-safe (parallel calls raise "Transport is
already connected"). ``SharedCaidoClient.call`` serializes access and
rebuilds the transport if it dies mid-scan. Returns ``None`` when no holder
is present (e.g. standalone tool invocation outside a scan run).
"""
def _ctx_client(ctx: RunContextWrapper) -> Client | None:
inner = ctx.context if isinstance(ctx.context, dict) else {}
proxy = inner.get("caido_client")
return proxy if isinstance(proxy, SharedCaidoClient) else None
return inner.get("caido_client")
async def _call[T](client: Client, fn: Callable[[Client], Awaitable[T]]) -> T:
"""Run ``fn`` against the shared client, serialized under ``_CAIDO_CALL_LOCK``."""
async with _CAIDO_CALL_LOCK:
return await fn(client)
def _to_tool_json(value: Any) -> Any:
@@ -93,39 +97,6 @@ def _err(name: str, exc: Exception) -> str:
)
_HTTPQL_HINT = (
"HTTPQL syntax: quote string values and leave integers unquoted; combine "
"terms with AND / OR (there is no NOT). Numeric fields (resp.code, req.port, "
"id, roundtrip) use eq/ne/gt/gte/lt/lte; text/byte fields (req.host, req.path, "
"req.method, req.raw, resp.raw) use cont/ncont/eq/ne/like/nlike/regex/nregex. "
"Example: 'resp.code.gte:200 AND resp.code.lt:300 AND req.host.cont:\"api\"'."
)
def _is_httpql_error(exc: Exception) -> bool:
message = str(exc).lower()
return "httpql" in message or ("filter" in message and "pars" in message)
def _httpql_error(exc: Exception, httpql_filter: str | None) -> str:
"""Return an actionable error for a rejected HTTPQL filter.
Preserves Caido's exact parser message and echoes the offending query so
the agent can self-correct instead of retrying the same broken filter.
"""
logger.info("list_requests rejected HTTPQL filter %r: %s", httpql_filter, exc)
return json.dumps(
{
"success": False,
"error": f"Invalid HTTPQL filter: {exc}",
"httpql_filter": httpql_filter,
"hint": _HTTPQL_HINT,
},
ensure_ascii=False,
default=str,
)
@function_tool(timeout=120)
async def list_requests(
ctx: RunContextWrapper,
@@ -184,12 +155,13 @@ async def list_requests(
sort_order: ``asc`` or ``desc``.
scope_id: Restrict to a Caido scope (managed via ``scope_rules``).
"""
proxy = _ctx_proxy(ctx)
if proxy is None:
client = _ctx_client(ctx)
if client is None:
return _no_client()
try:
connection = await proxy.call(
connection = await _call(
client,
lambda client: caido_api.list_requests_with_client(
client,
httpql_filter=httpql_filter,
@@ -198,7 +170,7 @@ async def list_requests(
sort_by=sort_by,
sort_order=sort_order,
scope_id=scope_id,
)
),
)
entries = []
@@ -252,8 +224,6 @@ async def list_requests(
default=str,
)
except Exception as exc: # noqa: BLE001
if httpql_filter and _is_httpql_error(exc):
return _httpql_error(exc, httpql_filter)
return _err("list_requests", exc)
@@ -291,13 +261,14 @@ async def view_request(
page: 1-indexed page number (only when no ``search_pattern``).
page_size: Lines per page.
"""
proxy = _ctx_proxy(ctx)
if proxy is None:
client = _ctx_client(ctx)
if client is None:
return _no_client()
try:
result = await proxy.call(
lambda client: caido_api.get_request_with_client(client, request_id, part=part)
result = await _call(
client,
lambda client: caido_api.get_request_with_client(client, request_id, part=part),
)
if result is None:
return json.dumps(
@@ -408,8 +379,8 @@ async def repeat_request(
- ``body`` — replace the body string entirely.
- ``cookies`` — dict of cookies to add/update.
"""
proxy = _ctx_proxy(ctx)
if proxy is None:
client = _ctx_client(ctx)
if client is None:
return _no_client()
mods = modifications or {}
@@ -431,9 +402,7 @@ async def repeat_request(
return await caido_api.replay_send_raw(client, raw=raw, connection=connection)
try:
# A replay mutates target state, so don't auto-retry on a mid-send
# transport failure (the request may already have been sent).
replay = await proxy.call(_do, idempotent=False)
replay = await _call(client, _do)
if replay is None:
return json.dumps(
{"success": False, "error": f"Request {request_id} not found"},
@@ -492,18 +461,19 @@ async def list_sitemap(
(recursive subtree). Only meaningful with ``parent_id``.
page: 1-indexed page (30 entries per page).
"""
proxy = _ctx_proxy(ctx)
if proxy is None:
client = _ctx_client(ctx)
if client is None:
return _no_client()
try:
payload = await proxy.call(
payload = await _call(
client,
lambda client: caido_api.list_sitemap_with_client(
client,
scope_id=scope_id,
parent_id=parent_id,
depth=depth,
page=page,
)
),
)
return json.dumps(payload, ensure_ascii=False, default=str)
except Exception as exc: # noqa: BLE001
@@ -525,12 +495,13 @@ async def view_sitemap_entry(
Args:
entry_id: ID from ``list_sitemap`` (or any nested entry).
"""
proxy = _ctx_proxy(ctx)
if proxy is None:
client = _ctx_client(ctx)
if client is None:
return _no_client()
try:
payload = await proxy.call(
lambda client: caido_api.view_sitemap_entry_with_client(client, entry_id)
payload = await _call(
client,
lambda client: caido_api.view_sitemap_entry_with_client(client, entry_id),
)
return json.dumps(payload, ensure_ascii=False, default=str)
except Exception as exc: # noqa: BLE001
@@ -583,13 +554,13 @@ async def scope_rules(
scope_id: Required for ``get`` / ``update`` / ``delete``.
scope_name: Required for ``create`` / ``update``.
"""
proxy = _ctx_proxy(ctx)
if proxy is None:
client = _ctx_client(ctx)
if client is None:
return _no_client()
try:
if action == "list":
scopes = await proxy.call(caido_api.scope_list)
scopes = await _call(client, caido_api.scope_list)
return json.dumps(
{"success": True, "scopes": [_to_tool_json(s) for s in scopes]},
ensure_ascii=False,
@@ -602,7 +573,7 @@ async def scope_rules(
ensure_ascii=False,
default=str,
)
scope = await proxy.call(lambda client: caido_api.scope_get(client, scope_id))
scope = await _call(client, lambda client: caido_api.scope_get(client, scope_id))
return json.dumps(
{"success": True, "scope": _to_tool_json(scope)},
ensure_ascii=False,
@@ -615,11 +586,11 @@ async def scope_rules(
ensure_ascii=False,
default=str,
)
scope = await proxy.call(
scope = await _call(
client,
lambda client: caido_api.scope_create(
client, name=scope_name, allowlist=allowlist, denylist=denylist
),
idempotent=False,
)
return json.dumps(
{"success": True, "scope": _to_tool_json(scope)},
@@ -636,11 +607,11 @@ async def scope_rules(
ensure_ascii=False,
default=str,
)
scope = await proxy.call(
scope = await _call(
client,
lambda client: caido_api.scope_update(
client, scope_id, name=scope_name, allowlist=allowlist, denylist=denylist
),
idempotent=False,
)
return json.dumps(
{"success": True, "scope": _to_tool_json(scope)},
@@ -653,7 +624,7 @@ async def scope_rules(
ensure_ascii=False,
default=str,
)
await proxy.call(lambda client: caido_api.scope_delete(client, scope_id), idempotent=False)
await _call(client, lambda client: caido_api.scope_delete(client, scope_id))
return json.dumps(
{
"success": True,
+109
View File
@@ -6,6 +6,7 @@ from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import litellm
import pytest
from strix.config.models import _configure_litellm_compatibility
from strix.report.state import litellm_cost_callback
@@ -42,3 +43,111 @@ def test_cost_callback_reads_usage_cost_from_mapping_response() -> None:
litellm_cost_callback({}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.125)
def test_cost_callback_reads_byok_upstream_inference_cost() -> None:
report_state = MagicMock()
response = SimpleNamespace(
usage=SimpleNamespace(
cost=0,
is_byok=True,
cost_details=SimpleNamespace(upstream_inference_cost=6.75e-06),
),
_hidden_params={},
)
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({"response_cost": None}, response)
report_state.record_observed_llm_cost.assert_called_once_with(6.75e-06)
def test_cost_callback_sums_usage_cost_and_upstream_inference_cost() -> None:
report_state = MagicMock()
response = {
"usage": {
"cost": 0.01,
"is_byok": True,
"cost_details": {"upstream_inference_cost": 0.2},
}
}
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({}, response)
report_state.record_observed_llm_cost.assert_called_once_with(pytest.approx(0.21))
def test_cost_callback_ignores_upstream_cost_for_non_byok_responses() -> None:
report_state = MagicMock()
response = {
"usage": {
"cost": 0.05,
"is_byok": False,
"cost_details": {"upstream_inference_cost": 0.04},
}
}
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.05)
def test_cost_callback_estimates_cost_with_provider_prefixed_model() -> None:
report_state = MagicMock()
response = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
kwargs = {
"response_cost": None,
"model": "anthropic/claude-sonnet-4.5",
"litellm_params": {"custom_llm_provider": "openrouter"},
}
def fake_completion_cost(**kwargs: object) -> float:
if kwargs["model"] == "openrouter/anthropic/claude-sonnet-4.5":
return 0.5
raise ValueError(kwargs["model"])
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=fake_completion_cost),
):
litellm_cost_callback(kwargs, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.5)
def test_cost_callback_estimates_cost_with_bare_model_fallback() -> None:
report_state = MagicMock()
response = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
kwargs = {
"response_cost": None,
"model": "openai/gpt-4o-mini",
"litellm_params": {"custom_llm_provider": "openrouter"},
}
def fake_completion_cost(**kwargs: object) -> float:
if kwargs["model"] == "gpt-4o-mini":
return 0.025
raise ValueError(kwargs["model"])
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=fake_completion_cost),
):
litellm_cost_callback(kwargs, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.025)
def test_cost_callback_records_nothing_when_no_cost_available() -> None:
report_state = MagicMock()
response = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
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": "x/y"}, response)
report_state.record_observed_llm_cost.assert_not_called()
+30
View File
@@ -155,3 +155,33 @@ def test_make_model_settings_forces_required_for_anyllm_routed_openai_model() ->
)
assert settings.tool_choice == "required"
def test_make_model_settings_sets_request_timeout() -> None:
settings = make_model_settings(
"none",
model_name="gpt-4o",
request_timeout=300.0,
)
assert settings.extra_args is not None
assert settings.extra_args["timeout"] == 300.0
def test_make_model_settings_omits_timeout_when_unset() -> None:
settings = make_model_settings("none", model_name="gpt-4o")
assert settings.extra_args is None
def test_make_model_settings_timeout_survives_reasoning_resolve() -> None:
# Reasoning is resolved via ModelSettings.resolve(); the timeout in extra_args
# must not be dropped when a reasoning override is merged in.
settings = make_model_settings(
"high",
model_name="openai/o3",
request_timeout=120.0,
)
assert settings.extra_args is not None
assert settings.extra_args["timeout"] == 120.0
+77
View File
@@ -0,0 +1,77 @@
"""Tests for the model retry policy used by every agent model call.
The SDK's built-in ``http_status`` policy only retries errors that carry a known
HTTP status code. Quota/billing (and other provider-side) failures often surface
*inside* a streamed response as a bare error with no status code, so Strix adds a
statusless retry policy to ``DEFAULT_MODEL_RETRY`` to keep them recoverable — the
behavior the pre-SDK engine had.
"""
from __future__ import annotations
import asyncio
from agents.retry import ModelRetryNormalizedError, RetryPolicyContext
from strix.config.models import DEFAULT_MODEL_RETRY, _retry_statusless_provider_errors
def _context(normalized: ModelRetryNormalizedError) -> RetryPolicyContext:
return RetryPolicyContext(
error=RuntimeError("boom"),
attempt=1,
max_retries=5,
stream=True,
normalized=normalized,
provider_advice=None,
)
def _retries(normalized: ModelRetryNormalizedError) -> bool:
"""Evaluate the composed DEFAULT_MODEL_RETRY policy for a normalized error."""
policy = DEFAULT_MODEL_RETRY.policy
assert policy is not None
decision = asyncio.run(policy(_context(normalized)))
return bool(getattr(decision, "retry", decision))
def test_statusless_error_is_retried() -> None:
# A mid-stream quota/billing error arrives with no HTTP status code.
assert _retries(ModelRetryNormalizedError(status_code=None)) is True
def test_statusless_abort_is_not_retried() -> None:
# A user/client cancellation must never be retried.
assert _retries(ModelRetryNormalizedError(status_code=None, is_abort=True)) is False
def test_client_error_is_not_retried() -> None:
# A definitive 4xx client error (bad request/auth) is not recoverable.
assert _retries(ModelRetryNormalizedError(status_code=400)) is False
def test_rate_limit_and_server_errors_are_retried() -> None:
for status in (429, 500, 502, 503, 504):
assert _retries(ModelRetryNormalizedError(status_code=status)) is True
def test_timeout_error_is_retried() -> None:
# A stalled model stream trips the per-request read/inactivity timeout, which
# the SDK normalizes as a timeout. DEFAULT_MODEL_RETRY must retry it so a hung
# turn recovers instead of silently wedging the agent.
assert _retries(ModelRetryNormalizedError(is_timeout=True)) is True
assert _retries(ModelRetryNormalizedError(is_network_error=True)) is True
def test_policy_helper_matches_statusless_only() -> None:
assert _retry_statusless_provider_errors(_context(ModelRetryNormalizedError())) is True
assert (
_retry_statusless_provider_errors(_context(ModelRetryNormalizedError(status_code=400)))
is False
)
assert (
_retry_statusless_provider_errors(
_context(ModelRetryNormalizedError(status_code=None, is_abort=True))
)
is False
)
+23 -1
View File
@@ -3,8 +3,13 @@
from __future__ import annotations
import pytest
from agents.model_settings import ModelSettings
from strix.config.models import RECOMMENDED_MODEL_NAMES, is_recommended_or_frontier_model
from strix.config.models import (
RECOMMENDED_MODEL_NAMES,
is_recommended_or_frontier_model,
request_timeout_extra_args,
)
@pytest.mark.parametrize("model_name", RECOMMENDED_MODEL_NAMES)
@@ -12,6 +17,23 @@ def test_recommended_models_are_accepted(model_name: str) -> None:
assert is_recommended_or_frontier_model(model_name)
def test_request_timeout_extra_args_positive() -> None:
assert request_timeout_extra_args(300) == {"timeout": 300}
assert request_timeout_extra_args(10) == {"timeout": 10}
def test_request_timeout_extra_args_survives_model_settings_json_dump() -> None:
"""The Chat Completions and LiteLLM paths pydantic-serialize ModelSettings for
their tracing span; a non-JSON-serializable timeout fails every turn there."""
settings = ModelSettings(extra_args=request_timeout_extra_args(300))
assert settings.to_json_dict()["extra_args"] == {"timeout": 300}
@pytest.mark.parametrize("value", [None, 0, -1])
def test_request_timeout_extra_args_disabled(value: float | None) -> None:
assert request_timeout_extra_args(value) is None
def test_recommended_models_are_matched_case_insensitively() -> None:
assert is_recommended_or_frontier_model("Vertex_AI/Gemini-3-Pro-Preview")
+17 -182
View File
@@ -1,20 +1,19 @@
"""Tests for the shared Caido client lifecycle and proxy error handling.
"""Tests for the shared Caido client lifecycle and proxy call serialization.
Covers the concurrency/reconnect guarantees of ``caido_api.call_with_client``
(the sandbox-imported path) and ``caido_api.SharedCaidoClient`` (the host-side
holder), plus the actionable HTTPQL errors in ``proxy.tools``.
Covers the caching + serialization guarantees of ``caido_api.call_with_client``
(the sandbox-imported path) and ``proxy.tools._call`` (the host-side path). The
Caido GraphQL transport is not concurrency-safe, so both paths must run one
call at a time against the shared client.
"""
from __future__ import annotations
import asyncio
import json
from typing import TYPE_CHECKING, Any, cast
import pytest
from strix.tools.proxy import caido_api, tools
from strix.tools.proxy.caido_api import SharedCaidoClient
if TYPE_CHECKING:
@@ -91,83 +90,15 @@ async def test_failed_init_does_not_poison_cache(monkeypatch: pytest.MonkeyPatch
assert "default" not in caido_api._CLIENT_CACHE
async def test_call_with_client_reconnects_and_closes_dead_transport(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dead = _FakeClient("dead")
fresh = _FakeClient("fresh")
caido_api._CLIENT_CACHE["default"] = cast("Any", dead)
new_calls = {"n": 0}
async def _new() -> Any:
new_calls["n"] += 1
return fresh
monkeypatch.setattr(caido_api, "_new_client", _new)
attempts: list[Any] = []
async def fn(client: Any) -> str:
attempts.append(client)
if len(attempts) == 1:
raise RuntimeError("Transport is already connected")
return "ok"
assert await caido_api.call_with_client(fn) == "ok"
assert attempts == [dead, fresh]
assert new_calls["n"] == 1
assert caido_api._CLIENT_CACHE["default"] is fresh
assert dead.closed is True # stale transport is not leaked
async def test_call_with_client_non_idempotent_rebuilds_but_reraises(
monkeypatch: pytest.MonkeyPatch,
) -> None:
dead = _FakeClient("dead")
fresh = _FakeClient("fresh")
caido_api._CLIENT_CACHE["default"] = cast("Any", dead)
async def _new() -> Any:
return fresh
monkeypatch.setattr(caido_api, "_new_client", _new)
calls = {"n": 0}
async def fn(_client: Any) -> str:
calls["n"] += 1
raise RuntimeError("Server disconnected")
# A mutation must not be auto-retried (it may already have applied), but the
# dead client is still healed so later calls succeed.
with pytest.raises(RuntimeError, match="Server disconnected"):
await caido_api.call_with_client(fn, idempotent=False)
assert calls["n"] == 1
assert caido_api._CLIENT_CACHE["default"] is fresh
assert dead.closed is True
async def test_call_with_client_does_not_retry_application_errors(
monkeypatch: pytest.MonkeyPatch,
) -> None:
async def test_call_with_client_propagates_errors() -> None:
cached = _FakeClient("cached")
caido_api._CLIENT_CACHE["default"] = cast("Any", cached)
async def _new() -> Any:
raise AssertionError("deterministic errors must not trigger a reconnect")
monkeypatch.setattr(caido_api, "_new_client", _new)
calls = {"n": 0}
async def fn(_client: Any) -> str:
calls["n"] += 1
raise ValueError("Invalid HTTPQL filter")
with pytest.raises(ValueError, match="Invalid HTTPQL"):
await caido_api.call_with_client(fn)
assert calls["n"] == 1
assert caido_api._CLIENT_CACHE["default"] is cached
@@ -177,7 +108,7 @@ async def test_call_with_client_serializes_concurrent_calls(
caido_api._CLIENT_CACHE["default"] = cast("Any", _FakeClient("shared"))
async def _new() -> Any:
raise AssertionError("no reconnect expected")
raise AssertionError("no new client expected")
monkeypatch.setattr(caido_api, "_new_client", _new)
@@ -194,61 +125,8 @@ async def test_call_with_client_serializes_concurrent_calls(
assert state["max"] == 1
async def test_shared_client_reconnects_and_closes_dead_transport() -> None:
dead = _FakeClient("dead")
fresh = _FakeClient("fresh")
async def _reconnect() -> Any:
return fresh
holder = SharedCaidoClient(cast("Any", dead), _reconnect)
attempts: list[Any] = []
async def fn(client: Any) -> str:
attempts.append(client)
if len(attempts) == 1:
raise RuntimeError("Connector is closed")
return "ok"
assert await holder.call(fn) == "ok"
assert attempts == [dead, fresh]
assert dead.closed is True
async def test_shared_client_non_idempotent_rebuilds_but_reraises() -> None:
dead = _FakeClient("dead")
fresh = _FakeClient("fresh")
async def _reconnect() -> Any:
return fresh
holder = SharedCaidoClient(cast("Any", dead), _reconnect)
calls = {"n": 0}
async def fn(_client: Any) -> str:
calls["n"] += 1
raise RuntimeError("Server disconnected")
with pytest.raises(RuntimeError, match="Server disconnected"):
await holder.call(fn, idempotent=False)
assert calls["n"] == 1
assert dead.closed is True
# The healthy client remains for the next call.
assert await holder.call(lambda _c: _ok()) == "ok"
async def _ok() -> str:
return "ok"
async def test_shared_client_serializes_concurrent_calls() -> None:
async def _reconnect() -> Any:
raise AssertionError("no reconnect expected")
holder = SharedCaidoClient(cast("Any", _FakeClient("shared")), _reconnect)
async def test_host_call_serializes_concurrent_calls() -> None:
client = _FakeClient("host")
state = {"active": 0, "max": 0}
async def fn(_client: Any) -> str:
@@ -258,64 +136,21 @@ async def test_shared_client_serializes_concurrent_calls() -> None:
state["active"] -= 1
return "ok"
await asyncio.gather(*(holder.call(fn) for _ in range(6)))
await asyncio.gather(*(tools._call(cast("Any", client), fn) for _ in range(6)))
assert state["max"] == 1
async def test_shared_client_passes_through_application_errors() -> None:
async def _reconnect() -> Any:
raise AssertionError("deterministic errors must not trigger a reconnect")
holder = SharedCaidoClient(cast("Any", _FakeClient("c")), _reconnect)
async def fn(_client: Any) -> str:
raise ValueError("Invalid HTTPQL filter")
with pytest.raises(ValueError, match="Invalid HTTPQL"):
await holder.call(fn)
def test_is_connection_error_matches_markers_and_causes() -> None:
assert caido_api._is_connection_error(RuntimeError("Transport is already connected"))
assert caido_api._is_connection_error(RuntimeError("Connector is closed"))
assert caido_api._is_connection_error(RuntimeError("Server disconnected"))
assert not caido_api._is_connection_error(ValueError("Invalid HTTPQL filter"))
nested = RuntimeError("wrapper")
nested.__cause__ = RuntimeError("connection reset by peer")
assert caido_api._is_connection_error(nested)
class _Ctx:
def __init__(self, context: Any) -> None:
self.context = context
def test_ctx_proxy_returns_holder_when_present() -> None:
async def _reconnect() -> Any:
raise AssertionError("unused")
holder = SharedCaidoClient(cast("Any", _FakeClient("c")), _reconnect)
got = tools._ctx_proxy(cast("Any", _Ctx({"caido_client": holder})))
assert got is holder
def test_ctx_client_returns_client_when_present() -> None:
client = _FakeClient("host")
got = tools._ctx_client(cast("Any", _Ctx({"caido_client": client})))
assert got is client
def test_ctx_proxy_returns_none_without_holder() -> None:
assert tools._ctx_proxy(cast("Any", _Ctx({}))) is None
assert tools._ctx_proxy(cast("Any", _Ctx(None))) is None
assert tools._ctx_proxy(cast("Any", _Ctx({"caido_client": object()}))) is None
def test_is_httpql_error_detection() -> None:
assert tools._is_httpql_error(RuntimeError("HTTPQL parse error at column 4"))
assert tools._is_httpql_error(RuntimeError("failed to parse filter"))
assert not tools._is_httpql_error(RuntimeError("Transport is already connected"))
def test_httpql_error_preserves_message_and_query() -> None:
exc = RuntimeError("HTTPQL parse error: unexpected token at column 12")
payload = json.loads(tools._httpql_error(exc, 'resp.code.eq:"200"'))
assert payload["success"] is False
assert "unexpected token at column 12" in payload["error"]
assert payload["httpql_filter"] == 'resp.code.eq:"200"'
assert "AND / OR" in payload["hint"]
def test_ctx_client_returns_none_without_client() -> None:
assert tools._ctx_client(cast("Any", _Ctx({}))) is None
assert tools._ctx_client(cast("Any", _Ctx(None))) is None
+1
View File
@@ -37,6 +37,7 @@ async def test_persistent_rate_limit_stops_gracefully(
model="openai/gpt-4o",
reasoning_effort="high",
force_required_tool_choice=False,
timeout=300,
),
runtime=types.SimpleNamespace(max_context_images=3),
)
+2 -2
View File
@@ -45,6 +45,7 @@ def _patch_engine_scaffold(
model="openai/gpt-4o",
reasoning_effort="high",
force_required_tool_choice=False,
timeout=300,
)
)
monkeypatch.setattr(runner, "load_settings", lambda: settings)
@@ -124,8 +125,7 @@ async def test_root_prompt_options_flow_into_root_agent(
assert "https://example.com" in instructions_override
assert "CUSTOM SCAN PROMPT" in instructions_override
assert (
"cannot expand, replace, or weaken authorized target constraints"
in instructions_override
"cannot expand, replace, or weaken authorized target constraints" in instructions_override
)
assert kwargs["system_prompt_context"] == {
**scope_context,
+172
View File
@@ -0,0 +1,172 @@
import hashlib
import io
import json
import platform
import time
from pathlib import Path
import pytest
from rich.console import Console
from strix.interface import update_check
@pytest.fixture(autouse=True)
def _isolated_cache(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(update_check, "_CACHE_PATH", tmp_path / "update-check.json")
monkeypatch.setattr(update_check, "_background_thread", None)
monkeypatch.delenv("STRIX_NO_UPDATE_CHECK", raising=False)
for key in ("CI", "GITHUB_ACTIONS", "GITLAB_CI", "JENKINS_URL", "BUILDKITE", "CIRCLECI"):
monkeypatch.delenv(key, raising=False)
@pytest.mark.parametrize(
("latest", "current", "expected"),
[
("1.2.0", "1.1.0", True),
("1.1.0", "1.1.0", False),
("1.0.9", "1.1.0", False),
("2.0.0", "1.99.99", True),
("1.10.0", "1.9.0", True),
("v1.2.0", "1.1.0", True),
("not-a-version", "1.1.0", False),
("1.2.0", "unknown", False),
],
)
def test_is_newer(latest: str, current: str, expected: bool) -> None:
assert update_check._is_newer(latest, current) is expected
def test_get_available_update_from_cache(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "9.9.9", "checked_at": time.time()})
)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
assert update_check.get_available_update() == "9.9.9"
def test_get_available_update_up_to_date(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "1.0.0", "checked_at": time.time()})
)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
assert update_check.get_available_update() is None
def test_get_available_update_disabled_by_env(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "9.9.9", "checked_at": time.time()})
)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
monkeypatch.setenv("STRIX_NO_UPDATE_CHECK", "1")
assert update_check.get_available_update() is None
def test_get_available_update_disabled_in_ci(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "9.9.9", "checked_at": time.time()})
)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
monkeypatch.setenv("CI", "true")
assert update_check.get_available_update() is None
def test_get_available_update_corrupt_cache(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text("{not json")
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
assert update_check.get_available_update() is None
def test_background_check_skipped_when_fresh(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "1.0.0", "checked_at": time.time()})
)
called = False
def fake_refresh() -> None:
nonlocal called
called = True
monkeypatch.setattr(update_check, "_refresh_cache", fake_refresh)
update_check.start_background_check()
assert update_check._background_thread is None
assert called is False
def test_background_check_runs_when_stale(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "1.0.0", "checked_at": time.time() - 2 * 24 * 60 * 60})
)
monkeypatch.setattr(update_check, "_fetch_latest_version", lambda: "1.2.3")
update_check.start_background_check()
assert update_check._background_thread is not None
update_check._background_thread.join(timeout=5)
cache = json.loads(update_check._CACHE_PATH.read_text())
assert cache["latest_version"] == "1.2.3"
def test_skipped_version_suppresses_update(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps({"latest_version": "9.9.9", "checked_at": time.time()})
)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
update_check.skip_version("9.9.9")
assert update_check.get_available_update() is None
assert update_check.get_available_update(respect_skip=False) == "9.9.9"
def test_newer_release_overrides_skipped_version(monkeypatch: pytest.MonkeyPatch) -> None:
update_check._CACHE_PATH.write_text(
json.dumps(
{"latest_version": "9.9.10", "checked_at": time.time(), "skipped_version": "9.9.9"}
)
)
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
assert update_check.get_available_update() == "9.9.10"
def test_write_cache_preserves_existing_fields() -> None:
update_check.skip_version("9.9.9")
update_check._write_cache(latest_version="1.2.3", checked_at=123.0)
cache = json.loads(update_check._CACHE_PATH.read_text())
assert cache == {"latest_version": "1.2.3", "checked_at": 123.0, "skipped_version": "9.9.9"}
def test_get_upgrade_command_all_methods() -> None:
assert update_check.get_upgrade_command("binary") == "strix --update"
assert update_check.get_upgrade_command("pipx") == "pipx upgrade strix-agent"
assert update_check.get_upgrade_command("uv") == "uv tool upgrade strix-agent"
assert update_check.get_upgrade_command("pip") == "pip install --upgrade strix-agent"
def test_self_update_non_binary_prints_command(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(update_check, "is_binary_install", lambda: False)
buffer = io.StringIO()
assert update_check.self_update(Console(file=buffer)) is False
assert "upgrade" in buffer.getvalue()
def test_self_update_already_latest(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(update_check, "is_binary_install", lambda: True)
monkeypatch.setattr(update_check, "_fetch_latest_version", lambda: "1.0.0")
monkeypatch.setattr(update_check, "get_version", lambda: "1.0.0")
assert update_check.self_update() is True
def test_sha256_file(tmp_path: Path) -> None:
path = tmp_path / "blob"
path.write_bytes(b"strix")
assert update_check._sha256_file(path) == hashlib.sha256(b"strix").hexdigest()
def test_release_target(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(platform, "system", lambda: "Linux")
monkeypatch.setattr(platform, "machine", lambda: "x86_64")
assert update_check._release_target() == "linux-x86_64"
monkeypatch.setattr(platform, "system", lambda: "Darwin")
monkeypatch.setattr(platform, "machine", lambda: "arm64")
assert update_check._release_target() == "macos-arm64"
monkeypatch.setattr(platform, "machine", lambda: "riscv64")
assert update_check._release_target() is None