mirror of
https://github.com/usestrix/strix.git
synced 2026-08-19 18:13:34 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
aa7a097037 | ||
|
|
df80fad960 | ||
|
|
a0bddbc13b | ||
|
|
b5e82edb83 | ||
|
|
7d5a67d234 | ||
|
|
88ad3e4472 | ||
|
|
cf7689e927 | ||
|
|
3bb95ab43d | ||
|
|
9aa151c687 | ||
|
|
b9c2592b53 | ||
|
|
f54ecb74f9 | ||
|
|
96ca7e544d |
@@ -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).
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user