mirror of
https://hubproxy.babadafafafafa.cn/https://github.com/usestrix/strix.git
synced 2026-09-20 16:13:44 +08:00
Compare commits
22 Commits
mcp-generi
...
docs/skill
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
38186a7e50 | ||
|
|
8fdf6a5c09 | ||
|
|
46cf2f52f3 | ||
|
|
3de9471431 | ||
|
|
a071022182 | ||
|
|
d26b1ab0de | ||
|
|
de730119f0 | ||
|
|
608ef4a37b | ||
|
|
f901d2a8bf | ||
|
|
944274e12f | ||
|
|
eeca404716 | ||
|
|
1df67c52e2 | ||
|
|
3c767cdd47 | ||
|
|
0a6e8b01bf | ||
|
|
1f3f9b31ae | ||
|
|
cf179d564e | ||
|
|
583af23d9a | ||
|
|
717ffc8f4c | ||
|
|
cbb0f57058 | ||
|
|
8b655de615 | ||
|
|
7d8d71beea | ||
|
|
a5856108a7 |
22
AGENTS.md
22
AGENTS.md
@@ -38,13 +38,25 @@ Target-specific workflows built on the same engine:
|
||||
|
||||
- **Managed cloud (app.strix.ai):** no Docker, no LLM key, no local install; adds team dashboards, scheduling, PR reviews, and downloadable PDF/DOCX reports (Enterprise plan). Best in sandboxed/CI environments and for teams. Use it when local infra isn't available.
|
||||
```bash
|
||||
# token from Settings → API Access; register the target as an asset, then:
|
||||
curl -sS https://app.strix.ai/api/v1/scans -H "Authorization: Bearer $STRIX_API_TOKEN" \
|
||||
-H "Content-Type: application/json" -d '{"engagement_type":"live_test","domain_ids":["<uuid>"]}'
|
||||
strix cloud login --scopes scans:read scans:write uploads:write billing:read
|
||||
strix cloud domains add --domain example.com --asset-type web_app
|
||||
strix cloud scans start --engagement-type live_test --domain-ids <uuid> --wait
|
||||
strix cloud scans start --source . --dry-run --show-files --json # review + capture source.archive_sha256
|
||||
SOURCE_SHA256="<reviewed source.archive_sha256>"
|
||||
strix cloud scans start --source . --approve-sha256 "$SOURCE_SHA256" --wait
|
||||
strix cloud vulns list --severity critical
|
||||
strix cloud billing topup --credits 20 --yes # explicit approval after exit code 5
|
||||
```
|
||||
- API docs: https://docs.app.strix.ai (OpenAPI: https://docs.app.strix.ai/openapi.json).
|
||||
- Account setup runs from the CLI too: `strix cloud workspaces list|create|use` (`workspace` is an alias and `use` accepts a displayed number, name, or ID), `strix cloud session scopes|scopes set`, `strix cloud org members invite`, `strix cloud billing subscribe --plan strix_cloud`, `strix cloud billing portal`, `strix cloud integrations install github`, and `strix cloud domains verify <id>`. Workspace switching preserves the server-side profile and can never widen past the login ceiling; ordinary switches do not reprompt. The last four end at a person: the command prints a link or a DNS record for the user to open or add, and it never completes the payment, installation, or DNS change for them.
|
||||
- Every REST operation has a `strix cloud <resource> <verb>` command. Run `strix cloud` to list them. Output is JSON when stdout is not a terminal (or with `--json`), and there are no prompts without a TTY. Binary downloads are the exception: redirect raw bytes intentionally, or combine `--output FILE --json` for structured download metadata. Exit codes: `0` success, `1` error, `2` usage, `4` auth or plan limit, `5` payment required. `--token` or `STRIX_API_TOKEN` is a stateless override and never replaces stored auth; set `--workspace-id`/`STRIX_WORKSPACE_ID` for an override CLI session. `--data` adds extra request fields as JSON, and accepts `@file` or `-` for standard input.
|
||||
- Local source uploads require `uploads:write`. For an agent/CI handoff, review `scans start --source . --dry-run --show-files --json`, capture `source.archive_sha256`, then rerun with the same `--source`, `--exclude`, and `--include-*` selection flags plus `--approve-sha256 HASH`. A changed snapshot is rejected. `--yes` approves only the snapshot built in that invocation, so reserve it for a deliberate human or one-shot approval rather than a digest-bound two-step handoff.
|
||||
- Git ignores, hidden files, `.git`, symlinks, dependency/build output, secret-like filenames, and nested archives are excluded by default; `.strixignore` and `--exclude` narrow the manifest further (a trailing `/` excludes a directory subtree). Limits: 20,000 files, 25 MiB/file, 250 MiB expanded, 50 MiB compressed. Source-only infers `code_review`; source plus a domain infers `live_test`.
|
||||
- The temporary local archive is always removed. A staged upload is deleted after a definitive rejection, but retained when a network error, `5xx`, malformed success response, or interruption leaves the scan launch ambiguous. JSON reports its `upload_id` with `launch_outcome_unknown: true`, or with `cleanup_unknown: true` when automatic deletion cannot be confirmed. Check `scans list` before retrying; if no scan is linked, run `uploads delete UPLOAD_ID`.
|
||||
- Non-Enterprise scans consume the scope estimate (a default-tier source-only review currently starts at 60 credits); Enterprise scans are plan-included. A rejected launch does not consume credits.
|
||||
- Human output is compact and numbered; non-TTY output and `--json` retain full records. Enable tab completion with `source <(strix completions zsh)` (or `bash`), or `strix completions fish | source`.
|
||||
- The REST API works directly too: https://docs.app.strix.ai (OpenAPI: https://docs.app.strix.ai/openapi.json).
|
||||
|
||||
- CLI docs index for LLMs: https://docs.strix.ai/llms.txt (full: https://docs.strix.ai/llms-full.txt).
|
||||
- CLI docs index for LLMs: https://docs.strix.ai/llms.txt (full: https://docs.strix.ai/llms-full.txt). Managed API docs for LLMs: https://docs.app.strix.ai/llms.txt.
|
||||
- Only scan targets the user is authorized to test.
|
||||
|
||||
## Contributing to this repo
|
||||
|
||||
102
README.md
102
README.md
@@ -320,6 +320,108 @@ strix auth status # show the active sign-in
|
||||
strix auth logout # forget the sign-in
|
||||
```
|
||||
|
||||
#### Use the managed platform: `strix cloud`
|
||||
|
||||
The `strix cloud` commands drive the managed platform ([app.strix.ai](https://app.strix.ai)) from the terminal. Sign in once with the device flow. The sign-in creates your account and workspace on first use and stores a personal API token in `~/.strix/platform-auth.json`:
|
||||
|
||||
```bash
|
||||
strix cloud login # browser approval, then workspace + scope profile
|
||||
strix cloud login --workspace "My Team" # select a workspace by name or ID
|
||||
strix cloud whoami # fast local account/workspace status
|
||||
strix cloud session # verify remote session + consent ceiling
|
||||
strix cloud logout # revoke remotely, then remove locally
|
||||
```
|
||||
|
||||
The default **Recommended** scope preset supports normal scan work, local source uploads,
|
||||
workspace switching, and user-approved credit top-ups. It excludes credential creation;
|
||||
request `tokens:write` explicitly (or choose Full) when needed. For strict least privilege, pass an explicit list such as
|
||||
`--scopes scans:read scans:write uploads:write billing:read`. Named automation
|
||||
profiles are also available with `--scope-profile minimal|recommended|full`.
|
||||
|
||||
Every operation of the [REST API](https://docs.app.strix.ai) has a matching command in the form `strix cloud <resource> <verb>`:
|
||||
|
||||
```bash
|
||||
strix cloud # list all resources
|
||||
strix cloud scans # run the safe default (`scans list`)
|
||||
strix cloud scans help # list the verbs of a resource
|
||||
strix cloud domains add --domain example.com --asset-type web_app
|
||||
strix cloud scans start --engagement-type live_test --domain-ids <uuid> --wait
|
||||
strix cloud scans start --source . --dry-run --show-files --json # review + capture source.archive_sha256
|
||||
SOURCE_SHA256="<reviewed source.archive_sha256>"
|
||||
strix cloud scans start --source . --approve-sha256 "$SOURCE_SHA256" --wait
|
||||
strix cloud vulns list --severity critical
|
||||
strix cloud credits # credit balance
|
||||
strix cloud billing topup --credits 20 --yes # explicitly approve agent payment after HTTP 402
|
||||
```
|
||||
|
||||
Workspaces and account setup also work from the terminal:
|
||||
|
||||
```bash
|
||||
strix cloud workspaces list # numbered list; `workspace` is also accepted
|
||||
strix cloud workspaces create --name "My Team" # admin + organizations:write
|
||||
strix cloud workspaces use 2 # switch by list number, exact name, or ID
|
||||
strix cloud session scopes # granted scopes + login ceiling
|
||||
strix cloud session scopes set minimal # narrow without another browser sign-in
|
||||
strix cloud billing subscribe --plan strix_cloud # opens the hosted checkout page
|
||||
strix cloud billing portal # opens the billing portal
|
||||
strix cloud integrations install github # opens the app installation page
|
||||
strix cloud domains verify <domain-id> # prints the DNS record to add
|
||||
```
|
||||
|
||||
The last four commands end at a person. Strix creates the link, opens the browser for an interactive terminal, and always prints the URL. The user enters the card, approves the installation, or adds the DNS record. Pass `--no-browser` to print the URL only.
|
||||
|
||||
The commands work for humans and agents: terminal output favors names, branches, lifecycle states, and numbered selectors, while redirected output (or `--json`) preserves complete machine-readable records and IDs. Human lists retain the selectors needed by follow-up commands but omit internal organization/user IDs; a selector too long for the compact table is repeated losslessly in a copyable block. Paginated lists print the next `--page` or `--offset`, and detail views preserve useful prose within a safe terminal bound; use `--json` for the complete record. Token lists distinguish API keys from named CLI device sessions. Binary downloads are the exception: intentionally redirect their raw bytes, or use `--output FILE --json` to write the file and receive structured download metadata. There are no prompts when stdin is not a terminal. Exit codes: `0` success, `1` error, `2` invalid usage, `4` authentication or plan limit, `5` payment required. `--token` and `STRIX_API_TOKEN` are stateless per-command overrides and never replace the stored sign-in; pair a CLI-session override with `--workspace-id` or `STRIX_WORKSPACE_ID`.
|
||||
|
||||
A browser sign-in creates one reusable credential per CLI installation. Logging in again on the
|
||||
same installation replaces its secret instead of accumulating keys. Workspace switches keep that
|
||||
credential and expiry, preserve the server-side scope preference, cap access by the target role,
|
||||
and can never exceed the login consent ceiling. Each process pins its starting workspace, so a
|
||||
concurrent switch fails safely instead of sending a stale command to another organization.
|
||||
`strix cloud logout` revokes the server session before deleting the local token; use
|
||||
`--local-only` only when you deliberately cannot reach the server.
|
||||
|
||||
Write commands take request fields as flags, and every write command also accepts one JSON object with `--data`:
|
||||
|
||||
```bash
|
||||
strix cloud scans start --data '{"engagement_type":"code_review"}' # literal JSON
|
||||
strix cloud scans start --data @request.json # read a file
|
||||
cat request.json | strix cloud scans start --data - # read standard input
|
||||
```
|
||||
|
||||
For an agent or CI local-source scan, run `--dry-run --show-files --json`, review the manifest,
|
||||
and capture `source.archive_sha256`. Rerun with the same `--source`, every `--exclude`, and any
|
||||
`--include-*` selection flags, replacing `--dry-run` with `--approve-sha256 HASH`; Strix
|
||||
rebuilds the archive and refuses to upload it if the digest changed. `--yes` instead approves
|
||||
only the snapshot built in that one invocation. It is suitable for a deliberate human or
|
||||
one-shot approval, not as a digest-bound two-step agent/CI handoff.
|
||||
|
||||
The safe default honors `.gitignore` and `.strixignore` and excludes hidden paths, secret-like
|
||||
files, VCS metadata, dependencies/build output, symlinks, and nested archives. Opt in
|
||||
separately with `--include-hidden`, `--include-sensitive`, or `--include-archives`. The client
|
||||
caps a bundle at 20,000 files, 25 MiB per file, 250 MiB expanded, and 50 MiB compressed, and
|
||||
the service independently validates the archive. Source alone infers a code review; adding a
|
||||
domain infers a live test. You can always pass `--engagement-type` explicitly.
|
||||
|
||||
Strix removes the temporary local archive after every invocation. It deletes a staged remote
|
||||
upload after a definitive scan rejection. If a network error, `5xx` response, malformed
|
||||
success response, or interruption makes the launch outcome ambiguous, it retains the upload and reports its `upload_id` with
|
||||
`launch_outcome_unknown: true`; if automatic deletion cannot be confirmed, it reports the ID
|
||||
with `cleanup_unknown: true`. Check `strix cloud scans list` before retrying. If no scan is
|
||||
linked to the retained upload, delete it with `strix cloud uploads delete UPLOAD_ID`.
|
||||
|
||||
Non-Enterprise scans consume the deterministic estimate shown for their scope (a source-only
|
||||
code review at the default `ultra` tier currently starts at 60 credits). Enterprise scans are
|
||||
plan-included and do not consume the credit wallet. Report downloads need Enterprise,
|
||||
schedules need Pro, and billing writes need an admin token. Plan blocks exit `4`; an
|
||||
insufficient credit wallet exits `5` without creating or charging a scan.
|
||||
|
||||
Enable native tab completion once per shell session:
|
||||
|
||||
```bash
|
||||
source <(strix completions zsh) # use bash instead of zsh when appropriate
|
||||
strix completions fish | source
|
||||
```
|
||||
|
||||
#### Connect your own MCP servers
|
||||
|
||||
Strix can connect to Model Context Protocol (MCP) servers you list and expose their tools to the agent during a run. Create `~/.strix/mcp-servers.json` with a JSON list of servers. Each entry is either a local `stdio` server that Strix launches as a subprocess, or a remote `http` server:
|
||||
|
||||
@@ -35,6 +35,25 @@ Skip the setup. Run Strix in the cloud at [app.strix.ai](https://app.strix.ai).
|
||||
2. Connect your repository or enter a target URL
|
||||
3. Launch your first scan
|
||||
|
||||
## Scan Local Source
|
||||
|
||||
Send a local working tree to the managed white-box scanner without connecting a source-control provider:
|
||||
|
||||
```bash
|
||||
# Review the exact file manifest and capture source.archive_sha256. Nothing is uploaded.
|
||||
strix cloud scans start --source . --dry-run --show-files --json
|
||||
SOURCE_SHA256="<reviewed source.archive_sha256>"
|
||||
|
||||
# Repeat the same source-selection flags and approve that exact snapshot.
|
||||
strix cloud scans start --source . --approve-sha256 "$SOURCE_SHA256" --wait
|
||||
```
|
||||
|
||||
In a Git repository, Strix includes tracked files and untracked files that are not ignored. Hidden files, `.git`, symlinks, dependencies and build output, secret-like filenames, and nested archives are excluded by default. Use `.strixignore` or repeat `--exclude GLOB` for project-specific exclusions. `--include-hidden`, `--include-sensitive`, and `--include-archives` are explicit opt-ins.
|
||||
|
||||
The CLI limits individual files, total expanded bytes, archive bytes, and file count. For an agent or CI handoff, repeat the same `--source`, `--exclude`, and `--include-*` flags with `--approve-sha256`; Strix refuses the upload if the rebuilt archive differs from the reviewed digest. `--yes` is a one-invocation approval for the snapshot built at that moment, not a digest-bound two-step approval.
|
||||
|
||||
The temporary local archive is always removed. After a definitive launch rejection, Strix also deletes the staged remote upload. If a network error, server error, or interruption makes the launch outcome ambiguous, it retains the upload and reports its ID; check `strix cloud scans list` before retrying, then delete an unlinked upload with `strix cloud uploads delete UPLOAD_ID`.
|
||||
|
||||
<Card title="Try Strix Cloud" icon="rocket" href="https://app.strix.ai">
|
||||
Run your first pentest in minutes.
|
||||
</Card>
|
||||
|
||||
@@ -36,13 +36,14 @@ npx skills use usestrix/strix@penetration-testing-with-strix | claude
|
||||
Both use the same engine and produce the same validated findings and SARIF, so agents can pick per situation or combine them:
|
||||
|
||||
- **Open-source CLI (self-hosted)** — runs locally in a Docker sandbox with your own LLM key. Free, fully local, air-gap capable. Best for local dev loops and full control.
|
||||
- **Managed cloud** — runs on Strix's infrastructure via the [app.strix.ai REST API](https://docs.app.strix.ai). No Docker, no LLM key, no local install; adds team dashboards, scheduling, PR reviews, and downloadable PDF/DOCX reports (Enterprise plan). Best in sandboxed/CI environments and for teams. Create an API token under **Settings → API Access**; the `managed-pentesting-with-strix` skill has the full flow.
|
||||
- **Managed cloud** — runs on Strix's infrastructure. Drive it with the `strix cloud` CLI (every REST operation has a `strix cloud <resource> <verb>` command) or the [app.strix.ai REST API](https://docs.app.strix.ai) directly. No Docker, no LLM key; adds team dashboards, scheduling, PR reviews, and downloadable PDF/DOCX reports (Enterprise plan). Best in sandboxed/CI environments and for teams. Sign in with `strix cloud login` (browser device sign-in, account created on first use) or create a token in the dashboard under **Settings → API Access**. The `managed-pentesting-with-strix` skill has the full flow.
|
||||
|
||||
## Agent-Friendly Interfaces
|
||||
|
||||
Everything an agent needs is machine-readable:
|
||||
|
||||
- **Headless CLI** — `strix -n` runs without the TUI and exits with `0` (clean), `1` (error), or `2` (vulnerabilities found).
|
||||
- **Cloud CLI** — `strix cloud` prints JSON when stdout is not a terminal (or with `--json`), never prompts without a TTY, and exits with `0` (success), `1` (error), `2` (usage), `4` (authentication required), or `5` (payment required). Credit top-ups pay the Stripe machine-payment challenge with an agent wallet (`strix cloud billing topup --credits N --yes`). Account setup also runs from the CLI: `strix cloud workspaces list|create|use`, `strix cloud org members invite`, `strix cloud billing subscribe`, `strix cloud billing portal`, and `strix cloud integrations install github`. The last three print a hosted link the user opens to finish the payment or approve the installation.
|
||||
- **REST API** — the managed platform exposes a documented [OpenAPI](https://docs.app.strix.ai/openapi.json) at `https://app.strix.ai/api/v1` (scans, vulnerabilities, assets, PR reviews, schedules, webhooks) with bearer tokens and scopes.
|
||||
- **Structured results** — every run writes `vulnerabilities.json`, `vulnerabilities.csv`, `findings.sarif` (SARIF 2.1.0), and per-finding Markdown under `strix_runs/<run-name>/`; the cloud exposes the same as JSON plus SARIF export.
|
||||
- **Budget controls** — `--max-budget` and `--max-turns` give agents hard cost/time caps.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "strix-agent"
|
||||
version = "1.5.3"
|
||||
version = "1.6.0"
|
||||
description = "Open-source AI Hackers for your apps"
|
||||
readme = "README.md"
|
||||
license = "Apache-2.0"
|
||||
@@ -43,6 +43,7 @@ dependencies = [
|
||||
"requests>=2.32.0",
|
||||
"cvss>=3.2",
|
||||
"caido-sdk-client>=0.2.0",
|
||||
"markdown-it-py>=3.0.0",
|
||||
"reportlab>=4.0",
|
||||
"pypdf>=5.0",
|
||||
# Cap <49: 49.x drops the universal2 macOS wheel (arm64-only), which breaks
|
||||
@@ -229,6 +230,7 @@ ignore = [
|
||||
# Test doubles use fixture tokens/passwords and match a callee signature whose
|
||||
# args they intentionally ignore.
|
||||
"tests/test_viewer_auth.py" = ["S105", "S106", "ARG001"]
|
||||
"tests/test_cloud_cli.py" = ["S105", "ARG001"]
|
||||
"tests/test_codex_auth.py" = ["S105", "S106", "SLF001"]
|
||||
# Hatchling loads the build hook by path, not as an importable package.
|
||||
"scripts/tui_sidecar_hook.py" = ["INP001"]
|
||||
@@ -255,6 +257,11 @@ ignore = [
|
||||
"strix/tools/notes/tools.py" = ["PLC0415", "TC002"]
|
||||
"strix/tools/finish/tool.py" = ["PLC0415", "TC002"]
|
||||
"strix/tools/reporting/tool.py" = ["PLC0415", "TC002"]
|
||||
# Lazy imports of strix.tools.mcp.client avoid a circular import (client imports
|
||||
# the session module at module load).
|
||||
"strix/tools/mcp/session.py" = ["PLC0415"]
|
||||
# call_mcp is a chain of guard clauses that each return an error string.
|
||||
"strix/tools/mcp/agent_tools.py" = ["PLR0911"]
|
||||
"strix/tools/**/*.py" = [
|
||||
"ARG001", # Unused function argument (tools may have unused args for interface consistency)
|
||||
]
|
||||
|
||||
@@ -11,7 +11,7 @@ metadata:
|
||||
|
||||
APIs fail differently from web UIs: there is no rendered surface to crawl, the interesting bugs are authorization-shaped rather than injection-shaped, and the same endpoint behaves differently per token. This workflow targets those specifics with Strix's autonomous agents, using the current [OWASP API Security Top 10 (2023)](https://owasp.org/API-Security/editions/2023/en/0x11-t10/) as the coverage checklist. For the web-app equivalent, the current edition is the OWASP Top 10:2025 — see **owasp-top-10-testing**.
|
||||
|
||||
Install, LLM setup, full CLI flags, and the managed-cloud path are in the **penetration-testing-with-strix** skill. Read it if `strix --version` fails or the target is not an API.
|
||||
Install, LLM setup, full CLI flags, and the managed-cloud path are in the **penetration-testing-with-strix** skill. Read it if `strix --version` fails or the target is not an API. For a run with no Docker and no LLM key, the same binary drives the managed platform: `strix cloud login`, then `strix cloud scans start ...` (details in **managed-pentesting-with-strix**).
|
||||
|
||||
## 1. Gather what the agents need
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ metadata:
|
||||
|
||||
Entry point for "make my application secure" requests, where the target is not yet a single URL or repo. The job here is to pick the right test per asset, run it, and produce one ranked plan — not to run everything at maximum depth.
|
||||
|
||||
Install, LLM setup, all CLI flags, and the managed-cloud path live in the **penetration-testing-with-strix** skill. Read it first if `strix --version` fails.
|
||||
Install, LLM setup, all CLI flags, and the managed-cloud path live in the **penetration-testing-with-strix** skill. Read it first if `strix --version` fails. For a run with no Docker and no LLM key, the same binary drives the managed platform: `strix cloud login`, then `strix cloud scans start ...` (details in **managed-pentesting-with-strix**).
|
||||
|
||||
Only test assets the user owns or is authorized to test. Confirm authorization before the first run, and prefer staging over production, because the agents send real exploit payloads and can change data.
|
||||
|
||||
|
||||
@@ -113,11 +113,11 @@ Gate the pipeline on the exit code (see the budget/fail-open caveat above — gi
|
||||
|
||||
# Option B — Managed platform (no runner infra)
|
||||
|
||||
No workflow file, no Docker, no LLM key. Two ways to use it:
|
||||
No workflow file, no Docker, no LLM key. Three ways to use it:
|
||||
|
||||
1. **PR-review app (zero code):** the user installs the Strix GitHub/GitLab/Bitbucket app and enables PR reviews for the repo in the app.strix.ai dashboard. Every PR is then reviewed automatically, with findings posted as PR comments. Nothing to add to the repo. This is the lowest-effort path — recommend it first when the user just wants PR gating.
|
||||
|
||||
2. **API-triggered from any pipeline:** if you want to trigger from an existing pipeline (or a system without the SCM app), call the API with a token that has `pr_reviews:write` (or `scans:write`). Store the token as a CI secret; ask the user to create it at **Settings → API Access**. Example GitHub Actions step:
|
||||
2. **CLI-triggered from any pipeline:** if you want to trigger from an existing pipeline (or a system without the SCM app), use the same `strix` binary with a token that has `pr_reviews:write`. Store the token as a CI secret and ask the user to create it at **Settings → API Access**. Read the repository's `provider` and `installation_id` once with `strix cloud repos list`. Example GitHub Actions step:
|
||||
|
||||
```yaml
|
||||
- name: Strix PR review (managed)
|
||||
@@ -125,12 +125,25 @@ No workflow file, no Docker, no LLM key. Two ways to use it:
|
||||
env:
|
||||
STRIX_API_TOKEN: ${{ secrets.STRIX_API_TOKEN }}
|
||||
run: |
|
||||
curl -sS --fail https://app.strix.ai/api/v1/pr-reviews/start \
|
||||
-H "Authorization: Bearer $STRIX_API_TOKEN" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d "{\"repository_full_name\":\"${{ github.repository }}\",\"pr_number\":${{ github.event.pull_request.number }}}"
|
||||
curl -sSL https://strix.ai/install | bash
|
||||
strix cloud pr-reviews start \
|
||||
--provider github \
|
||||
--installation-id "${{ vars.STRIX_INSTALLATION_ID }}" \
|
||||
--repository-full-name "${{ github.repository }}" \
|
||||
--pr-number "${{ github.event.pull_request.number }}"
|
||||
```
|
||||
|
||||
To gate the build on results, poll the PR review / scan status and fail on unresolved criticals/highs. Full endpoints (PR reviews, scans, SARIF export, schedules for scheduled deep scans) are in the **managed-pentesting-with-strix** skill.
|
||||
Output is JSON when stdout is not a terminal, and there are no prompts without a TTY. To gate the build on results, poll `strix cloud pr-reviews get <id> --json` and fail on unresolved criticals or highs. The raw REST endpoint (`POST /api/v1/pr-reviews/start`) works too when the pipeline cannot install the CLI.
|
||||
|
||||
3. **Source upload from a pipeline without an SCM app:** upload the checked-out tree as a cloud code review (`scans:write` and `uploads:write`). The two-step digest handoff keeps a human in control of what leaves the runner:
|
||||
|
||||
```bash
|
||||
strix cloud scans start --source . --dry-run --show-files --json # review, capture source.archive_sha256
|
||||
strix cloud scans start --source . --approve-sha256 "$SOURCE_SHA256" --wait
|
||||
```
|
||||
|
||||
Exit codes: `0` success, `4` auth or plan limit, `5` payment required. Non-Enterprise scans consume credits.
|
||||
|
||||
Full CLI coverage (PR reviews, scans, SARIF export, schedules) is in the **managed-pentesting-with-strix** skill.
|
||||
|
||||
Recommend Option B for most teams (no maintenance, central dashboard); use Option A when scans must stay entirely within your own infrastructure.
|
||||
|
||||
@@ -11,7 +11,7 @@ metadata:
|
||||
|
||||
White-box security review with Strix: the agents read the source to build a model of routes, sinks, and authorization checks, then attempt real exploitation. Findings come with a proof-of-concept, so the output is a short list of proven issues rather than the hundreds of "potential" hits a pattern-matching scanner produces.
|
||||
|
||||
Install, LLM setup, all flags, and the managed-cloud path are in the **penetration-testing-with-strix** skill.
|
||||
Install, LLM setup, all flags, and the managed-cloud path are in the **penetration-testing-with-strix** skill. For a run with no Docker and no LLM key, the same binary drives the managed platform: `strix cloud login`, then `strix cloud scans start ...` (details in **managed-pentesting-with-strix**).
|
||||
|
||||
## Run it
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ Get the findings from wherever the scan ran:
|
||||
- **OSS CLI** — artifacts in `strix_runs/<run-name>/`:
|
||||
- `vulnerabilities/*.md` — one finding per file: description, severity, PoC steps or script, affected code locations, remediation guidance.
|
||||
- `vulnerabilities.json` — the same findings as JSON (ids, severity, CWE/CVE, `code_locations` with `fix_before`/`fix_after` suggestions when available).
|
||||
- **Cloud (app.strix.ai)** — fetch the scan's `vulnerabilities[]` via `GET /api/v1/scans/{scanId}` (or `GET /api/v1/vulnerabilities` org-wide). Each carries `severity, cwe, endpoint, method, impact, technical_analysis, poc_description, poc_script_code` and, for code findings, `code_file`/`code_diff`/`code_before`/`code_after`. See the **managed-pentesting-with-strix** skill for auth.
|
||||
- **Cloud (app.strix.ai)** — pull findings with the CLI: `strix cloud vulns list --scan-id <scan-id> --json` (or `strix cloud scans get <scan-id> --json | jq '.vulnerabilities'`, or `strix cloud vulns list --severity critical` org-wide). Each finding carries `severity, cwe, endpoint, method, impact, technical_analysis, poc_description, poc_script_code` and, for code findings, `code_file`/`code_diff`/`code_before`/`code_after`. After a fix is verified, mark it with `strix cloud vulns update <id> --status fixed`. See the **managed-pentesting-with-strix** skill for `strix cloud login` and scopes.
|
||||
|
||||
Order work by severity: critical → high → medium → low. Every Strix finding was validated with a working proof-of-concept, so do not dismiss findings as false positives without re-testing the PoC yourself.
|
||||
|
||||
|
||||
@@ -1,23 +1,59 @@
|
||||
---
|
||||
name: managed-pentesting-with-strix
|
||||
description: Run a managed pentest of a web app or API through the app.strix.ai REST API — no local Docker, LLM key, or install needed. Create an API token, register domain/repository assets, launch and poll scans, triage vulnerabilities, export SARIF, download PDF/DOCX pentest reports for SOC 2 and other compliance evidence (Enterprise plan), start PR reviews, and set up schedules and webhooks. Use when the user wants continuous or scheduled pentesting-as-a-service, an auditor-ready pentest report, scans tracked in a team dashboard, or security testing from a sandboxed agent/CI environment with no infrastructure.
|
||||
description: Run a managed pentest of a web app, API, repository, or local workspace on the app.strix.ai platform with the `strix cloud` CLI or REST API — no local Docker or LLM key needed. Safely review and upload local source, register assets, launch and poll scans, triage vulnerabilities, export SARIF, download compliance reports, start PR reviews, buy credits, and set up schedules or webhooks. Use for managed, continuous, scheduled, team-tracked, or sandboxed-agent security testing.
|
||||
license: Apache-2.0
|
||||
metadata:
|
||||
author: usestrix
|
||||
homepage: https://docs.app.strix.ai
|
||||
---
|
||||
|
||||
# Strix Cloud API (managed, no local infra)
|
||||
# Strix Cloud (managed, no local infra)
|
||||
|
||||
Use this when you want Strix's autonomous pentesting **without running Docker or an LLM yourself** — the scan runs on Strix's infrastructure and results are tracked in a team dashboard. This is the right choice in sandboxed/hosted agent and CI environments, for teams, and for scheduled/continuous testing (downloadable PDF/DOCX reports are an Enterprise-plan feature). For fully local, free, air-gapped, or BYO-LLM runs, use the open-source CLI in the **penetration-testing-with-strix** skill instead — both share the same engine and SARIF output, so you can mix them.
|
||||
|
||||
Full reference: **[docs.app.strix.ai](https://docs.app.strix.ai)** · OpenAPI: `https://docs.app.strix.ai/openapi.json`
|
||||
There are two equivalent interfaces. Prefer the CLI:
|
||||
|
||||
## Setup
|
||||
- **`strix cloud` CLI** — every REST operation has a command in the form `strix cloud <resource> <verb>`. Install with `curl -sSL https://strix.ai/install | bash`. Run `strix cloud` to list all resources and `strix cloud <resource> help` (or `-h`) to list a resource's verbs; a bare resource with a safe read operation runs its documented default.
|
||||
- **REST API** — base URL `https://app.strix.ai/api/v1`, `Authorization: Bearer <token>` on every request. Full reference: **[docs.app.strix.ai](https://docs.app.strix.ai)** · agent index: `https://docs.app.strix.ai/llms.txt` · OpenAPI: `https://docs.app.strix.ai/openapi.json`.
|
||||
|
||||
- **Base URL:** `https://app.strix.ai/api/v1`
|
||||
- **Auth:** every request sends `Authorization: Bearer <token>`. Tokens are **org-scoped**.
|
||||
- **Get a token:** the user creates one in the dashboard at **Settings → API Access** (app.strix.ai). Ask them for it; never hardcode, log, or commit it. Store it in an env var or the CI secret store.
|
||||
The CLI is equally usable by agents and people. Output is complete JSON when stdout is not a terminal, or when you pass `--json`; terminal tables favor names, branches, lifecycle states, and numbered selectors. Human lists retain the selectors needed by follow-up commands but omit internal organization/user IDs; a selector too long for the compact table is repeated losslessly in a copyable block. Paginated lists print the next `--page` or `--offset`, and detail views preserve useful prose within a safe terminal bound; use `--json` for the complete record. Token lists label credentials as active, expired, or revoked. Binary downloads are the exception: redirect raw bytes intentionally, or use `--output FILE --json` to write the file and receive structured metadata. There are no interactive prompts when stdin is not a terminal. Exit codes: `0` success, `1` request/runtime error, `2` invalid usage, `4` authentication or plan limit, `5` payment required.
|
||||
|
||||
Every resource group with a safe read operation has a useful default action, and `-h` or `help` always shows its verbs. Native tab completion includes resources, verbs, flags, workspace commands, and local paths:
|
||||
|
||||
```bash
|
||||
source <(strix completions zsh) # current zsh session
|
||||
source <(strix completions bash) # current bash session
|
||||
strix completions fish | source # current fish session
|
||||
```
|
||||
|
||||
Write commands take request fields as flags. Every write command also accepts one JSON object with `--data`, which is the way to send fields that have no flag:
|
||||
|
||||
```bash
|
||||
strix cloud scans start --data '{"engagement_type":"code_review"}' # literal JSON
|
||||
strix cloud scans start --data @request.json # read a file
|
||||
cat request.json | strix cloud scans start --data - # read standard input
|
||||
```
|
||||
|
||||
The platform enforces plan and role limits, and the CLI passes the platform message through. Report downloads need the Enterprise plan. Schedules need the Pro plan. Billing writes need an admin token. A blocked command exits with code `4`.
|
||||
|
||||
## Setup: sign in
|
||||
|
||||
Run the device sign-in. It creates the user's account and workspace on first use and stores a personal API token in `~/.strix/platform-auth.json`:
|
||||
|
||||
```bash
|
||||
strix cloud login
|
||||
# Non-interactive least-privilege example:
|
||||
strix cloud login --scopes scans:read scans:write uploads:write billing:read vulnerabilities:read assets:read assets:write
|
||||
# Or use a stable named profile:
|
||||
strix cloud login --scope-profile recommended
|
||||
```
|
||||
|
||||
The user approves the sign-in in the browser. With `--scopes` (and optionally `--workspace <name-or-id>`) there are no terminal prompts, so the command works from a non-interactive agent shell. In an interactive terminal without flags, the CLI offers a workspace picker and scope presets (Recommended, Full access, Minimal, Custom). Recommended covers ordinary scans, source uploads, workspace switching, and user-approved credit top-ups; it excludes `tokens:write`, which must be requested explicitly when credential management is required. Use explicit scopes for a narrower automation token.
|
||||
|
||||
- `strix cloud whoami` is the fast local status. `strix cloud session --json` verifies the remote device session; `strix cloud session scopes` shows both effective access and the immutable login ceiling.
|
||||
- `strix cloud logout` revokes the remote session before removing the local token. On a network or server failure it keeps the token so the user can retry; `--local-only` deliberately skips revocation.
|
||||
- Every other `strix cloud` command uses the stored token automatically. `--token <token>` or `STRIX_API_TOKEN` is a stateless per-command override and never overwrites the stored account. For an override that is itself a CLI session, also pass `--workspace-id` or set `STRIX_WORKSPACE_ID`.
|
||||
- Never hardcode, log, or commit the token. Store it in an env var or the CI secret store.
|
||||
- **Scopes (least-privilege):** assign only what the integration needs and rotate regularly:
|
||||
|
||||
| Scope | Grants |
|
||||
@@ -28,15 +64,105 @@ Full reference: **[docs.app.strix.ai](https://docs.app.strix.ai)** · OpenAPI: `
|
||||
| `schedules:read` / `:write` | read schedules · create/trigger recurring scans |
|
||||
| `pr_reviews:write` | trigger PR security reviews |
|
||||
| `webhooks:read` / `:write` | manage webhook subscriptions |
|
||||
| `tokens:write` | create/revoke API tokens |
|
||||
| `uploads:write` | upload local source or documents for a scan |
|
||||
| `organizations:read` | read organization details (listing/switching the signed-in user's workspaces needs no API scope) |
|
||||
| `organizations:write` | create/update workspaces (admin) |
|
||||
| `tokens:write` | create/revoke ordinary API tokens (not needed to manage the current CLI session) |
|
||||
| `knowledge:read` / `:write` | read/update organization knowledge |
|
||||
| `audit:read` | read/export the Enterprise audit log |
|
||||
| `billing:read` / `billing:write` | read credit balance & auto top-up settings · buy credits (admin) |
|
||||
|
||||
HTTP errors map to messages and exit codes: `401` bad/expired token (exit `4`), `402` out of credits (exit `5`), `403` scope/plan-tier limit (exit `4`), `422` validation error (exit `1`).
|
||||
|
||||
Create a time-limited automation token with `strix cloud tokens create`. Use
|
||||
`--rbac-scopes` to restrict it to target IDs, tags, or business units; the value is a
|
||||
JSON array of `{ "type": "target|tag|business_unit", "value": "..." }` objects:
|
||||
|
||||
```bash
|
||||
export STRIX_API_TOKEN="<token>"
|
||||
BASE=https://app.strix.ai/api/v1
|
||||
auth=(-H "Authorization: Bearer $STRIX_API_TOKEN")
|
||||
strix cloud tokens create --type service --name staging-ci \
|
||||
--expires-at 2026-12-31T23:59:59Z \
|
||||
--scopes scans:read scans:write \
|
||||
--rbac-scopes '[{"type":"tag","value":"staging"}]'
|
||||
```
|
||||
|
||||
All examples use `jq` to parse JSON. Handle HTTP errors: `401` bad/expired token, `402` out of credits, `403` scope/plan-tier limit, `422` validation error.
|
||||
The token secret is returned once. Store it directly in a secret manager and do not
|
||||
print or commit it. `--expires-at` and `--expires-in-days` are mutually exclusive.
|
||||
|
||||
## 0. Credits & top-ups
|
||||
|
||||
Non-Enterprise scans consume org credits. Enterprise engagements are plan-included and do not debit the wallet. Check the balance before a scan (`billing:read`):
|
||||
|
||||
```bash
|
||||
strix cloud credits
|
||||
```
|
||||
|
||||
When the balance is too low, buy credits with `strix cloud billing topup` (`billing:write`, admin token). The server answers the first request with **HTTP 402 and a machine-payment challenge** (Stripe Machine Payments Protocol). The CLI pays the challenge with the Stripe Link wallet client when Node.js is available — the user approves the spend in the [Link app](https://link.com/agents). The response returns the receipt (`credits_granted`, `duplicate`, `reference`) and the new balance.
|
||||
|
||||
A default-tier source-only code review currently starts at 60 credits. Source uploads are not free: they launch an ordinary `code_review` and use the same deterministic scope estimator. The service checks the full balance before launch, reserves credits atomically only after validation succeeds, and does not create or charge a rejected scan. Retests and Enterprise scans are exempt.
|
||||
|
||||
```bash
|
||||
strix cloud billing topup --credits 20 --yes # explicit approval; skips the TTY prompt
|
||||
strix cloud billing topup --credits 20 --no-pay # print the 402 challenge without paying
|
||||
```
|
||||
|
||||
The default payment path is the Stripe Link wallet. When no wallet is connected, an interactive `strix cloud billing topup` starts the Link sign-in for the user and prints the verification link. The user approves the connection one time in the Link app, and then approves each payment there. No keys or variables are necessary. In a non-interactive process, the command stops and tells the user to connect the wallet at [link.com/agents](https://link.com/agents) or to use the hosted checkout link.
|
||||
|
||||
In a non-interactive agent or CI process, payment never proceeds unless the command includes `--yes`. Show the challenge or estimated spend to the user and obtain approval before adding it. `--no-pay` always stops after printing the challenge.
|
||||
|
||||
If the user does not want a wallet, create a hosted checkout link with `strix cloud billing subscribe --plan strix_top_up` and give the link to the user. The user pays in the browser.
|
||||
|
||||
Automatic top-ups (admin): `strix cloud billing auto-topup` shows the setting. Enable it with:
|
||||
|
||||
```bash
|
||||
strix cloud billing auto-topup update --enabled --topup-credits 20 --monthly-cap-credits 200
|
||||
```
|
||||
|
||||
An omitted `--monthly-cap-credits` keeps the stored cap. Pass `--no-monthly-cap` to remove the cap.
|
||||
|
||||
### Workspaces and account setup
|
||||
|
||||
Manage workspaces with a personal token from `strix cloud login`:
|
||||
|
||||
```bash
|
||||
strix cloud workspaces list # numbered name/role/current list
|
||||
strix cloud workspaces create --name "My Team" # admin + organizations:write
|
||||
strix cloud workspaces use 2 # displayed number, exact name, or ID
|
||||
strix cloud workspace use "My Team" # singular `workspace` alias also works
|
||||
strix cloud session scopes # effective scopes + consent ceiling
|
||||
strix cloud session scopes set minimal # narrow the session
|
||||
strix cloud org members invite --email dev@example.com --role analyst
|
||||
```
|
||||
|
||||
`workspaces use` retargets the current personal token to a workspace the user already belongs to and stores the updated workspace metadata; the bearer secret and expiry stay unchanged. It does not reprompt during ordinary switches: the server preserves the chosen profile, enforces the immutable login ceiling, and caps effective scopes by the target role. Use `--scope-profile` or `--scopes` to narrow within that ceiling; broader consent requires `strix cloud login` again. The CLI pins each process to the workspace it started in, so concurrent shells fail with a recoverable conflict instead of silently crossing organizations.
|
||||
|
||||
### Handoffs a person must finish
|
||||
|
||||
Four steps end at the user. The command creates the link or the record and prints it. Strix opens the browser only in an interactive terminal. Pass `--no-browser` to print the URL only.
|
||||
|
||||
```bash
|
||||
strix cloud billing subscribe --plan strix_cloud # hosted checkout page for the Cloud plan
|
||||
strix cloud billing portal # billing portal for the card and the plan
|
||||
strix cloud integrations install github # GitHub App or Slack installation page
|
||||
strix cloud domains verify <domain-id> # DNS record to add, then run it again
|
||||
```
|
||||
|
||||
Give the printed URL or DNS record to the user and wait. Do not claim that the payment, the installation, or the DNS change is complete. Confirm the result afterwards with `strix cloud credits`, `strix cloud integrations list`, or `strix cloud domains list`. All four commands need an admin token, except `domains verify`, which needs `assets:write`.
|
||||
|
||||
### Organization knowledge
|
||||
|
||||
Agents can manage the organization knowledge base without the dashboard (`knowledge:read` / `knowledge:write`):
|
||||
|
||||
```bash
|
||||
strix cloud knowledge list --search authentication
|
||||
strix cloud knowledge add --title "Authentication" --content "Staging uses SSO."
|
||||
strix cloud knowledge update <document-id> --content "Staging uses SSO and TOTP."
|
||||
strix cloud knowledge delete <document-id>
|
||||
strix cloud knowledge policies add --key staging-only --content "Never test production."
|
||||
strix cloud knowledge policies delete staging-only
|
||||
strix cloud knowledge repos entries usestrix/strix
|
||||
```
|
||||
|
||||
Knowledge policy writes require an admin token. Repository names are passed as normal `owner/name` values; the CLI handles URL encoding. The `costs` and `llm-settings` commands target on-prem installations and return `404` on app.strix.ai.
|
||||
|
||||
## 1. Register the target as an asset
|
||||
|
||||
@@ -44,109 +170,152 @@ Scans run against **registered assets**, not raw URLs. Register once, then reuse
|
||||
|
||||
```bash
|
||||
# Domain (black-box / live target). Requires domain verification before external scanning.
|
||||
# asset_type must be one of: web_app | api | attack_surface.
|
||||
curl -sS "$BASE/domains" "${auth[@]}" -H "Content-Type: application/json" \
|
||||
-d '{"domain":"staging.example.com","asset_type":"web_app"}' | jq '{id:.domain.id, status, reachable, verification}'
|
||||
# --asset-type must be one of: web_app | api | attack_surface.
|
||||
strix cloud domains add --domain staging.example.com --asset-type web_app
|
||||
|
||||
# Repository (white-box / code review). `full_name` is "owner/name".
|
||||
# Send one repository object, or a bare JSON array for several — not an object
|
||||
# wrapping a "repositories" key (that is rejected with 400).
|
||||
curl -sS "$BASE/repositories" "${auth[@]}" -H "Content-Type: application/json" \
|
||||
-d '[{"full_name":"org/app","provider":"github"}]' | jq '.repositories[] | {id, full_name}'
|
||||
strix cloud repos add --data '{"full_name":"org/app","provider":"github"}'
|
||||
```
|
||||
|
||||
Look up existing assets instead of re-adding: `GET /domains`, `GET /repositories` (both `assets:read`, paginated with `?page=&limit=`).
|
||||
Look up existing assets instead of re-adding: `strix cloud domains list`, `strix cloud repos list` (both `assets:read`).
|
||||
|
||||
## 2. Launch a scan
|
||||
|
||||
`POST /scans` (`scans:write`). Provide at least one target via `domain_ids`, `repository_ids`, or `internal_targets` (internal infra needs a network connector — see docs).
|
||||
`strix cloud scans start` (`scans:write`). Provide at least one target with `--domain-ids`, `--repository-ids`, or `--internal-targets` (internal infra needs a network connector — see docs).
|
||||
|
||||
```bash
|
||||
scan_id=$(curl -sS "$BASE/scans" "${auth[@]}" -H "Content-Type: application/json" -d '{
|
||||
"engagement_type": "live_test",
|
||||
"domain_ids": ["<domain-uuid>"],
|
||||
"focus": "IDOR, auth bypass, SSRF",
|
||||
"context": "Staging. Test account creds are configured as a test user.",
|
||||
"notify_on_completion": true
|
||||
}' | jq -r .scan_id)
|
||||
echo "$scan_id"
|
||||
strix cloud scans start \
|
||||
--engagement-type live_test \
|
||||
--domain-ids <domain-uuid> \
|
||||
--focus "IDOR, auth bypass, SSRF" \
|
||||
--context "Staging. Test account creds are configured as a test user." \
|
||||
--notify-on-completion
|
||||
```
|
||||
|
||||
Useful `CreateScanRequest` fields:
|
||||
Useful flags (each maps to a `CreateScanRequest` field):
|
||||
|
||||
| Field | Purpose |
|
||||
| Flag | Purpose |
|
||||
|---|---|
|
||||
| `engagement_type` | `live_test` (default), `code_review`, `internal_infra`, `compliance_pentest` |
|
||||
| `domain_ids` / `repository_ids` / `internal_targets` | targets (at least one) |
|
||||
| `domain_paths` / `repository_branches` | narrow to specific paths / branches |
|
||||
| `credentials` | authenticated scanning, incl. `mfa_method` (`totp`/`email_otp`/…) + `totp_secret` |
|
||||
| `headers` | extra HTTP headers (API keys, for example) for the target |
|
||||
| `focus` / `concerns` / `context` | steer the agents |
|
||||
| `upload_ids` | attach uploaded source/docs archives for white-box context |
|
||||
| `notify_on_completion` / `notification_emails` | email when done |
|
||||
| `--engagement-type` | `live_test` (default), `code_review`, `internal_infra`, `compliance_pentest` |
|
||||
| `--domain-ids` / `--repository-ids` / `--internal-targets` | targets (at least one) |
|
||||
| `--domain-paths` / `--repository-branches` | narrow to specific paths / branches (JSON maps) |
|
||||
| `--credentials` | authenticated scanning, incl. `mfa_method` (`totp`/`email_otp`/…) + `totp_secret` (JSON list) |
|
||||
| `--headers` | extra target HTTP headers as a JSON array of header objects |
|
||||
| `--focus` / `--concerns` / `--context` | free-form strings that steer the agents |
|
||||
| `--upload-ids` | attach uploaded source/docs archives for white-box context |
|
||||
| `--notify-on-completion` / `--notification-emails` | email when done |
|
||||
|
||||
Response is `{ scan_id, title, status }` with `status` = `pending`.
|
||||
Without `--source`, the response is `{ scan_id, title, status }` with `status` = `pending`.
|
||||
Local-source success wraps that platform response as
|
||||
`{ source, upload_id, scan: { scan_id, title, status } }`, so automation can retain the exact
|
||||
approved manifest and staged-upload identifier alongside the created scan.
|
||||
|
||||
## 3. Poll to completion
|
||||
### Scan a local workspace in the cloud
|
||||
|
||||
`GET /scans/{scanId}` (`scans:read`). Status flow: `pending → running → completed` (or `failed` / `cancelled`). Poll on an interval — scans take minutes to hours. Do not block.
|
||||
For an agent or CI workflow, bind approval to the exact source snapshot that was reviewed. Run
|
||||
the dry run with the intended source-selection flags, review the manifest and selected paths,
|
||||
and capture `source.archive_sha256`. Then repeat the same `--source`, every `--exclude`, and
|
||||
any `--include-hidden`, `--include-sensitive`, or `--include-archives` flags with
|
||||
`--approve-sha256`:
|
||||
|
||||
```bash
|
||||
while :; do
|
||||
s=$(curl -sS "$BASE/scans/$scan_id" "${auth[@]}" | jq -r .status)
|
||||
echo "status=$s"; [[ "$s" =~ ^(completed|failed|cancelled)$ ]] && break
|
||||
sleep 60
|
||||
done
|
||||
strix cloud scans start --source . --exclude 'private/' --dry-run --show-files --json
|
||||
# After reviewing the output, capture its source.archive_sha256 value:
|
||||
SOURCE_SHA256="<reviewed source.archive_sha256>"
|
||||
# Repeat every source-selection flag unchanged; a source-only scan infers code_review.
|
||||
strix cloud scans start --source . --exclude 'private/' \
|
||||
--approve-sha256 "$SOURCE_SHA256" --wait
|
||||
```
|
||||
|
||||
The CLI rebuilds the archive and refuses the upload if its SHA-256 no longer matches. `--yes`
|
||||
has deliberately narrower semantics: it approves only the snapshot built during that one
|
||||
invocation. Use it for a deliberate human or one-shot approval, not as the second half of a
|
||||
digest-bound agent/CI review. Without a TTY, a source upload requires either matching
|
||||
`--approve-sha256` approval or `--yes`; an interactive terminal can instead show the summary,
|
||||
the selected filenames when `--show-files` is set, and a `[y/N]` confirmation for its current
|
||||
snapshot.
|
||||
|
||||
The default selection is privacy-conscious: in a Git worktree it includes tracked files plus untracked files that are not ignored; it honors `.gitignore`, excludes every hidden path component, always excludes `.git`, symlinks, dependencies/build output, secret-like filenames, and nested archives. Add project exclusions to `.strixignore` (one exclude glob per line) or repeat `--exclude GLOB`; a trailing slash such as `private/` excludes that directory subtree.
|
||||
|
||||
The client refuses more than 20,000 files, a file over 25 MiB, more than 250 MiB expanded, or a ZIP over 50 MiB. The service then stream-inflates the ZIP and independently rejects malformed or unsupported entries, unsafe paths, too many entries, oversized entries, excessive expanded data, and oversized compressed input, so an untrusted client cannot bypass the ZIP-bomb controls by forging metadata.
|
||||
|
||||
Only use `--include-hidden`, `--include-sensitive`, or `--include-archives` after the dry-run manifest shows that the scan needs them. Hidden and sensitive files are separate opt-ins: for example, including `.env` requires both `--include-hidden` and `--include-sensitive`.
|
||||
|
||||
The CLI removes its private temporary local archive after every invocation. Once a remote
|
||||
upload is staged, a definitive scan rejection causes the CLI to delete it. A network failure,
|
||||
`5xx` response, malformed success response, or interruption after scan launch begins is
|
||||
ambiguous—the platform may have accepted the scan—so the CLI retains the upload and returns
|
||||
its `upload_id` with `launch_outcome_unknown: true`. If an automatic deletion attempt cannot
|
||||
be confirmed, it instead returns the retained `upload_id` with `cleanup_unknown: true`.
|
||||
Before retrying, run `strix cloud scans list` to avoid a duplicate scan or charge. If no scan
|
||||
is linked to the retained upload, remove it with `strix cloud uploads delete UPLOAD_ID`;
|
||||
linked uploads cannot be deleted.
|
||||
|
||||
With no explicit type, source alone infers `code_review`. Any domain target wins and infers `live_test`, so source plus a deployed domain is the normal white-box live-test workflow. Pass `--engagement-type` when you need to override the inference.
|
||||
|
||||
## 3. Wait for completion
|
||||
|
||||
Pass `--wait` to `scans start` to poll until the scan reaches a final state, or poll yourself with `strix cloud scans get <scan-id>` (`scans:read`). Bound automation with `--wait-timeout SECONDS`; timeout exits cleanly without cancelling the remote scan. Status flow: `pending → running → completed` (or `failed` / `cancelled`). Scans take minutes to hours — poll on an interval, do not block indefinitely.
|
||||
|
||||
## 4. Read findings
|
||||
|
||||
The scan-detail response includes `executive_summary`, `methodology`, `recommendations`, a `findings` severity roll-up, and a `vulnerabilities[]` array. Each vulnerability carries `title, severity, status, cvss, cwe, endpoint, method, impact, technical_analysis, poc_description, poc_script_code`, and (for code findings) `code_file`/`code_diff`/`code_before`/`code_after`.
|
||||
|
||||
```bash
|
||||
curl -sS "$BASE/scans/$scan_id" "${auth[@]}" \
|
||||
strix cloud scans get <scan-id> --json \
|
||||
| jq '["critical","high","medium","low","info"] as $order
|
||||
| .vulnerabilities
|
||||
| sort_by(.severity as $s | $order | index($s))
|
||||
| .[] | {title, severity, endpoint, cwe}'
|
||||
```
|
||||
|
||||
Cloud severities are `critical | high | medium | low` and statuses are `open | in_progress | fixed | ignored`. Sort by an explicit severity order rather than `sort_by(.severity)`, which sorts alphabetically (critical, high, low, medium).
|
||||
Cloud severities are `critical | high | medium | low` and statuses are `open | in_progress | snoozed | fixed | ignored | not_affected`. Sort by an explicit severity order rather than `sort_by(.severity)`, which sorts alphabetically (critical, high, low, medium).
|
||||
|
||||
Org-wide triage across scans: `GET /vulnerabilities` (`vulnerabilities:read`; filter by severity/status). Update triage state with the vulnerabilities `:write` endpoints. To remediate, hand off to the **fix-security-vulnerabilities-with-strix** skill.
|
||||
Org-wide triage across scans: `strix cloud vulns list --severity critical` (`vulnerabilities:read`, and it also filters by `--status`, `--scan-id`, and more). Update triage state with `strix cloud vulns update <id> --status fixed`. To remediate, hand off to the **fix-security-vulnerabilities-with-strix** skill.
|
||||
|
||||
## 5. Export & report
|
||||
|
||||
```bash
|
||||
# SARIF 2.1.0 for GitHub code scanning / ASPM ingestion
|
||||
curl -sS "$BASE/scans/$scan_id/sarif" "${auth[@]}" -o findings.sarif
|
||||
strix cloud scans sarif <scan-id> --output findings.sarif
|
||||
|
||||
# Report. The format and file type are query params (`Accept` is ignored):
|
||||
# format=technical (default) | retest | attestation | executive_summary
|
||||
# type=pdf (default) | docx
|
||||
# Any report download requires the Enterprise plan; formats beyond `technical`,
|
||||
# Report. Formats: technical (default) | retest | attestation | executive_summary
|
||||
# Types: pdf (default) | docx
|
||||
# Any report download requires the Enterprise plan. Formats beyond `technical`,
|
||||
# DOCX, and white-label branding are Enterprise-only too. Scan must be completed.
|
||||
curl -sS "$BASE/scans/$scan_id/report?format=technical&type=pdf" "${auth[@]}" -o strix-report.pdf
|
||||
strix cloud scans report <scan-id> --format technical --type pdf --output strix-report.pdf
|
||||
```
|
||||
|
||||
Downloads refuse to replace a file unless `--force` is explicit. Enterprise audit logs can be streamed as JSON or exported without trying to JSON-decode the body:
|
||||
|
||||
```bash
|
||||
strix cloud audit list --format csv --all --output audit.csv
|
||||
strix cloud audit list --format ndjson --all --output audit.ndjson
|
||||
```
|
||||
|
||||
## 6. PR reviews
|
||||
|
||||
Trigger an automated security review of a pull request (`pr_reviews:write`); results appear as PR comments and in the dashboard:
|
||||
Trigger an automated security review of a pull request (`pr_reviews:write`). Read the repository's `provider` and `installation_id` with `strix cloud repos list`; both identify the installed source-control integration. The results appear as PR comments and in the dashboard:
|
||||
|
||||
```bash
|
||||
curl -sS "$BASE/pr-reviews/start" "${auth[@]}" -H "Content-Type: application/json" \
|
||||
-d '{"repository_full_name":"org/app","pr_number":123}'
|
||||
strix cloud pr-reviews start \
|
||||
--provider github \
|
||||
--installation-id <installation-id> \
|
||||
--repository-full-name org/app \
|
||||
--pr-number 123
|
||||
```
|
||||
|
||||
List/inspect via `GET /pr-reviews` and `GET /pr-reviews/{id}`. Repo-level PR-review behavior is configured with the repository-settings endpoint.
|
||||
List/inspect with `strix cloud pr-reviews list` and `strix cloud pr-reviews get <id>`. Repo-level PR-review behavior is configured with `strix cloud pr-reviews settings`.
|
||||
|
||||
## 7. Continuous testing (schedules & webhooks)
|
||||
|
||||
- **Schedules** (`schedules:write`, Pro plan): create recurring scans and trigger them on demand — the managed equivalent of a cron-driven CLI loop.
|
||||
- **Webhooks** (`webhooks:write`): subscribe to pentest/vulnerability lifecycle events such as `scan.completed` and `vulnerability.created` to push results into Slack, ticketing, or your own pipeline instead of polling.
|
||||
- **Schedules** (`schedules:write`, Pro plan): `strix cloud schedules create` makes recurring scans, and `strix cloud schedules trigger <id>` runs one on demand — the managed equivalent of a cron-driven CLI loop.
|
||||
- **Webhooks** (`webhooks:write`): `strix cloud webhooks create` subscribes to pentest/vulnerability lifecycle events such as `scan.completed` and `vulnerability.created` to push results into Slack, ticketing, or your own pipeline instead of polling.
|
||||
|
||||
See the schedules and webhooks sections at [docs.app.strix.ai](https://docs.app.strix.ai) for payloads.
|
||||
|
||||
Network connectors are Enterprise-only. `strix cloud connectors create` may return a one-time enrollment command containing credentials; do not paste it into logs, and request it with `--include-command` only when the user is ready to install it. Browser checkout, source-control installation, DNS verification, connector installation, chat sharing, and publishing SARIF to an external provider are user handoffs or explicit external mutations—prepare the command/link, then obtain the appropriate approval before completing them.
|
||||
|
||||
## Safety
|
||||
|
||||
Only scan assets the user's organization owns or is authorized to test. External domain scans require verification (DNS/file/meta-tag) enforced by the platform — do not try to bypass it.
|
||||
|
||||
@@ -13,7 +13,7 @@ The OWASP Top 10 is a taxonomy of risk categories, not a test suite — "OWASP T
|
||||
|
||||
**Use the current edition: [OWASP Top 10:2025](https://owasp.org/Top10/)** (8th installment, superseding 2021). Ask the user before targeting an older edition — some compliance checklists still reference 2021, and a report labelled with the wrong edition is misleading. Key differences from 2021: **SSRF is folded into A01**, **A03 Software Supply Chain Failures** expands the old "Vulnerable and Outdated Components", and **A10 Mishandling of Exceptional Conditions** is new; A02 Security Misconfiguration moved 5→2.
|
||||
|
||||
Install, LLM setup, and the managed-cloud alternative: **penetration-testing-with-strix**.
|
||||
Install, LLM setup, and the managed-cloud alternative: **penetration-testing-with-strix**. For a run with no Docker and no LLM key, the same binary drives the managed platform: `strix cloud login`, then `strix cloud scans start ...` (details in **managed-pentesting-with-strix**).
|
||||
|
||||
## What is and is not testable by an agent
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ metadata:
|
||||
Strix runs autonomous AI pentesting agents that dynamically exploit a target and only report findings validated with a working proof-of-concept. There are **two ways to run it, built on the same engine and producing the same findings** — pick per situation, and mix them freely:
|
||||
|
||||
- **Open-source CLI** (self-hosted) — runs on your machine in a Docker sandbox with your own LLM key. Free, fully local, BYO-LLM, air-gap capable. Docs: [docs.strix.ai](https://docs.strix.ai).
|
||||
- **Cloud API** (managed) — runs on Strix's infrastructure via `https://app.strix.ai/api/v1`. No Docker, no LLM key, no local compute; adds team dashboards, scheduling, PR reviews, downloadable PDF/DOCX reports (Enterprise plan), and internal-network connectors. Docs: [docs.app.strix.ai](https://docs.app.strix.ai). Full workflow in the **managed-pentesting-with-strix** skill.
|
||||
- **Managed cloud** — runs on Strix's infrastructure, driven from the same CLI (`strix cloud ...`) or the REST API at `https://app.strix.ai/api/v1`. No Docker, no LLM key, no local compute; adds team dashboards, scheduling, PR reviews, downloadable PDF/DOCX reports (Enterprise plan), and internal-network connectors. Docs: [docs.app.strix.ai](https://docs.app.strix.ai). Full workflow in the **managed-pentesting-with-strix** skill.
|
||||
|
||||
## Which one? (decide, do not default)
|
||||
|
||||
@@ -122,27 +122,33 @@ Artifacts land in `strix_runs/<run-name>/`:
|
||||
|
||||
---
|
||||
|
||||
# Option B — Cloud API (managed, no local infra)
|
||||
# Option B — Managed cloud (no local infra)
|
||||
|
||||
Full details, asset registration, polling, reports, PR reviews, schedules, and webhooks are in the **managed-pentesting-with-strix** skill. Minimal launch-and-poll:
|
||||
The same `strix` binary drives the managed platform. Every command starts with `strix cloud`. Full details — asset registration, source uploads, reports, PR reviews, schedules, webhooks, and billing — are in the **managed-pentesting-with-strix** skill. Minimal flow:
|
||||
|
||||
```bash
|
||||
export STRIX_API_TOKEN="<token>" # org-scoped bearer, from Settings → API Access at app.strix.ai
|
||||
BASE=https://app.strix.ai/api/v1
|
||||
# 1. Sign in (device flow — the user confirms a code in the browser; this also
|
||||
# creates the account and workspace when needed)
|
||||
strix cloud login
|
||||
|
||||
# 1. Launch a scan against an already-registered domain/repo asset
|
||||
scan_id=$(curl -sS "$BASE/scans" \
|
||||
-H "Authorization: Bearer $STRIX_API_TOKEN" -H "Content-Type: application/json" \
|
||||
-d '{"engagement_type":"live_test","domain_ids":["<domain-uuid>"]}' | jq -r .scan_id)
|
||||
# If you need specific scopes, request them with --scopes:
|
||||
# strix cloud login --scopes scans:read scans:write assets:read assets:write \
|
||||
# vulnerabilities:read billing:read billing:write
|
||||
|
||||
# 2. Poll until terminal (pending → running → completed/failed/cancelled)
|
||||
curl -sS "$BASE/scans/$scan_id" -H "Authorization: Bearer $STRIX_API_TOKEN" | jq '.status'
|
||||
# 2. Register and verify the target domain (verification prints a DNS record for the user)
|
||||
strix cloud domains add --domain staging.example.com --asset-type web_app
|
||||
strix cloud domains verify <domain-id>
|
||||
|
||||
# 3. Read validated findings from the scan detail's `vulnerabilities[]`, or export SARIF
|
||||
curl -sS "$BASE/scans/$scan_id/sarif" -H "Authorization: Bearer $STRIX_API_TOKEN" -o findings.sarif
|
||||
# 3. Launch and wait
|
||||
strix cloud scans start --engagement-type live_test --domain-ids <domain-id> --wait
|
||||
|
||||
# 4. Read validated findings
|
||||
strix cloud vulns list --severity critical
|
||||
```
|
||||
|
||||
Ask the user to create the token (and register the target as a domain/repository asset) if they have not. If Docker/local prerequisites are not already satisfied, use this path instead of trying to install infra.
|
||||
For a local repository, `strix cloud scans start --source .` uploads the working tree (needs `uploads:write`) and infers a code review. When credits run out, `strix cloud billing topup` starts an agent-payable Stripe challenge — the managed skill covers the payment flow. Output is JSON when stdout is not a terminal, so the commands compose in scripts.
|
||||
|
||||
The raw REST API works too (`https://app.strix.ai/api/v1`, org-scoped bearer token — see [docs.app.strix.ai](https://docs.app.strix.ai)). If Docker or local prerequisites are not already satisfied, use this path instead of trying to install infra.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -11,7 +11,7 @@ metadata:
|
||||
|
||||
Black-box (and optionally source-assisted) penetration testing of a running web app with Strix's autonomous agents. Every reported finding is validated with a working exploit, so there are no signature-based false positives to triage.
|
||||
|
||||
Install, LLM setup, all CLI flags, and the managed-cloud alternative are covered in the **penetration-testing-with-strix** skill — read it if the target is not a running web app, or if `strix --version` fails. This skill is the web-app-specific workflow.
|
||||
Install, LLM setup, all CLI flags, and the managed-cloud alternative are covered in the **penetration-testing-with-strix** skill — read it if the target is not a running web app, or if `strix --version` fails. For a run with no Docker and no LLM key, the same binary drives the managed platform: `strix cloud login`, then `strix cloud scans start ...` (details in **managed-pentesting-with-strix**). This skill is the web-app-specific workflow.
|
||||
|
||||
## 1. Confirm authorization and scope
|
||||
|
||||
|
||||
@@ -52,6 +52,7 @@ from strix.tools.reporting.tool import (
|
||||
create_vulnerability_report,
|
||||
get_report,
|
||||
list_reports,
|
||||
update_vulnerability_report,
|
||||
)
|
||||
from strix.tools.respond.tool import respond_to_user
|
||||
from strix.tools.thinking.tool import think
|
||||
@@ -580,6 +581,7 @@ _BASE_TOOLS: tuple[Tool, ...] = (
|
||||
web_search,
|
||||
create_vulnerability_report,
|
||||
create_dependency_report,
|
||||
update_vulnerability_report,
|
||||
list_reports,
|
||||
get_report,
|
||||
list_requests,
|
||||
|
||||
@@ -235,11 +235,12 @@ VALIDATION REQUIREMENTS:
|
||||
- CLOSURE DISCIPLINE: every candidate you open ends in exactly one explicit state — `confirmed` (working PoC, or a complete source→control→sink→impact trace that is reachable), `ruled_out` (you can name the SPECIFIC control, at a location, that runs on every attacker-reachable path before the sink), or `open_proof_gap` (plausible, unconfirmed, and you could NOT name such a control). "I moved on" is not a closure state. Silently dropping an uncertain candidate is mislabelling an `open_proof_gap` as `ruled_out` and is how real bugs get missed.
|
||||
- Missing information is NOT proof of safety: no caller found, can't tell if deployed/exposed, couldn't stand up the service, build failed — each is an `open_proof_gap`, never a reason to mark a candidate clean. Difficulty is a reason to defer, not to suppress.
|
||||
- COVERAGE: record every surface you assess with `record_coverage` (surface + risk area + outcome + evidence), including the ones that came back clean — a report that only lists findings cannot say what was reviewed and cleared. Use the `needs_follow_up` outcome for anything left in an `open_proof_gap` state, and carry the same items up in `agent_finish(open_items=[...])`. The ledger is shared and mutable: when you resolve a surface another agent left open — or find that a closed one is not — move that entry with `update_coverage` instead of recording a second one for the same surface. The root agent reconciles all of it via `list_coverage` before `finish_scan`.
|
||||
- THREAT MODEL: before you start testing, call `get_threat_model` on the target you were pointed at — it is the scan's shared answer to who the attacker is, where the trust boundaries sit, and what counts as critical here, and it is cached per target rather than per scan. Read it instead of re-deriving trust boundaries yourself; where your testing disproves it — a boundary it calls trusted turns out to be attacker-reachable, a role it did not know about, a host or endpoint it never listed — record that with `amend_threat_model` so the agents after you inherit the correction. Amending is not optional politeness: a model nobody corrects turns the first agent's guesses into everyone's assumptions.
|
||||
- THREAT MODEL: before you start testing, call `get_threat_model` on the target you were pointed at — it is the scan's shared answer to who the attacker is, where the trust boundaries sit, and what counts as critical here. It is scoped to this scan and nothing carries over from an earlier run, so `found: false` means no agent on this run has derived one yet. Read it instead of re-deriving trust boundaries yourself; where your testing disproves it — a boundary it calls trusted turns out to be attacker-reachable, a role it did not know about, a host or endpoint it never listed — record that with `amend_threat_model` so the agents after you inherit the correction. Amending is not optional politeness: a model nobody corrects turns the first agent's guesses into everyone's assumptions.
|
||||
- Before filing any report, run the counterevidence pass: argue the strongest case AGAINST the finding, record what you found in the `counterevidence` field, set `confidence` honestly (a static-only trace you couldn't execute is at best `medium`), and state what evidence would change the severity. See the counterevidence and severity-calibration knowledge above.
|
||||
- A vulnerability is ONLY considered reported when a reporting agent uses create_vulnerability_report (or create_dependency_report for known-CVE dependency/supply-chain findings) with full details. Mentions in agent_finish, finish_scan, or generic messages are NOT sufficient
|
||||
- Reporting and fixing are ONE step, not two: when source is available, the reporting agent derives the concrete fix and files it INLINE via create_vulnerability_report (`code_locations` with `fix_before`/`fix_after` + `fix_pr_body`) — the report is not complete without it. Do NOT report first and then spawn a separate downstream agent to re-derive and re-apply the same patch; that just re-does the analysis and wastes tokens. (Do not silently patch a finding WITHOUT filing a report — the report, with its embedded fix, is the deliverable.)
|
||||
- DEDUPLICATION: The create_vulnerability_report tool uses LLM-based deduplication. If it rejects your report as a duplicate, DO NOT attempt to re-submit the same vulnerability. Accept the rejection and move on to testing other areas. The vulnerability has already been reported by another agent
|
||||
- DEDUPLICATION: The create_vulnerability_report tool uses LLM-based deduplication. If it rejects your report as a duplicate, DO NOT attempt to re-submit the same vulnerability. Accept the rejection and move on to testing other areas. The vulnerability has already been reported by another agent. If your evidence proves more than the finding it matched (a working exploit where that one had only a static trace, a chain that raises the impact), revise that finding with update_vulnerability_report using the duplicate_of id — never re-file it.
|
||||
- REVISING A FINDING: use update_vulnerability_report (report id + the fields you want to replace + update_reason) when you learn something a finding already on file does not carry — you built the PoC after filing it, a chain raised its impact, further testing weakened it, or its counterevidence/remediation/code locations were wrong. Editing a finding needs no duplicate verdict, and it is always better than filing a second report for the same issue. Read the finding first with get_report, and pass only the fields that change.
|
||||
- REVIEWING FILED FINDINGS (orchestrator/root agent): use list_reports to see every vulnerability filed so far in this scan (by any agent, root or child) — metadata-first with per-severity counts — and get_report to read one finding in full by its id. These are read-only orchestration tools: the root agent uses them to track coverage, avoid dispatching work on already-covered ground, assemble the finish_scan executive summary, and reason about attack-chaining across confirmed findings. Leaf/specialist agents should NOT call them — just do your assigned testing and file findings. Each entry shows which agent filed it (agent_name), and your own entries are flagged by_you. list_notes/get_note do the same for notes.
|
||||
|
||||
STATE & COORDINATION TOOLS (when and how):
|
||||
|
||||
@@ -18,7 +18,7 @@ from agents import (
|
||||
)
|
||||
from agents.model_settings import ModelSettings
|
||||
from agents.models.fake_id import FAKE_RESPONSES_ID
|
||||
from agents.models.interface import Model
|
||||
from agents.models.interface import Model, ModelProvider
|
||||
from agents.models.multi_provider import MultiProvider
|
||||
from agents.models.openai_responses import OpenAIResponsesModel
|
||||
from agents.retry import (
|
||||
@@ -48,7 +48,7 @@ if TYPE_CHECKING:
|
||||
from agents.agent_output import AgentOutputSchemaBase
|
||||
from agents.handoffs import Handoff
|
||||
from agents.items import ModelResponse, TResponseInputItem, TResponseStreamEvent
|
||||
from agents.models.interface import ModelProvider, ModelTracing
|
||||
from agents.models.interface import ModelTracing
|
||||
from agents.retry import ModelRetryAdvice, ModelRetryAdviceRequest
|
||||
from agents.tool import Tool
|
||||
from agents.usage import Usage
|
||||
@@ -445,12 +445,61 @@ def _response_usage(usage: Usage | None) -> ResponseUsage | None:
|
||||
)
|
||||
|
||||
|
||||
class _CredentialedLitellmProvider(ModelProvider):
|
||||
"""LiteLLM route bound to one endpoint's credentials.
|
||||
|
||||
``LitellmProvider`` reads them from the process-wide LiteLLM globals, which
|
||||
belong to the main model; a secondary endpoint needs its own.
|
||||
"""
|
||||
|
||||
def __init__(self, api_key: str | None, base_url: str | None) -> None:
|
||||
self._api_key = api_key
|
||||
self._base_url = base_url
|
||||
|
||||
def get_model(self, model_name: str | None) -> Model:
|
||||
from agents.extensions.models.litellm_model import LitellmModel
|
||||
from agents.models.default_models import get_default_model
|
||||
|
||||
return LitellmModel(
|
||||
model=model_name or get_default_model(),
|
||||
api_key=self._api_key,
|
||||
base_url=self._base_url,
|
||||
)
|
||||
|
||||
|
||||
class StrixProvider(MultiProvider):
|
||||
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
|
||||
so users type ``deepseek/deepseek-chat`` rather than
|
||||
``litellm/deepseek/deepseek-chat``.
|
||||
|
||||
``api_key``/``base_url`` bind every route this provider resolves to one
|
||||
endpoint, for a secondary model (the dedupe judge) whose endpoint differs
|
||||
from the main model's process-wide defaults.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
openai_api_key=api_key,
|
||||
openai_base_url=base_url,
|
||||
# A custom endpoint is OpenAI-compatible, i.e. chat completions; the
|
||||
# global default is the main model's and may say otherwise.
|
||||
openai_use_responses=False if base_url else None,
|
||||
**kwargs,
|
||||
)
|
||||
self._override_api_key = api_key
|
||||
self._override_base_url = base_url
|
||||
|
||||
def _create_fallback_provider(self, prefix: str) -> ModelProvider:
|
||||
if prefix == "litellm" and (self._override_api_key or self._override_base_url):
|
||||
return _CredentialedLitellmProvider(self._override_api_key, self._override_base_url)
|
||||
return super()._create_fallback_provider(prefix)
|
||||
|
||||
def _resolve_prefixed_model(
|
||||
self,
|
||||
*,
|
||||
@@ -864,6 +913,22 @@ def is_claude_model(model_name: str) -> bool:
|
||||
return "claude" in (model_name or "").strip().lower()
|
||||
|
||||
|
||||
def routes_through_litellm(model_name: str | None) -> bool:
|
||||
"""Whether :class:`StrixProvider` sends this model through LiteLLM.
|
||||
|
||||
Bare names and the ``openai/``/``any-llm/`` prefixes are served by the SDK's
|
||||
own clients, which raise ``TypeError`` on request fields they do not know,
|
||||
so LiteLLM-only fields must not be attached there. A bare ``claude-...``
|
||||
name is exactly that case: an ``LLM_API_BASE`` pointing at an
|
||||
OpenAI-compatible gateway in front of Claude.
|
||||
"""
|
||||
name = (model_name or "").strip()
|
||||
if not name or codex.subscription_model(name):
|
||||
return False
|
||||
prefix, _, rest = name.partition("/")
|
||||
return bool(rest) and prefix.lower() not in {"openai", "any-llm"}
|
||||
|
||||
|
||||
def is_bedrock_route(model_name: str) -> bool:
|
||||
name = (model_name or "").strip().lower()
|
||||
return name.startswith("bedrock/") or "anthropic." in name
|
||||
|
||||
@@ -291,6 +291,12 @@ class AgentCoordinator:
|
||||
self.pending_counts[target_agent_id] = self.pending_counts.get(target_agent_id, 0) + 1
|
||||
if from_user:
|
||||
runtime.user_wake_required = False
|
||||
self.errors.pop(target_agent_id, None)
|
||||
self.wait_kinds.pop(target_agent_id, None)
|
||||
self.recovery_counts.pop(target_agent_id, None)
|
||||
self.idle_resume_counts.pop(target_agent_id, None)
|
||||
self._parent_notified.discard(target_agent_id)
|
||||
self.statuses[target_agent_id] = "waiting"
|
||||
runtime.wake.set()
|
||||
stream = runtime.stream
|
||||
interrupt_on_message = runtime.interrupt_on_message
|
||||
|
||||
@@ -18,6 +18,7 @@ from strix.config.models import (
|
||||
is_openrouter_model,
|
||||
model_supports_reasoning,
|
||||
request_timeout_extra_args,
|
||||
routes_through_litellm,
|
||||
)
|
||||
from strix.core.sessions import scrub_images_from_items
|
||||
|
||||
@@ -267,7 +268,7 @@ def make_model_settings(
|
||||
and model_supports_reasoning(model_name)
|
||||
):
|
||||
model_settings = model_settings.resolve(
|
||||
_reasoning_settings(reasoning_effort, model_settings.extra_args),
|
||||
_reasoning_settings(reasoning_effort),
|
||||
)
|
||||
if force_required_tool_choice and _accepts_required_tool_choice(model_name):
|
||||
model_settings = model_settings.resolve(ModelSettings(tool_choice="required"))
|
||||
@@ -293,20 +294,19 @@ def _request_headers(
|
||||
return headers or None
|
||||
|
||||
|
||||
def _reasoning_settings(
|
||||
effort: ReasoningEffort,
|
||||
extra_args: dict[str, Any] | None,
|
||||
) -> ModelSettings:
|
||||
def _reasoning_settings(effort: ReasoningEffort) -> ModelSettings:
|
||||
"""``max`` is not in the OpenAI SDK's ``Reasoning.effort`` enum, so send it as
|
||||
a raw body field instead — also keeping it clear of LiteLLM's DeepSeek mapping,
|
||||
which collapses every ``reasoning_effort`` level to plain thinking-enabled.
|
||||
Providers that don't support ``max`` reject the request.
|
||||
|
||||
It goes in ``extra_body``, the field every model implementation forwards as the
|
||||
request's ``extra_body``; the same value under ``extra_args`` collides with that
|
||||
keyword and raises before a request is ever sent.
|
||||
"""
|
||||
if effort != "max":
|
||||
return ModelSettings(reasoning=Reasoning(effort=effort))
|
||||
return ModelSettings(
|
||||
extra_args={**(extra_args or {}), "extra_body": {"reasoning_effort": "max"}},
|
||||
)
|
||||
return ModelSettings(extra_body={"reasoning_effort": "max"})
|
||||
|
||||
|
||||
def _prompt_cache_extra_args(model_name: str) -> dict[str, Any] | None:
|
||||
@@ -317,8 +317,13 @@ def _prompt_cache_extra_args(model_name: str) -> dict[str, Any] | None:
|
||||
it — elsewhere it leaks onto the wire and native Anthropic 400s). Unmapped
|
||||
Bedrock models get no points at all: Bedrock rejects the passed-through
|
||||
field outright.
|
||||
|
||||
The field is LiteLLM's own, consumed by its transform, so it only goes to
|
||||
routes LiteLLM serves. A bare ``claude-...`` name is served by the SDK's
|
||||
OpenAI client instead (a gateway in front of Claude), and that client raises
|
||||
``TypeError`` on request kwargs it does not know.
|
||||
"""
|
||||
if not is_claude_model(model_name):
|
||||
if not is_claude_model(model_name) or not routes_through_litellm(model_name):
|
||||
return None
|
||||
if is_bedrock_route(model_name) and not bedrock_route_supports_prompt_caching(model_name):
|
||||
return None
|
||||
|
||||
@@ -53,18 +53,43 @@ from strix.tools.output_store import (
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agents.mcp import MCPServer
|
||||
from agents.memory import SQLiteSession
|
||||
from agents.result import RunResultBase
|
||||
|
||||
from strix.runtime.status import StatusSink
|
||||
from strix.tools.mcp import ConnectedMcpServer, McpConnectionRequest
|
||||
from strix.tools.mcp import (
|
||||
ConnectedMcpServer,
|
||||
McpConnectionRequest,
|
||||
McpRegistry,
|
||||
SupervisedMcpSession,
|
||||
)
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
StreamEventSink = Callable[[str, Any], None]
|
||||
|
||||
# Receives the run's MCP connection roster as a list of non-secret status dicts
|
||||
# ({"name", "provider", "tool_count", "dead"}), once when the connections are
|
||||
# established and again each time a connection transitions to dead. An interface
|
||||
# can persist it, render it, or forward it on as connection status. Kept as a
|
||||
# snapshot of the whole roster (not a per-
|
||||
# connection delta) so every call carries a consistent, current picture.
|
||||
McpStatusSink = Callable[[list[dict[str, Any]]], None]
|
||||
|
||||
|
||||
def _mcp_roster_payload(registry: McpRegistry) -> list[dict[str, Any]]:
|
||||
"""The run's MCP roster as non-secret status dicts (name/provider/tool_count/dead)."""
|
||||
return [
|
||||
{
|
||||
"name": status.name,
|
||||
"provider": status.provider,
|
||||
"tool_count": status.tool_count,
|
||||
"dead": status.dead,
|
||||
}
|
||||
for status in registry.statuses()
|
||||
]
|
||||
|
||||
|
||||
def _mcp_startup_summary(connections: list[ConnectedMcpServer]) -> str:
|
||||
"""One user-facing line summarizing the MCP servers that connected."""
|
||||
@@ -91,6 +116,21 @@ def _record_mcp_connections(connections: list[ConnectedMcpServer]) -> None:
|
||||
report_state.record_mcp_connections([connection.name for connection in connections])
|
||||
|
||||
|
||||
def _persist_mcp_status(roster: list[dict[str, Any]]) -> None:
|
||||
"""Write the run's non-secret MCP connection status roster to run.json.
|
||||
|
||||
The viewer rebuilds its display by re-reading the run's files from disk, so
|
||||
it cannot see the in-memory ``mcp_status_sink`` the TUI consumes. Persisting
|
||||
the same non-secret roster (name / provider / tool_count / dead) gives the
|
||||
viewer a source it can poll. Runs regardless of whether an interface sink is
|
||||
attached, so the standalone / non-TUI CLI path records health too.
|
||||
"""
|
||||
report_state = get_global_report_state()
|
||||
if report_state is None:
|
||||
return
|
||||
report_state.record_mcp_connection_status(roster)
|
||||
|
||||
|
||||
def _merge_root_prompt_context(
|
||||
scope_context: dict[str, Any],
|
||||
extra_system_prompt_context: dict[str, Any] | None,
|
||||
@@ -157,6 +197,7 @@ async def run_strix_scan(
|
||||
extra_system_prompt_context: dict[str, Any] | None = None,
|
||||
status_sink: StatusSink | None = None,
|
||||
mcp_connection_requests: list[McpConnectionRequest] | None = None,
|
||||
mcp_status_sink: McpStatusSink | None = None,
|
||||
) -> RunResultBase | None:
|
||||
"""Run or resume one Strix scan against a sandbox.
|
||||
|
||||
@@ -222,11 +263,13 @@ async def run_strix_scan(
|
||||
|
||||
from strix.tools.coverage.tools import hydrate_coverage_from_disk
|
||||
from strix.tools.notes.tools import hydrate_notes_from_disk
|
||||
from strix.tools.threat_model.tools import hydrate_threat_models_from_disk
|
||||
from strix.tools.todo.tools import hydrate_todos_from_disk
|
||||
|
||||
hydrate_todos_from_disk(state_dir)
|
||||
hydrate_notes_from_disk(state_dir)
|
||||
hydrate_coverage_from_disk(state_dir)
|
||||
hydrate_threat_models_from_disk(state_dir)
|
||||
|
||||
root_id: str | None = None
|
||||
if is_resume:
|
||||
@@ -295,7 +338,7 @@ async def run_strix_scan(
|
||||
configure_spill_writer(_spill_to_workspace)
|
||||
|
||||
sessions_to_close: list[SQLiteSession] = []
|
||||
mcp_servers: list[MCPServer] = []
|
||||
mcp_sessions: list[SupervisedMcpSession] = []
|
||||
|
||||
try:
|
||||
targets = scan_config.get("targets") or []
|
||||
@@ -364,7 +407,7 @@ async def run_strix_scan(
|
||||
mcp_requests = mcp_connection_requests
|
||||
if mcp_requests:
|
||||
connections = await attach_mcp_requests(mcp_requests, mcp_registry)
|
||||
mcp_servers = [c.server for c in connections]
|
||||
mcp_sessions = [c.session for c in connections]
|
||||
# Recorded even when nothing connected, so a resumed run does not
|
||||
# keep attributing tool calls to servers it no longer has.
|
||||
_record_mcp_connections(connections)
|
||||
@@ -385,6 +428,31 @@ async def run_strix_scan(
|
||||
}
|
||||
for summary in mcp_registry.summaries()
|
||||
]
|
||||
# Feed a non-secret connection roster (name / provider /
|
||||
# tool_count / dead) to two consumers: once now (all
|
||||
# currently healthy) and again whenever a connection later
|
||||
# dies. It is always persisted to run.json so the viewer,
|
||||
# which re-reads the run's files from disk, can render the
|
||||
# MCP connections panel and health without an in-memory
|
||||
# sink. When an interface sink is attached (the TUI backend,
|
||||
# or pro forwarding into the app's event stream) it also
|
||||
# receives the same snapshot. In-use is derived separately by
|
||||
# each interface from the connection-tagged tool-call events,
|
||||
# so it is not carried here.
|
||||
def _emit_mcp_status() -> None:
|
||||
roster = _mcp_roster_payload(mcp_registry)
|
||||
_persist_mcp_status(roster)
|
||||
if mcp_status_sink is not None:
|
||||
try:
|
||||
mcp_status_sink(roster)
|
||||
except Exception:
|
||||
logger.exception("MCP status sink failed")
|
||||
|
||||
for connection_name in mcp_registry.names():
|
||||
entry = mcp_registry.get(connection_name)
|
||||
if entry is not None:
|
||||
entry.session.set_on_dead(_emit_mcp_status)
|
||||
_emit_mcp_status()
|
||||
except Exception:
|
||||
logger.exception("Failed to connect user MCP servers; continuing without them")
|
||||
|
||||
@@ -579,9 +647,9 @@ async def run_strix_scan(
|
||||
for s in sessions_to_close:
|
||||
with contextlib.suppress(Exception):
|
||||
s.close()
|
||||
for mcp_server in mcp_servers:
|
||||
for mcp_session in mcp_sessions:
|
||||
with contextlib.suppress(Exception):
|
||||
await mcp_server.cleanup() # type: ignore[no-untyped-call]
|
||||
await mcp_session.aclose()
|
||||
with contextlib.suppress(Exception):
|
||||
await coordinator._maybe_snapshot()
|
||||
if cleanup_on_exit:
|
||||
|
||||
@@ -105,14 +105,15 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915
|
||||
report_state.set_scan_config(scan_config)
|
||||
report_state.save_run_data()
|
||||
|
||||
def display_vulnerability(report: dict[str, Any]) -> None:
|
||||
def display_vulnerability(report: dict[str, Any], *, updated: bool = False) -> None:
|
||||
report_id = report.get("id", "unknown")
|
||||
|
||||
vuln_text = format_vulnerability_report(report)
|
||||
|
||||
suffix = " (updated)" if updated else ""
|
||||
vuln_panel = Panel(
|
||||
vuln_text,
|
||||
title=f"[bold red]{report_id.upper()}",
|
||||
title=f"[bold red]{report_id.upper()}{suffix}",
|
||||
title_align="left",
|
||||
border_style="red",
|
||||
padding=(1, 2),
|
||||
@@ -122,6 +123,9 @@ async def run_cli(args: Any) -> None: # noqa: PLR0915
|
||||
console.print()
|
||||
|
||||
report_state.vulnerability_found_callback = display_vulnerability
|
||||
report_state.vulnerability_updated_callback = lambda report: display_vulnerability(
|
||||
report, updated=True
|
||||
)
|
||||
|
||||
def cleanup_on_exit() -> None:
|
||||
report_state.cleanup()
|
||||
|
||||
169
strix/interface/cloud/__init__.py
Normal file
169
strix/interface/cloud/__init__.py
Normal file
@@ -0,0 +1,169 @@
|
||||
"""`strix cloud` — the managed Strix platform (app.strix.ai) from the terminal.
|
||||
|
||||
Every command maps to one operation of the public REST API. Output is JSON
|
||||
when stdout is not a terminal, so agents can parse every result. Exit codes:
|
||||
0 success, 1 error, 2 invalid usage, 4 authentication required, 5 payment
|
||||
required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
from strix.interface.cloud.render import json_mode
|
||||
from strix.interface.cloud.runner import resolve, run
|
||||
from strix.interface.cloud.session import run_session
|
||||
from strix.interface.cloud.spec import DEFAULT_VERBS, GROUP_HELP, SPEC
|
||||
from strix.interface.cloud.workspaces import run_workspace_use
|
||||
from strix.interface.platform_cli import run_login
|
||||
from strix.interface.terminal_text import sanitize_terminal_text
|
||||
|
||||
|
||||
_USAGE_HEADER = """[bold]Usage:[/] strix cloud <command> [arguments]
|
||||
|
||||
[bold]Session commands:[/]
|
||||
login Sign in to the managed platform and store an API token
|
||||
logout Remove the stored API token
|
||||
whoami Show the stored account, workspace, and token state
|
||||
session Inspect or narrow the remote CLI session
|
||||
credits Show the credit balance of the workspace
|
||||
|
||||
[bold]Resource commands:[/]"""
|
||||
|
||||
_USAGE_FOOTER = """
|
||||
Run [bold]strix cloud <command> help[/] to list its verbs. Common read-only
|
||||
commands may also run their default verb when no verb is given.
|
||||
Every REST resource command accepts [bold]--json[/] and [bold]--token[/]. Write
|
||||
commands accept [bold]--data[/] with a JSON object of extra request fields.
|
||||
Login is an interactive device flow; [bold]whoami[/] and [bold]logout[/] also
|
||||
produce JSON automatically when output is redirected.
|
||||
API reference: https://docs.app.strix.ai"""
|
||||
|
||||
_HELP_TOKENS = frozenset({"-h", "--help", "help"})
|
||||
|
||||
|
||||
def _is_help_request(argv: list[str]) -> bool:
|
||||
"""Recognize a help token with an optional JSON-output flag in either order."""
|
||||
return sum(argument in _HELP_TOKENS for argument in argv) == 1 and all(
|
||||
argument in _HELP_TOKENS or argument == "--json" for argument in argv
|
||||
)
|
||||
|
||||
|
||||
def run_cloud(argv: list[str]) -> int:
|
||||
"""Run a managed-cloud command without ever leaking a Ctrl-C traceback."""
|
||||
try:
|
||||
return _run_cloud(argv)
|
||||
except KeyboardInterrupt:
|
||||
if json_mode(flag="--json" in argv):
|
||||
sys.stdout.write(json.dumps({"error": "Interrupted.", "interrupted": True}) + "\n")
|
||||
else:
|
||||
Console(stderr=True).print("[yellow]Interrupted.[/]")
|
||||
return 130
|
||||
|
||||
|
||||
def _run_cloud(argv: list[str]) -> int: # noqa: PLR0911, PLR0912
|
||||
"""Entry point for ``strix cloud …``. Returns a process exit code."""
|
||||
console = Console()
|
||||
as_json = json_mode(flag="--json" in argv)
|
||||
if not argv or _is_help_request(argv):
|
||||
if as_json:
|
||||
_print_usage_json()
|
||||
else:
|
||||
_print_usage(console)
|
||||
return 0
|
||||
if argv == ["--json"]:
|
||||
_print_usage_json()
|
||||
return 0
|
||||
|
||||
group, rest = argv[0], argv[1:]
|
||||
if group == "workspace":
|
||||
group = "workspaces"
|
||||
if group in ("login", "logout", "whoami"):
|
||||
return _run_session(console, group, rest)
|
||||
if group == "session":
|
||||
return run_session(rest)
|
||||
if group == "credits":
|
||||
group, rest = "billing", ["credits", *rest]
|
||||
if group == "workspaces" and rest and rest[0] == "use":
|
||||
try:
|
||||
return run_workspace_use(rest[1:])
|
||||
except http.CloudError as exc:
|
||||
if "--json" in rest:
|
||||
sys.stdout.write(json.dumps({"error": str(exc)}) + "\n")
|
||||
else:
|
||||
console.print(f"[red]Error:[/] {escape(sanitize_terminal_text(exc))}")
|
||||
return exc.exit_code
|
||||
|
||||
if group not in SPEC:
|
||||
if as_json:
|
||||
sys.stdout.write(json.dumps({"error": f"unknown command: {group}"}) + "\n")
|
||||
return 2
|
||||
console.print(f"[red]Unknown command:[/] {escape(sanitize_terminal_text(group))}")
|
||||
_print_usage(console)
|
||||
return 2
|
||||
group_help = _is_help_request(rest)
|
||||
resolved = None if group_help else resolve(group, rest)
|
||||
if resolved is None:
|
||||
help_tokens: set[str] = set(_HELP_TOKENS) if group_help else set()
|
||||
invalid = [arg for arg in rest if arg != "--json" and arg not in help_tokens]
|
||||
_print_verbs(console, group, as_json=as_json, error="unknown verb" if invalid else None)
|
||||
return 2 if invalid else 0
|
||||
cmd, remaining = resolved
|
||||
verb_label = " ".join(rest[: len(rest) - len(remaining)]) or DEFAULT_VERBS.get(group, "")
|
||||
return run(group, verb_label, cmd, remaining)
|
||||
|
||||
|
||||
def _run_session(_console: Console, group: str, rest: list[str]) -> int:
|
||||
if rest and rest[0] == "help":
|
||||
rest = ["--help", *rest[1:]]
|
||||
session_argv = {
|
||||
"login": rest,
|
||||
"logout": ["logout", *rest],
|
||||
"whoami": ["status", *rest],
|
||||
}
|
||||
return run_login(session_argv[group])
|
||||
|
||||
|
||||
def _print_usage(console: Console) -> None:
|
||||
console.print(_USAGE_HEADER)
|
||||
for group in SPEC:
|
||||
console.print(f" {group:<14}{GROUP_HELP.get(group, '')}")
|
||||
console.print(_USAGE_FOOTER)
|
||||
|
||||
|
||||
def _print_verbs(
|
||||
console: Console, group: str, *, as_json: bool = False, error: str | None = None
|
||||
) -> None:
|
||||
if as_json:
|
||||
verbs: list[dict[str, str]] = [
|
||||
{"name": verb, "help": command.help} for verb, command in SPEC[group].items()
|
||||
]
|
||||
if group == "workspaces":
|
||||
verbs.append({"name": "use", "help": "Switch the stored token to another workspace."})
|
||||
payload: dict[str, object] = {
|
||||
"command": f"strix cloud {group}",
|
||||
"verbs": verbs,
|
||||
}
|
||||
if error:
|
||||
payload["error"] = error
|
||||
sys.stdout.write(json.dumps(payload, indent=2) + "\n")
|
||||
return
|
||||
console.print(f"[bold]strix cloud {group}[/] verbs:")
|
||||
for verb, cmd in SPEC[group].items():
|
||||
console.print(f" {verb:<28}{cmd.help}")
|
||||
if group == "workspaces":
|
||||
console.print(f" {'use':<28}Switch the stored token to another workspace.")
|
||||
|
||||
|
||||
def _print_usage_json() -> None:
|
||||
payload = {
|
||||
"command": "strix cloud",
|
||||
"session_commands": ["login", "logout", "whoami", "session", "credits"],
|
||||
"resource_commands": [{"name": group, "help": GROUP_HELP.get(group, "")} for group in SPEC],
|
||||
}
|
||||
sys.stdout.write(json.dumps(payload, indent=2) + "\n")
|
||||
18
strix/interface/cloud/arguments.py
Normal file
18
strix/interface/cloud/arguments.py
Normal file
@@ -0,0 +1,18 @@
|
||||
"""Argument parsing that reports managed-cloud usage errors through one contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from typing import NoReturn
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
|
||||
|
||||
class CloudArgumentParser(argparse.ArgumentParser):
|
||||
"""Raise a typed usage error instead of printing argparse prose and exiting."""
|
||||
|
||||
def error(self, message: str) -> NoReturn:
|
||||
raise http.CloudError(
|
||||
f"invalid arguments for {self.prog}: {message}",
|
||||
exit_code=http.EXIT_USAGE,
|
||||
)
|
||||
718
strix/interface/cloud/billing.py
Normal file
718
strix/interface/cloud/billing.py
Normal file
@@ -0,0 +1,718 @@
|
||||
"""Billing top-up and agent-wallet execution for ``strix cloud``."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import webbrowser
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
from strix.interface.cloud.payment_proxy import WalletUpstreamResponse, wallet_payment_bridge
|
||||
from strix.interface.cloud.render import emit
|
||||
from strix.interface.terminal_text import sanitize_terminal_text
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import argparse
|
||||
|
||||
from rich.console import Console
|
||||
|
||||
|
||||
_MAX_WALLET_DETAIL_CHARS = 2_000
|
||||
# Keep the wallet client on the exact protocol implementation used by the
|
||||
# platform. This version is also old enough to remain installable in npm
|
||||
# environments that apply a short package-publication safety window.
|
||||
_MPPX_PACKAGE = "mppx@0.8.17"
|
||||
# Stripe's own wallet client. It runs the complete challenge flow: it creates a
|
||||
# spend request, waits for the person to approve it in the Link app, and retries
|
||||
# the payment with the approved credential.
|
||||
_LINK_CLI_PACKAGE = "@stripe/link-cli@0.13.1"
|
||||
_LINK_CLI_CLIENT_NAME = "Strix CLI"
|
||||
_LINK_LOGIN_TIMEOUT_S = 300
|
||||
# Poll every 2 seconds while the person approves the spend request in the Link
|
||||
# app. 150 attempts give the person 5 minutes.
|
||||
_LINK_APPROVAL_POLL_INTERVAL_S = 2
|
||||
_LINK_APPROVAL_MAX_ATTEMPTS = 150
|
||||
# Bound every wallet subprocess so a stalled npm download or wallet request
|
||||
# cannot block the top-up command forever. The poll step gets the full
|
||||
# approval window plus this margin.
|
||||
_WALLET_STEP_TIMEOUT_S = 300
|
||||
_LINK_APPROVAL_TIMEOUT_S = (
|
||||
_LINK_APPROVAL_POLL_INTERVAL_S * _LINK_APPROVAL_MAX_ATTEMPTS + _WALLET_STEP_TIMEOUT_S
|
||||
)
|
||||
_NPM_REGISTRY = "https://registry.npmjs.org"
|
||||
_WALLET_ENV_NAMES = frozenset(
|
||||
{
|
||||
"ALL_PROXY",
|
||||
"APPDATA",
|
||||
"COLORTERM",
|
||||
"COMSPEC",
|
||||
"FORCE_COLOR",
|
||||
"HOME",
|
||||
"HTTPS_PROXY",
|
||||
"HTTP_PROXY",
|
||||
"LANG",
|
||||
"LC_ALL",
|
||||
"LC_CTYPE",
|
||||
"LOCALAPPDATA",
|
||||
"NO_COLOR",
|
||||
"NO_PROXY",
|
||||
"PATH",
|
||||
"PATHEXT",
|
||||
"SSL_CERT_DIR",
|
||||
"SSL_CERT_FILE",
|
||||
"SYSTEMROOT",
|
||||
"TEMP",
|
||||
"TERM",
|
||||
"TMP",
|
||||
"TMPDIR",
|
||||
"USERPROFILE",
|
||||
"XDG_CONFIG_HOME",
|
||||
"XDG_DATA_HOME",
|
||||
"XDG_STATE_HOME",
|
||||
"all_proxy",
|
||||
"http_proxy",
|
||||
"https_proxy",
|
||||
"no_proxy",
|
||||
}
|
||||
)
|
||||
_AUTHORIZATION_SECRET = re.compile(r"(?i)((?:bearer|payment)\s+)[^\s\"']+")
|
||||
_LOOPBACK_NO_PROXY = ("127.0.0.1", "localhost", "::1")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _WalletClientResult:
|
||||
process: subprocess.CompletedProcess[str]
|
||||
upstream_responses: tuple[WalletUpstreamResponse, ...]
|
||||
|
||||
|
||||
def run_topup( # noqa: PLR0911, PLR0912, PLR0915
|
||||
console: Console,
|
||||
args: argparse.Namespace,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
as_json: bool,
|
||||
token: str | None,
|
||||
) -> int:
|
||||
"""Handle the HTTP 402 challenge and optional agent-wallet payment."""
|
||||
response = http.request("POST", "/billing/topup", token=token, body=body)
|
||||
if response.status_code != 402:
|
||||
emit(console, http.check(response), as_json=as_json)
|
||||
return http.EXIT_OK
|
||||
|
||||
challenge = http.parsed(response)
|
||||
if getattr(args, "no_pay", False):
|
||||
emit(
|
||||
console,
|
||||
{"error": "Payment required", "challenge": challenge},
|
||||
as_json=as_json,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
|
||||
credit_count = body.get("credits")
|
||||
if not getattr(args, "yes", False):
|
||||
if as_json or not (sys.stdin.isatty() and sys.stdout.isatty()):
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": (
|
||||
"Payment requires explicit approval in non-interactive mode. "
|
||||
"Review the challenge, then re-run with --yes to authorize payment."
|
||||
),
|
||||
"challenge": challenge,
|
||||
},
|
||||
as_json=as_json,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
answer = console.input(f"Buy {credit_count} credit(s) now? [y/N]: ").strip().lower()
|
||||
if answer not in ("y", "yes"):
|
||||
console.print("[yellow]Payment cancelled.[/]")
|
||||
return http.EXIT_PAYMENT
|
||||
|
||||
npx = shutil.which("npx")
|
||||
if npx is None:
|
||||
message = (
|
||||
"Payment requires a wallet client. Install Node.js and run the command again, "
|
||||
"or pay the challenge with an MPP wallet client."
|
||||
)
|
||||
if as_json:
|
||||
emit(
|
||||
console,
|
||||
{"error": message, "challenge": challenge},
|
||||
as_json=True,
|
||||
)
|
||||
else:
|
||||
emit(console, challenge, as_json=False)
|
||||
console.print(f"[yellow]Payment required.[/] {message}")
|
||||
return http.EXIT_PAYMENT
|
||||
|
||||
payment_method = getattr(args, "payment_method", None) or os.environ.get(
|
||||
"MPPX_STRIPE_PAYMENT_METHOD"
|
||||
)
|
||||
use_link_wallet = payment_method is None and not _mppx_wallet_configured()
|
||||
if use_link_wallet:
|
||||
setup_error = _prepare_link_wallet(console, npx, as_json=as_json)
|
||||
if setup_error is not None:
|
||||
emit(
|
||||
console,
|
||||
{"error": setup_error, "challenge": challenge},
|
||||
as_json=as_json,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
try:
|
||||
wallet_result = _run_wallet_client(
|
||||
console,
|
||||
npx,
|
||||
args,
|
||||
body,
|
||||
token=token,
|
||||
payment_method=payment_method,
|
||||
use_link_wallet=use_link_wallet,
|
||||
capture_output=as_json,
|
||||
)
|
||||
except KeyboardInterrupt:
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": (
|
||||
"Payment was interrupted after the wallet started. The outcome is unknown; "
|
||||
"run `strix cloud billing credits` and check the balance before retrying."
|
||||
),
|
||||
"interrupted": True,
|
||||
"payment_outcome_unknown": True,
|
||||
},
|
||||
as_json=as_json,
|
||||
)
|
||||
return 130
|
||||
except OSError:
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": "Could not start the wallet client securely.",
|
||||
"challenge": challenge,
|
||||
},
|
||||
as_json=as_json,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
|
||||
result = wallet_result.process
|
||||
confirmed_receipt = _confirmed_topup_receipt(wallet_result.upstream_responses)
|
||||
if confirmed_receipt is not None:
|
||||
emit(console, confirmed_receipt, as_json=as_json)
|
||||
return http.EXIT_OK
|
||||
|
||||
stdout = str(getattr(result, "stdout", "") or "").strip()
|
||||
stderr = str(getattr(result, "stderr", "") or "").strip()
|
||||
if not as_json:
|
||||
console.print(
|
||||
"[yellow]The wallet exited without a confirmed receipt. The payment outcome is "
|
||||
"unknown; run `strix cloud billing credits` before retrying.[/]"
|
||||
)
|
||||
detail = _wallet_detail(stderr or stdout or "")
|
||||
if detail:
|
||||
console.print(f"[dim]Wallet output: {detail}[/]")
|
||||
return http.EXIT_PAYMENT
|
||||
if result.returncode == 0:
|
||||
try:
|
||||
receipt = json.loads(stdout)
|
||||
except (TypeError, ValueError):
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": (
|
||||
"The wallet reported success but did not return JSON. Check the credit "
|
||||
"balance before retrying payment."
|
||||
),
|
||||
"detail": _wallet_detail(stdout or stderr or "No wallet output was returned."),
|
||||
"payment_outcome_unknown": True,
|
||||
},
|
||||
as_json=True,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
if not _valid_topup_receipt(receipt):
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": (
|
||||
"The wallet returned an invalid top-up receipt. Check the credit balance "
|
||||
"before retrying payment."
|
||||
),
|
||||
"detail": _wallet_detail(stdout),
|
||||
"payment_outcome_unknown": True,
|
||||
},
|
||||
as_json=True,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": (
|
||||
"The wallet returned a receipt, but the Strix billing endpoint did not "
|
||||
"confirm it. Check the credit balance before retrying payment."
|
||||
),
|
||||
"detail": _wallet_detail(stdout),
|
||||
"payment_outcome_unknown": True,
|
||||
},
|
||||
as_json=True,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
|
||||
emit(
|
||||
console,
|
||||
{
|
||||
"error": (
|
||||
"The wallet exited without a confirmed receipt. The payment outcome is unknown; "
|
||||
"run `strix cloud billing credits` and check the balance before retrying."
|
||||
),
|
||||
"detail": _wallet_detail(
|
||||
stderr or stdout or f"Wallet client exited with status {result.returncode}."
|
||||
),
|
||||
"wallet_exit_code": result.returncode,
|
||||
"payment_outcome_unknown": True,
|
||||
},
|
||||
as_json=True,
|
||||
)
|
||||
return http.EXIT_PAYMENT
|
||||
|
||||
|
||||
def _run_wallet_client(
|
||||
console: Console,
|
||||
npx: str,
|
||||
args: argparse.Namespace,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
token: str | None,
|
||||
payment_method: str | None,
|
||||
use_link_wallet: bool,
|
||||
capture_output: bool,
|
||||
) -> _WalletClientResult:
|
||||
"""Run the wallet through the loopback bridge without exposing the API token."""
|
||||
upstream_url = f"{http.app_url()}/api/v1/billing/topup"
|
||||
body_json = json.dumps(body)
|
||||
wallet_env = _wallet_environment()
|
||||
upstream_responses: list[WalletUpstreamResponse] = []
|
||||
with tempfile.TemporaryDirectory(prefix="strix-wallet-") as wallet_cwd:
|
||||
wallet_root = Path(wallet_cwd)
|
||||
user_config = wallet_root / "user.npmrc"
|
||||
global_config = wallet_root / "global.npmrc"
|
||||
user_config.touch(mode=0o600)
|
||||
global_config.touch(mode=0o600)
|
||||
npx_prefix = _npx_prefix(npx, wallet_root)
|
||||
with wallet_payment_bridge(
|
||||
upstream_url=upstream_url,
|
||||
api_token=http.api_token(token),
|
||||
workspace_id=http.expected_workspace_id(token_override=token is not None),
|
||||
expected_body=body_json.encode(),
|
||||
timeout=getattr(args, "timeout", None),
|
||||
response_observer=upstream_responses.append,
|
||||
) as wallet_url:
|
||||
if use_link_wallet:
|
||||
process = _run_link_wallet_flow(
|
||||
console,
|
||||
npx_prefix,
|
||||
wallet_url,
|
||||
body,
|
||||
body_json,
|
||||
wallet_env,
|
||||
wallet_root,
|
||||
quiet=capture_output,
|
||||
)
|
||||
else:
|
||||
command = [
|
||||
*npx_prefix,
|
||||
_MPPX_PACKAGE,
|
||||
wallet_url,
|
||||
"--fail",
|
||||
"-J",
|
||||
body_json,
|
||||
]
|
||||
if payment_method:
|
||||
command += ["-M", f"paymentMethod={payment_method}"]
|
||||
try:
|
||||
process = subprocess.run( # noqa: S603
|
||||
command,
|
||||
check=False,
|
||||
capture_output=capture_output,
|
||||
text=True,
|
||||
env=wallet_env,
|
||||
cwd=wallet_root,
|
||||
timeout=_LINK_APPROVAL_TIMEOUT_S,
|
||||
)
|
||||
except subprocess.TimeoutExpired as timeout_error:
|
||||
process = subprocess.CompletedProcess(
|
||||
args=command,
|
||||
returncode=1,
|
||||
stdout=_decoded_stream(timeout_error.stdout),
|
||||
stderr=(
|
||||
"The wallet step did not complete within "
|
||||
f"{_LINK_APPROVAL_TIMEOUT_S} seconds."
|
||||
),
|
||||
)
|
||||
return _WalletClientResult(process=process, upstream_responses=tuple(upstream_responses))
|
||||
|
||||
|
||||
def _run_link_wallet_flow(
|
||||
console: Console,
|
||||
npx_prefix: list[str],
|
||||
wallet_url: str,
|
||||
body: dict[str, Any],
|
||||
body_json: str,
|
||||
wallet_env: dict[str, str],
|
||||
wallet_root: Path,
|
||||
*,
|
||||
quiet: bool,
|
||||
) -> subprocess.CompletedProcess[str]:
|
||||
"""Create the spend request, wait for approval in the Link app, then pay."""
|
||||
|
||||
def run_step(
|
||||
arguments: list[str],
|
||||
progress_message: str,
|
||||
timeout: int = _WALLET_STEP_TIMEOUT_S,
|
||||
) -> subprocess.CompletedProcess[str]:
|
||||
command = [*npx_prefix, _LINK_CLI_PACKAGE, *arguments]
|
||||
|
||||
def run() -> subprocess.CompletedProcess[str]:
|
||||
try:
|
||||
return subprocess.run( # noqa: S603
|
||||
command,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=wallet_env,
|
||||
cwd=wallet_root,
|
||||
timeout=timeout,
|
||||
)
|
||||
except subprocess.TimeoutExpired as timeout_error:
|
||||
return subprocess.CompletedProcess(
|
||||
args=command,
|
||||
returncode=1,
|
||||
stdout=_decoded_stream(timeout_error.stdout),
|
||||
stderr=f"The wallet step did not complete within {timeout} seconds.",
|
||||
)
|
||||
|
||||
if quiet:
|
||||
return run()
|
||||
with console.status(progress_message):
|
||||
return run()
|
||||
|
||||
created = run_step(
|
||||
[
|
||||
"mpp",
|
||||
"pay",
|
||||
wallet_url,
|
||||
"--method",
|
||||
"POST",
|
||||
"--data",
|
||||
body_json,
|
||||
"--context",
|
||||
_payment_context(body),
|
||||
"--format",
|
||||
"json",
|
||||
],
|
||||
"Starting the Stripe Link wallet…",
|
||||
)
|
||||
spend_request = _pending_spend_request(created.stdout)
|
||||
if spend_request is None:
|
||||
return created
|
||||
request_id, approval_url = spend_request
|
||||
|
||||
if not quiet:
|
||||
console.print(f"[yellow]Approve the payment in the Link app:[/] {approval_url}")
|
||||
if sys.stdin.isatty() and sys.stdout.isatty() and approval_url.startswith("https://"):
|
||||
with suppress(Exception):
|
||||
webbrowser.open(approval_url)
|
||||
polled = run_step(
|
||||
[
|
||||
"spend-request",
|
||||
"retrieve",
|
||||
request_id,
|
||||
"--interval",
|
||||
str(_LINK_APPROVAL_POLL_INTERVAL_S),
|
||||
"--max-attempts",
|
||||
str(_LINK_APPROVAL_MAX_ATTEMPTS),
|
||||
"--format",
|
||||
"jsonl",
|
||||
],
|
||||
"Waiting for the approval in the Link app…",
|
||||
timeout=_LINK_APPROVAL_TIMEOUT_S,
|
||||
)
|
||||
if _final_spend_request_status(polled.stdout) != "approved":
|
||||
return polled
|
||||
|
||||
return run_step(
|
||||
[
|
||||
"mpp",
|
||||
"pay",
|
||||
wallet_url,
|
||||
"--spend-request-id",
|
||||
request_id,
|
||||
"--method",
|
||||
"POST",
|
||||
"--data",
|
||||
body_json,
|
||||
"--format",
|
||||
"json",
|
||||
],
|
||||
"Completing the payment…",
|
||||
)
|
||||
|
||||
|
||||
def _decoded_stream(stream: str | bytes | None) -> str:
|
||||
"""Return captured subprocess output as text."""
|
||||
if stream is None:
|
||||
return ""
|
||||
if isinstance(stream, bytes):
|
||||
return stream.decode(errors="replace")
|
||||
return stream
|
||||
|
||||
|
||||
def _embedded_json_documents(text: str) -> list[Any]:
|
||||
"""Extract JSON documents from wallet output that can contain other text."""
|
||||
documents: list[Any] = []
|
||||
decoder = json.JSONDecoder()
|
||||
position = 0
|
||||
while position < len(text):
|
||||
start_candidates = [
|
||||
index for index in (text.find("[", position), text.find("{", position)) if index != -1
|
||||
]
|
||||
if not start_candidates:
|
||||
break
|
||||
start = min(start_candidates)
|
||||
try:
|
||||
document, end = decoder.raw_decode(text, start)
|
||||
except ValueError:
|
||||
position = start + 1
|
||||
continue
|
||||
documents.append(document)
|
||||
position = end
|
||||
return documents
|
||||
|
||||
|
||||
def _spend_request_records(stdout: str) -> list[dict[str, Any]]:
|
||||
"""Parse spend-request records from JSON or JSON-lines wallet output."""
|
||||
records: list[dict[str, Any]] = []
|
||||
for candidate in _embedded_json_documents((stdout or "").strip()):
|
||||
items = candidate if isinstance(candidate, list) else [candidate]
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
record = cast("dict[str, Any]", item)
|
||||
data = record.get("data")
|
||||
if isinstance(data, dict):
|
||||
record = cast("dict[str, Any]", data)
|
||||
records.append(record)
|
||||
return records
|
||||
|
||||
|
||||
def _pending_spend_request(stdout: str) -> tuple[str, str] | None:
|
||||
"""Find a spend request that waits for approval in the Link app."""
|
||||
for record in _spend_request_records(stdout):
|
||||
request_id = record.get("id")
|
||||
approval_url = record.get("approval_url")
|
||||
if (
|
||||
record.get("status") == "pending_approval"
|
||||
and isinstance(request_id, str)
|
||||
and request_id
|
||||
and isinstance(approval_url, str)
|
||||
):
|
||||
return request_id, approval_url
|
||||
return None
|
||||
|
||||
|
||||
def _final_spend_request_status(stdout: str) -> str | None:
|
||||
"""Return the last reported status from the approval poll output."""
|
||||
status: str | None = None
|
||||
for record in _spend_request_records(stdout):
|
||||
value = record.get("status")
|
||||
if isinstance(value, str):
|
||||
status = value
|
||||
return status
|
||||
|
||||
|
||||
def _npx_prefix(npx: str, wallet_root: Path) -> list[str]:
|
||||
"""Install the wallet client from a fixed registry without lifecycle scripts."""
|
||||
return [
|
||||
npx,
|
||||
"--yes",
|
||||
f"--registry={_NPM_REGISTRY}",
|
||||
"--ignore-scripts",
|
||||
f"--userconfig={wallet_root / 'user.npmrc'}",
|
||||
f"--globalconfig={wallet_root / 'global.npmrc'}",
|
||||
f"--cache={_wallet_npm_cache()}",
|
||||
]
|
||||
|
||||
|
||||
def _wallet_npm_cache() -> Path:
|
||||
"""Keep one private npm cache so the pinned wallet client installs once."""
|
||||
cache = Path.home() / ".strix" / "wallet-npm-cache"
|
||||
cache.mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
return cache
|
||||
|
||||
|
||||
def _payment_context(body: dict[str, Any]) -> str:
|
||||
"""Describe the purchase for the person who approves it in the Link app."""
|
||||
credits_requested = body.get("credits")
|
||||
return (
|
||||
f"Strix scan credits. The Strix command line interface asks to buy "
|
||||
f"{credits_requested} scan credit(s) for the selected Strix workspace on "
|
||||
"app.strix.ai. Strix spends the credits on managed penetration test scans "
|
||||
"that the user starts."
|
||||
)
|
||||
|
||||
|
||||
def _mppx_wallet_configured() -> bool:
|
||||
"""Report whether the person already configured the mppx wallet client."""
|
||||
return bool(os.environ.get("MPPX_ACCOUNT") or os.environ.get("MPPX_STRIPE_SECRET_KEY"))
|
||||
|
||||
|
||||
def _run_link_cli(
|
||||
npx: str,
|
||||
arguments: list[str],
|
||||
*,
|
||||
capture_output: bool,
|
||||
timeout: float | None = None,
|
||||
) -> subprocess.CompletedProcess[str]:
|
||||
"""Run one Stripe Link wallet command in an isolated npm environment."""
|
||||
with tempfile.TemporaryDirectory(prefix="strix-wallet-") as wallet_cwd:
|
||||
wallet_root = Path(wallet_cwd)
|
||||
(wallet_root / "user.npmrc").touch(mode=0o600)
|
||||
(wallet_root / "global.npmrc").touch(mode=0o600)
|
||||
return subprocess.run( # noqa: S603
|
||||
[*_npx_prefix(npx, wallet_root), _LINK_CLI_PACKAGE, *arguments],
|
||||
check=False,
|
||||
capture_output=capture_output,
|
||||
text=True,
|
||||
env=_wallet_environment(),
|
||||
cwd=wallet_root,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
|
||||
def _link_wallet_authenticated(npx: str) -> bool:
|
||||
"""Report whether a Link wallet is already connected to this machine."""
|
||||
try:
|
||||
result = _run_link_cli(
|
||||
npx,
|
||||
["auth", "status", "--format", "json"],
|
||||
capture_output=True,
|
||||
timeout=_LINK_LOGIN_TIMEOUT_S,
|
||||
)
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
return False
|
||||
try:
|
||||
payload = json.loads(result.stdout or "null")
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
if isinstance(payload, list):
|
||||
payload = payload[0] if payload else None
|
||||
return bool(isinstance(payload, dict) and payload.get("authenticated"))
|
||||
|
||||
|
||||
def _prepare_link_wallet(console: Console, npx: str, *, as_json: bool) -> str | None:
|
||||
"""Connect a Link wallet when none is present. Return an error message on failure."""
|
||||
if _link_wallet_authenticated(npx):
|
||||
return None
|
||||
|
||||
manual_setup = (
|
||||
"Payment needs a Stripe Link wallet. Run `strix cloud billing topup` in an "
|
||||
"interactive terminal to connect one, or set up the wallet at "
|
||||
"https://link.com/agents. For a browser checkout instead, run "
|
||||
"`strix cloud billing subscribe --plan strix_top_up`."
|
||||
)
|
||||
if as_json or not (sys.stdin.isatty() and sys.stdout.isatty()):
|
||||
return manual_setup
|
||||
|
||||
console.print(
|
||||
"[yellow]No Stripe Link wallet is connected.[/] Strix starts the Link sign-in now. "
|
||||
"Approve the connection in the Link app, then Strix continues the payment. "
|
||||
"The user approves every payment in the Link app."
|
||||
)
|
||||
try:
|
||||
_run_link_cli(
|
||||
npx,
|
||||
[
|
||||
"auth",
|
||||
"login",
|
||||
"--client-name",
|
||||
_LINK_CLI_CLIENT_NAME,
|
||||
"--interval",
|
||||
"3",
|
||||
"--timeout",
|
||||
str(_LINK_LOGIN_TIMEOUT_S),
|
||||
],
|
||||
capture_output=False,
|
||||
timeout=_LINK_LOGIN_TIMEOUT_S + 30,
|
||||
)
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
return manual_setup
|
||||
if _link_wallet_authenticated(npx):
|
||||
return None
|
||||
return manual_setup
|
||||
|
||||
|
||||
def _wallet_environment() -> dict[str, str]:
|
||||
"""Pass only platform essentials and explicit wallet variables to npm/mppx."""
|
||||
environment = {
|
||||
name: value
|
||||
for name, value in os.environ.items()
|
||||
if name in _WALLET_ENV_NAMES or name.startswith(("LINK_", "MPPX_"))
|
||||
}
|
||||
for name in ("NO_PROXY", "no_proxy"):
|
||||
entries = [entry.strip() for entry in environment.get(name, "").split(",") if entry.strip()]
|
||||
normalized = {entry.lower().strip("[]") for entry in entries}
|
||||
entries.extend(host for host in _LOOPBACK_NO_PROXY if host not in normalized)
|
||||
environment[name] = ",".join(entries)
|
||||
return environment
|
||||
|
||||
|
||||
def _wallet_detail(value: str) -> str:
|
||||
"""Bound and redact third-party wallet diagnostics before returning JSON."""
|
||||
redacted = _AUTHORIZATION_SECRET.sub(r"\1[redacted]", sanitize_terminal_text(value))
|
||||
if len(redacted) <= _MAX_WALLET_DETAIL_CHARS:
|
||||
return redacted
|
||||
return redacted[: _MAX_WALLET_DETAIL_CHARS - 1] + "…"
|
||||
|
||||
|
||||
def _valid_topup_receipt(value: Any) -> bool:
|
||||
"""Require the documented success shape before reporting a paid top-up."""
|
||||
if not isinstance(value, dict):
|
||||
return False
|
||||
fields = cast("dict[str, Any]", value)
|
||||
credits_granted = fields.get("credits_granted")
|
||||
balance = fields.get("balance")
|
||||
return (
|
||||
isinstance(credits_granted, int)
|
||||
and not isinstance(credits_granted, bool)
|
||||
and credits_granted >= 0
|
||||
and isinstance(fields.get("duplicate"), bool)
|
||||
and isinstance(fields.get("reference"), str)
|
||||
and bool(fields["reference"])
|
||||
and isinstance(balance, int)
|
||||
and not isinstance(balance, bool)
|
||||
and balance >= 0
|
||||
)
|
||||
|
||||
|
||||
def _confirmed_topup_receipt(
|
||||
responses: tuple[WalletUpstreamResponse, ...],
|
||||
) -> dict[str, Any] | None:
|
||||
"""Return a receipt only when the trusted bridge observed its successful response."""
|
||||
for response in reversed(responses):
|
||||
if not 200 <= response.status_code < 300:
|
||||
continue
|
||||
try:
|
||||
receipt = json.loads(response.body)
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if _valid_topup_receipt(receipt):
|
||||
return cast("dict[str, Any]", receipt)
|
||||
return None
|
||||
361
strix/interface/cloud/http.py
Normal file
361
strix/interface/cloud/http.py
Normal file
@@ -0,0 +1,361 @@
|
||||
"""HTTP client for the managed Strix platform API (app.strix.ai)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from urllib.parse import SplitResult, urlsplit
|
||||
|
||||
import requests
|
||||
|
||||
from strix.config import load_settings
|
||||
from strix.interface.platform_cli import read_record
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
_DEFAULT_TIMEOUT_S = 120
|
||||
_SUPABASE_STORAGE_HOST = re.compile(r"^[a-z0-9-]+\.supabase\.co$")
|
||||
_STORAGE_PATH_PREFIX = "/storage/v1/"
|
||||
_app_url_override: str | None = None
|
||||
_token_override_active = False
|
||||
_workspace_id_override: str | None = None
|
||||
_timeout_s: float = _DEFAULT_TIMEOUT_S
|
||||
|
||||
EXIT_OK = 0
|
||||
EXIT_ERROR = 1
|
||||
EXIT_USAGE = 2
|
||||
EXIT_AUTH = 4
|
||||
EXIT_PAYMENT = 5
|
||||
|
||||
|
||||
class CloudError(Exception):
|
||||
"""A failed cloud command. Carries the process exit code."""
|
||||
|
||||
def __init__(self, message: str, *, exit_code: int = EXIT_ERROR, payload: Any = None) -> None:
|
||||
super().__init__(message)
|
||||
self.exit_code = exit_code
|
||||
self.payload = payload
|
||||
|
||||
|
||||
class CloudTransportError(CloudError):
|
||||
"""A request may have reached the platform, but no response was received."""
|
||||
|
||||
|
||||
def configure(
|
||||
*,
|
||||
base_url: str | None = None,
|
||||
timeout: float | None = None,
|
||||
token_override: bool = False,
|
||||
workspace_id: str | None = None,
|
||||
) -> None:
|
||||
"""Set the platform URL and the request timeout for this process."""
|
||||
global _app_url_override, _timeout_s, _token_override_active # noqa: PLW0603
|
||||
global _workspace_id_override # noqa: PLW0603
|
||||
_app_url_override = base_url.rstrip("/") if base_url else None
|
||||
_token_override_active = token_override
|
||||
explicit_workspace = workspace_id or os.environ.get("STRIX_WORKSPACE_ID")
|
||||
if explicit_workspace:
|
||||
_workspace_id_override = explicit_workspace.strip()
|
||||
elif not token_override and not os.environ.get("STRIX_API_TOKEN"):
|
||||
record = read_record()
|
||||
stored_workspace = record.get("organization_id") if record is not None else None
|
||||
_workspace_id_override = (
|
||||
stored_workspace.strip()
|
||||
if isinstance(stored_workspace, str) and stored_workspace.strip()
|
||||
else None
|
||||
)
|
||||
else:
|
||||
_workspace_id_override = None
|
||||
if timeout is not None:
|
||||
if not math.isfinite(timeout) or timeout <= 0:
|
||||
raise CloudError(
|
||||
"request timeout must be a finite number greater than 0.",
|
||||
exit_code=EXIT_USAGE,
|
||||
)
|
||||
_timeout_s = timeout
|
||||
|
||||
|
||||
def app_url() -> str:
|
||||
if _app_url_override:
|
||||
return _app_url_override
|
||||
viewer = load_settings().viewer
|
||||
configured = viewer.app_url.rstrip("/")
|
||||
explicitly_configured = bool(os.environ.get("STRIX_APP_URL")) or "app_url" in getattr(
|
||||
viewer, "model_fields_set", set[str]()
|
||||
)
|
||||
if explicitly_configured or _token_override_active or os.environ.get("STRIX_API_TOKEN"):
|
||||
return configured
|
||||
record = read_record()
|
||||
stored = record.get("app_url") if record is not None else None
|
||||
if isinstance(stored, str) and stored:
|
||||
try:
|
||||
_parse_origin_url(stored, label="stored platform URL")
|
||||
except CloudError:
|
||||
pass
|
||||
else:
|
||||
return stored.rstrip("/")
|
||||
return configured
|
||||
|
||||
|
||||
def api_token(override: str | None = None) -> str:
|
||||
token = override or os.environ.get("STRIX_API_TOKEN")
|
||||
if not token:
|
||||
record = read_record()
|
||||
if record is not None:
|
||||
stored = record.get("api_token")
|
||||
if isinstance(stored, str):
|
||||
_validate_stored_token_origin(record)
|
||||
token = stored
|
||||
if not token or not token.strip():
|
||||
raise CloudError(
|
||||
"not signed in. Run `strix cloud login`, or set STRIX_API_TOKEN.",
|
||||
exit_code=EXIT_AUTH,
|
||||
)
|
||||
return token.strip()
|
||||
|
||||
|
||||
def _validate_stored_token_origin(record: dict[str, Any]) -> None:
|
||||
"""Never send a stored bearer token to an origin other than its issuer."""
|
||||
stored_url = record.get("app_url")
|
||||
if not isinstance(stored_url, str) or not stored_url:
|
||||
raise CloudError(
|
||||
"the stored sign-in is not bound to a trusted platform. Run `strix cloud login` "
|
||||
"again before using it.",
|
||||
exit_code=EXIT_AUTH,
|
||||
)
|
||||
try:
|
||||
stored_origin = _origin(_parse_origin_url(stored_url, label="stored platform URL"))
|
||||
active_origin = _origin(_parse_origin_url(app_url(), label="configured platform URL"))
|
||||
except CloudError as exc:
|
||||
raise CloudError(
|
||||
"the stored sign-in has an invalid platform binding. Run `strix cloud login` again.",
|
||||
exit_code=EXIT_AUTH,
|
||||
) from exc
|
||||
if stored_origin != active_origin:
|
||||
raise CloudError(
|
||||
"the stored sign-in belongs to a different platform. Refusing to send its token; "
|
||||
"run `strix cloud login` for the configured platform or supply an explicit token.",
|
||||
exit_code=EXIT_AUTH,
|
||||
)
|
||||
|
||||
|
||||
def request(
|
||||
method: str,
|
||||
path: str,
|
||||
*,
|
||||
token: str | None = None,
|
||||
query: dict[str, Any] | None = None,
|
||||
body: dict[str, Any] | None = None,
|
||||
stream: bool = False,
|
||||
idempotency_key: str | None = None,
|
||||
) -> requests.Response:
|
||||
url = f"{app_url()}/api/v1{path}"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_token(token)}",
|
||||
}
|
||||
workspace_id = expected_workspace_id(token_override=token is not None)
|
||||
if workspace_id:
|
||||
headers["X-Strix-Workspace"] = workspace_id
|
||||
if idempotency_key is not None:
|
||||
headers["Idempotency-Key"] = idempotency_key
|
||||
try:
|
||||
response = requests.request(
|
||||
method,
|
||||
url,
|
||||
headers=headers,
|
||||
params={
|
||||
key: ("true" if value else "false") if isinstance(value, bool) else value
|
||||
for key, value in (query or {}).items()
|
||||
if value is not None
|
||||
}
|
||||
or None,
|
||||
json=body,
|
||||
timeout=_timeout_s,
|
||||
stream=stream,
|
||||
allow_redirects=False,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
raise CloudTransportError(f"could not reach {app_url()}: {exc}") from exc
|
||||
return response
|
||||
|
||||
|
||||
def expected_workspace_id(*, token_override: bool) -> str | None:
|
||||
"""Pin every request in this process to the workspace selected at startup."""
|
||||
if _workspace_id_override:
|
||||
return _workspace_id_override
|
||||
if token_override or _token_override_active or os.environ.get("STRIX_API_TOKEN"):
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def upload_file(signed_url: str, upload_token: str, path: Path) -> None:
|
||||
"""Stream a file to a platform-issued storage URL."""
|
||||
_validate_upload_url(signed_url)
|
||||
response: requests.Response | None = None
|
||||
try:
|
||||
with path.open("rb") as stream:
|
||||
response = requests.put(
|
||||
signed_url,
|
||||
data=stream,
|
||||
headers={
|
||||
"Authorization": f"Bearer {upload_token}",
|
||||
"Content-Type": "application/zip",
|
||||
},
|
||||
timeout=_timeout_s,
|
||||
allow_redirects=False,
|
||||
)
|
||||
except (OSError, requests.RequestException) as exc:
|
||||
raise CloudError(f"source upload failed: {exc}") from exc
|
||||
try:
|
||||
if 300 <= response.status_code < 400:
|
||||
raise CloudError("source upload refused an unexpected redirect")
|
||||
if not response.ok:
|
||||
detail = ""
|
||||
try:
|
||||
payload = response.json()
|
||||
if isinstance(payload, dict):
|
||||
fields = cast("dict[str, Any]", payload)
|
||||
detail = str(fields.get("message") or fields.get("error") or "")
|
||||
except ValueError:
|
||||
pass
|
||||
raise CloudError(detail or f"source upload failed (HTTP {response.status_code})")
|
||||
finally:
|
||||
response.close()
|
||||
|
||||
|
||||
def _validate_upload_url(signed_url: str) -> None:
|
||||
"""Allow uploads only to the trusted app origin or managed Supabase storage."""
|
||||
# Supabase signed upload URLs carry their signature in the query string.
|
||||
# Keep every origin/path restriction below, but allow that opaque query on
|
||||
# this one platform-issued URL type.
|
||||
target = _parse_origin_url(
|
||||
signed_url,
|
||||
label="source upload URL",
|
||||
allow_query=True,
|
||||
)
|
||||
if not target.path.startswith(_STORAGE_PATH_PREFIX):
|
||||
raise CloudError("source upload refused a URL outside the storage API")
|
||||
|
||||
configured_app = _parse_origin_url(app_url(), label="configured platform URL")
|
||||
if _origin(target) == _origin(configured_app):
|
||||
return
|
||||
if _is_loopback_host(configured_app.hostname or "") and _is_loopback_host(
|
||||
target.hostname or ""
|
||||
):
|
||||
return
|
||||
|
||||
hostname = target.hostname or ""
|
||||
if (
|
||||
target.scheme == "https"
|
||||
and target.port in (None, 443)
|
||||
and _SUPABASE_STORAGE_HOST.fullmatch(hostname)
|
||||
):
|
||||
return
|
||||
raise CloudError(
|
||||
"source upload refused an untrusted storage origin; only the configured platform "
|
||||
"origin and managed Supabase storage are allowed"
|
||||
)
|
||||
|
||||
|
||||
def _parse_origin_url(
|
||||
value: str,
|
||||
*,
|
||||
label: str,
|
||||
allow_query: bool = False,
|
||||
) -> SplitResult:
|
||||
try:
|
||||
parsed = urlsplit(value)
|
||||
port = parsed.port
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise CloudError(f"{label} is invalid") from exc
|
||||
hostname = parsed.hostname
|
||||
if (
|
||||
parsed.scheme not in {"http", "https"}
|
||||
or not hostname
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or (parsed.query and not allow_query)
|
||||
or parsed.fragment
|
||||
or "\\" in value
|
||||
or any(character.isspace() for character in value)
|
||||
or "%" in parsed.netloc
|
||||
):
|
||||
raise CloudError(f"{label} is invalid")
|
||||
try:
|
||||
hostname.encode("ascii")
|
||||
except UnicodeEncodeError as exc:
|
||||
raise CloudError(f"{label} contains a non-ASCII hostname") from exc
|
||||
if port is not None and not 1 <= port <= 65535:
|
||||
raise CloudError(f"{label} is invalid")
|
||||
return parsed
|
||||
|
||||
|
||||
def _origin(parsed: SplitResult) -> tuple[str, str, int]:
|
||||
default_port = 443 if parsed.scheme == "https" else 80
|
||||
return parsed.scheme, (parsed.hostname or "").lower(), parsed.port or default_port
|
||||
|
||||
|
||||
def _is_loopback_host(hostname: str) -> bool:
|
||||
normalized = hostname.lower().rstrip(".")
|
||||
if normalized == "localhost" or normalized.endswith(".localhost"):
|
||||
return True
|
||||
try:
|
||||
return ipaddress.ip_address(normalized).is_loopback
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def parsed(response: requests.Response) -> Any:
|
||||
content_type = response.headers.get("content-type", "")
|
||||
if "application/json" in content_type:
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError:
|
||||
return response.text
|
||||
return response.text
|
||||
|
||||
|
||||
def check(response: requests.Response) -> Any:
|
||||
data = parsed(response)
|
||||
if 200 <= response.status_code < 300:
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
if "application/json" not in content_type:
|
||||
raise CloudError(
|
||||
"the server returned a non-JSON response. Check STRIX_APP_URL and preview "
|
||||
"access, then retry."
|
||||
)
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError as exc:
|
||||
raise CloudError(
|
||||
"the server returned malformed JSON. Check STRIX_APP_URL and preview "
|
||||
"access, then retry."
|
||||
) from exc
|
||||
detail = ""
|
||||
error_code = ""
|
||||
if isinstance(data, dict):
|
||||
raw = cast("dict[str, Any]", data)
|
||||
detail = str(raw.get("detail") or raw.get("error") or "")
|
||||
error_code = str(raw.get("code") or raw.get("error_code") or "")
|
||||
nested_error = raw.get("error")
|
||||
if isinstance(nested_error, dict):
|
||||
nested = cast("dict[str, Any]", nested_error)
|
||||
error_code = error_code or str(nested.get("code") or "")
|
||||
detail = str(nested.get("message") or detail)
|
||||
message = detail or f"HTTP {response.status_code}"
|
||||
if error_code == "scan_credit_limit_reached":
|
||||
raise CloudError(message, exit_code=EXIT_PAYMENT, payload=data)
|
||||
if response.status_code in (401, 403):
|
||||
raise CloudError(message, exit_code=EXIT_AUTH, payload=data)
|
||||
if response.status_code == 402:
|
||||
hint = detail or (
|
||||
"not enough credits. Run `strix cloud billing topup --credits N` to buy credits."
|
||||
)
|
||||
raise CloudError(hint, exit_code=EXIT_PAYMENT, payload=data)
|
||||
raise CloudError(message, exit_code=EXIT_ERROR, payload=data)
|
||||
286
strix/interface/cloud/payment_proxy.py
Normal file
286
strix/interface/cloud/payment_proxy.py
Normal file
@@ -0,0 +1,286 @@
|
||||
"""Loopback bridge for wallet clients that only accept secrets in argv.
|
||||
|
||||
The ``mppx`` CLI accepts custom HTTP headers through ``-H`` only. Passing a
|
||||
Strix API token that way exposes it to process-listing tools. This module keeps
|
||||
the token in the Strix process and injects it while forwarding the wallet's few
|
||||
requests (challenge probes and the paid retry) to the fixed billing endpoint.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import threading
|
||||
from contextlib import contextmanager, suppress
|
||||
from dataclasses import dataclass, field
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable, Generator
|
||||
|
||||
|
||||
_DEFAULT_REQUEST_TIMEOUT_S = 120.0
|
||||
_MAX_REQUEST_BODY_BYTES = 64 * 1024
|
||||
_MAX_UPSTREAM_RESPONSE_BYTES = 1024 * 1024
|
||||
_MAX_WALLET_REQUESTS = 3
|
||||
_HOP_BY_HOP_HEADERS = frozenset(
|
||||
{
|
||||
"connection",
|
||||
"keep-alive",
|
||||
"proxy-authenticate",
|
||||
"proxy-authorization",
|
||||
"proxy-connection",
|
||||
"te",
|
||||
"trailer",
|
||||
"transfer-encoding",
|
||||
"upgrade",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _BridgeState:
|
||||
upstream_url: str
|
||||
authorization: str
|
||||
workspace_id: str | None
|
||||
expected_body: bytes
|
||||
path: str
|
||||
timeout: float
|
||||
response_observer: Callable[[WalletUpstreamResponse], None] | None = None
|
||||
request_count: int = 0
|
||||
lock: threading.Lock = field(default_factory=threading.Lock)
|
||||
|
||||
def claim_request(self) -> bool:
|
||||
"""Allow only the challenge probes and the one paid retry."""
|
||||
with self.lock:
|
||||
if self.request_count >= _MAX_WALLET_REQUESTS:
|
||||
return False
|
||||
self.request_count += 1
|
||||
return True
|
||||
|
||||
|
||||
class _ResponseTooLargeError(Exception):
|
||||
"""The fixed billing endpoint returned more data than a wallet needs."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WalletUpstreamResponse:
|
||||
"""A bounded upstream response observed by the trusted loopback bridge."""
|
||||
|
||||
status_code: int
|
||||
body: bytes
|
||||
|
||||
|
||||
def _bounded_response_body(response: requests.Response) -> bytes:
|
||||
content_length = response.headers.get("Content-Length")
|
||||
if content_length:
|
||||
try:
|
||||
if int(content_length) > _MAX_UPSTREAM_RESPONSE_BYTES:
|
||||
raise _ResponseTooLargeError
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
for chunk in response.iter_content(chunk_size=64 * 1024):
|
||||
if not chunk:
|
||||
continue
|
||||
total += len(chunk)
|
||||
if total > _MAX_UPSTREAM_RESPONSE_BYTES:
|
||||
raise _ResponseTooLargeError
|
||||
chunks.append(chunk)
|
||||
return b"".join(chunks)
|
||||
|
||||
|
||||
def _connection_header_names(handler: BaseHTTPRequestHandler) -> set[str]:
|
||||
value = handler.headers.get("Connection", "")
|
||||
return {item.strip().lower() for item in value.split(",") if item.strip()}
|
||||
|
||||
|
||||
def _forward_request_headers(handler: BaseHTTPRequestHandler) -> dict[str, str]:
|
||||
blocked = {
|
||||
*_HOP_BY_HOP_HEADERS,
|
||||
*_connection_header_names(handler),
|
||||
"content-length",
|
||||
"forwarded",
|
||||
"host",
|
||||
"true-client-ip",
|
||||
"x-forwarded-for",
|
||||
"x-forwarded-host",
|
||||
"x-forwarded-proto",
|
||||
"x-real-ip",
|
||||
"x-strix-authorization",
|
||||
"x-strix-workspace",
|
||||
"x-vercel-forwarded-for",
|
||||
}
|
||||
return {name: value for name, value in handler.headers.items() if name.lower() not in blocked}
|
||||
|
||||
|
||||
def _send_json_error(handler: BaseHTTPRequestHandler, status: int, message: str) -> None:
|
||||
body = f'{{"error": "{message}"}}'.encode()
|
||||
handler.close_connection = True
|
||||
handler.send_response(status)
|
||||
handler.send_header("Content-Type", "application/json")
|
||||
handler.send_header("Content-Length", str(len(body)))
|
||||
handler.send_header("Cache-Control", "no-store")
|
||||
handler.send_header("Connection", "close")
|
||||
handler.end_headers()
|
||||
with suppress(BrokenPipeError, ConnectionResetError):
|
||||
handler.wfile.write(body)
|
||||
|
||||
|
||||
def _make_handler(state: _BridgeState) -> type[BaseHTTPRequestHandler]:
|
||||
class WalletBridgeHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def log_message(self, format: str, *args: Any) -> None: # noqa: A002
|
||||
"""Do not write wallet request metadata to stderr."""
|
||||
del format, args
|
||||
|
||||
def do_POST(self) -> None: # noqa: PLR0911, PLR0912
|
||||
if self.path != state.path:
|
||||
_send_json_error(self, 404, "Not found")
|
||||
return
|
||||
if self.headers.get("Transfer-Encoding"):
|
||||
_send_json_error(self, 400, "Chunked request bodies are not supported")
|
||||
return
|
||||
try:
|
||||
content_length = int(self.headers.get("Content-Length", ""))
|
||||
except ValueError:
|
||||
_send_json_error(self, 411, "A valid Content-Length is required")
|
||||
return
|
||||
if content_length < 0 or content_length > _MAX_REQUEST_BODY_BYTES:
|
||||
_send_json_error(self, 413, "Request body is too large")
|
||||
return
|
||||
body = self.rfile.read(content_length)
|
||||
if body != state.expected_body:
|
||||
_send_json_error(self, 403, "Request body did not match the approved top-up")
|
||||
return
|
||||
if not state.claim_request():
|
||||
_send_json_error(self, 429, "Wallet request limit reached")
|
||||
return
|
||||
|
||||
headers = _forward_request_headers(self)
|
||||
headers["X-Strix-Authorization"] = state.authorization
|
||||
if state.workspace_id:
|
||||
headers["X-Strix-Workspace"] = state.workspace_id
|
||||
try:
|
||||
response = requests.request(
|
||||
"POST",
|
||||
state.upstream_url,
|
||||
headers=headers,
|
||||
data=body,
|
||||
timeout=state.timeout,
|
||||
allow_redirects=False,
|
||||
stream=True,
|
||||
)
|
||||
try:
|
||||
response_body = _bounded_response_body(response)
|
||||
response_status = response.status_code
|
||||
response_headers = dict(response.headers)
|
||||
finally:
|
||||
response.close()
|
||||
except _ResponseTooLargeError:
|
||||
_send_json_error(self, 502, "Strix billing response was too large")
|
||||
return
|
||||
except requests.RequestException:
|
||||
_send_json_error(self, 502, "Could not reach the Strix billing endpoint")
|
||||
return
|
||||
|
||||
if state.response_observer is not None:
|
||||
with suppress(Exception):
|
||||
state.response_observer(
|
||||
WalletUpstreamResponse(status_code=response_status, body=response_body)
|
||||
)
|
||||
|
||||
if 300 <= response_status < 400:
|
||||
_send_json_error(self, 502, "Strix billing refused an unexpected redirect")
|
||||
return
|
||||
|
||||
self.send_response(response_status)
|
||||
response_connection_headers = {
|
||||
item.strip().lower()
|
||||
for item in response_headers.get("Connection", "").split(",")
|
||||
if item.strip()
|
||||
}
|
||||
blocked_response_headers = {
|
||||
*_HOP_BY_HOP_HEADERS,
|
||||
*response_connection_headers,
|
||||
"cache-control",
|
||||
"content-encoding",
|
||||
"content-length",
|
||||
"location",
|
||||
}
|
||||
for name, value in response_headers.items():
|
||||
if (
|
||||
name.lower() not in blocked_response_headers
|
||||
and "\r" not in value
|
||||
and "\n" not in value
|
||||
):
|
||||
self.send_header(name, value)
|
||||
self.send_header("Content-Length", str(len(response_body)))
|
||||
self.send_header("Cache-Control", "no-store")
|
||||
self.end_headers()
|
||||
with suppress(BrokenPipeError, ConnectionResetError):
|
||||
self.wfile.write(response_body)
|
||||
|
||||
def do_GET(self) -> None:
|
||||
_send_json_error(self, 405, "Method not allowed")
|
||||
|
||||
def do_PUT(self) -> None:
|
||||
_send_json_error(self, 405, "Method not allowed")
|
||||
|
||||
def do_PATCH(self) -> None:
|
||||
_send_json_error(self, 405, "Method not allowed")
|
||||
|
||||
def do_DELETE(self) -> None:
|
||||
_send_json_error(self, 405, "Method not allowed")
|
||||
|
||||
return WalletBridgeHandler
|
||||
|
||||
|
||||
@contextmanager
|
||||
def wallet_payment_bridge(
|
||||
*,
|
||||
upstream_url: str,
|
||||
api_token: str,
|
||||
workspace_id: str | None = None,
|
||||
expected_body: bytes,
|
||||
timeout: float | None = None,
|
||||
response_observer: Callable[[WalletUpstreamResponse], None] | None = None,
|
||||
) -> Generator[str]:
|
||||
"""Yield a one-run loopback URL that injects the Strix API token upstream.
|
||||
|
||||
The random path prevents accidental cross-process requests and limits local
|
||||
denial-of-service races. It is not an authentication boundary against a
|
||||
same-user process that can inspect another process's argv.
|
||||
"""
|
||||
capability = secrets.token_urlsafe(32)
|
||||
path = f"/topup/{capability}"
|
||||
state = _BridgeState(
|
||||
upstream_url=upstream_url,
|
||||
authorization=f"Bearer {api_token}",
|
||||
workspace_id=workspace_id,
|
||||
expected_body=expected_body,
|
||||
path=path,
|
||||
timeout=timeout or _DEFAULT_REQUEST_TIMEOUT_S,
|
||||
response_observer=response_observer,
|
||||
)
|
||||
server = ThreadingHTTPServer(("127.0.0.1", 0), _make_handler(state))
|
||||
server.daemon_threads = True
|
||||
thread = threading.Thread(
|
||||
target=server.serve_forever,
|
||||
kwargs={"poll_interval": 0.05},
|
||||
name="strix-wallet-bridge",
|
||||
daemon=True,
|
||||
)
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_port}{path}"
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=1)
|
||||
1759
strix/interface/cloud/render.py
Normal file
1759
strix/interface/cloud/render.py
Normal file
File diff suppressed because it is too large
Load Diff
1128
strix/interface/cloud/runner.py
Normal file
1128
strix/interface/cloud/runner.py
Normal file
File diff suppressed because it is too large
Load Diff
167
strix/interface/cloud/session.py
Normal file
167
strix/interface/cloud/session.py
Normal file
@@ -0,0 +1,167 @@
|
||||
"""Inspect and safely narrow a managed Strix CLI session."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
from strix.interface.cloud.arguments import CloudArgumentParser
|
||||
from strix.interface.cloud.render import emit, json_mode
|
||||
from strix.interface.platform_cli import read_record, save_record
|
||||
from strix.interface.terminal_text import sanitize_terminal_text
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import argparse
|
||||
|
||||
|
||||
def run_session(argv: list[str]) -> int:
|
||||
console = Console()
|
||||
normalized = ["show", *argv] if not argv or argv[0].startswith("-") else list(argv)
|
||||
if normalized[0] == "help":
|
||||
normalized = ["--help", *normalized[1:]]
|
||||
if normalized[0] in {"-h", "--help"}:
|
||||
_print_help(console)
|
||||
return 0
|
||||
verb = normalized.pop(0)
|
||||
if verb == "scopes" and normalized and normalized[0] == "set":
|
||||
normalized.pop(0)
|
||||
return _run_scopes_set(console, normalized)
|
||||
if verb not in {"show", "scopes"}:
|
||||
console.print(f"[red]Unknown session command:[/] {escape(sanitize_terminal_text(verb))}")
|
||||
_print_help(console)
|
||||
return http.EXIT_USAGE
|
||||
return _run_show(console, normalized, scopes_only=verb == "scopes")
|
||||
|
||||
|
||||
def _common(parser: argparse.ArgumentParser) -> None:
|
||||
parser.add_argument("--json", action="store_true", help="Print the raw JSON response.")
|
||||
parser.add_argument("--show-scopes", action="store_true", help="Print every granted scope.")
|
||||
parser.add_argument("--token", default=None, help="API token override.")
|
||||
parser.add_argument("--workspace-id", default=None, metavar="ORG_ID")
|
||||
parser.add_argument("--app-url", default=None, metavar="URL")
|
||||
parser.add_argument("--timeout", default=None, type=float, metavar="SECONDS")
|
||||
|
||||
|
||||
def _configure(args: argparse.Namespace) -> bool:
|
||||
external = args.token is not None or bool(os.environ.get("STRIX_API_TOKEN", "").strip())
|
||||
http.configure(
|
||||
base_url=args.app_url,
|
||||
timeout=args.timeout,
|
||||
token_override=bool(args.token),
|
||||
workspace_id=args.workspace_id,
|
||||
)
|
||||
return external
|
||||
|
||||
|
||||
def _run_show(console: Console, argv: list[str], *, scopes_only: bool) -> int:
|
||||
parser = CloudArgumentParser(prog=f"strix cloud session {'scopes' if scopes_only else 'show'}")
|
||||
_common(parser)
|
||||
as_json = json_mode(flag="--json" in argv)
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
_configure(args)
|
||||
payload = http.check(http.request("GET", "/cli/session", token=args.token))
|
||||
except SystemExit as exc:
|
||||
return int(exc.code or 0)
|
||||
except http.CloudError as exc:
|
||||
return _error(console, exc, as_json=as_json)
|
||||
if not isinstance(payload, dict):
|
||||
return _error(console, http.CloudError("invalid CLI session response"), as_json=as_json)
|
||||
record = cast("dict[str, Any]", payload)
|
||||
if as_json:
|
||||
emit(console, record, as_json=True)
|
||||
return http.EXIT_OK
|
||||
scopes = _string_list(record.get("scopes"))
|
||||
ceiling = _string_list(record.get("scope_ceiling"))
|
||||
profile = str(record.get("scope_profile") or "custom").title()
|
||||
if not scopes_only:
|
||||
device_name = escape(str(record.get("device_name") or "this device"))
|
||||
console.print(f"[green]Active CLI session[/] on [bold]{device_name}[/]")
|
||||
console.print(f" Workspace: {escape(str(record.get('organization_id') or 'unknown'))}")
|
||||
console.print(f" Access: {profile} · {len(scopes)} scopes granted · {len(ceiling)} maximum")
|
||||
if args.show_scopes or scopes_only:
|
||||
console.print(f" Granted: [dim]{escape(' '.join(scopes))}[/]")
|
||||
console.print(f" Ceiling: [dim]{escape(' '.join(ceiling))}[/]")
|
||||
return http.EXIT_OK
|
||||
|
||||
|
||||
def _run_scopes_set(console: Console, argv: list[str]) -> int:
|
||||
parser = CloudArgumentParser(
|
||||
prog="strix cloud session scopes set",
|
||||
description="Change scopes within the access approved at browser sign-in.",
|
||||
)
|
||||
mode = parser.add_mutually_exclusive_group(required=True)
|
||||
mode.add_argument("profile", nargs="?", choices=("minimal", "recommended", "full"))
|
||||
mode.add_argument("--scopes", nargs="+", metavar="SCOPE")
|
||||
_common(parser)
|
||||
as_json = json_mode(flag="--json" in argv)
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
external = _configure(args)
|
||||
body = (
|
||||
{"scope_profile": args.profile}
|
||||
if args.profile
|
||||
else {"scope_profile": "custom", "scopes": args.scopes}
|
||||
)
|
||||
payload = http.check(http.request("PATCH", "/cli/session", token=args.token, body=body))
|
||||
except SystemExit as exc:
|
||||
return int(exc.code or 0)
|
||||
except http.CloudError as exc:
|
||||
return _error(console, exc, as_json=as_json)
|
||||
if not isinstance(payload, dict):
|
||||
return _error(console, http.CloudError("invalid CLI session response"), as_json=as_json)
|
||||
result = cast("dict[str, Any]", payload)
|
||||
if not external:
|
||||
stored = read_record()
|
||||
if stored is not None:
|
||||
stored.update(
|
||||
{
|
||||
key: result[key]
|
||||
for key in ("scopes", "requested_scopes", "scope_ceiling", "scope_profile")
|
||||
if key in result
|
||||
}
|
||||
)
|
||||
save_record(stored)
|
||||
if as_json:
|
||||
emit(console, result, as_json=True)
|
||||
else:
|
||||
scopes = _string_list(result.get("scopes"))
|
||||
profile = str(result.get("scope_profile") or "custom").title()
|
||||
console.print(f"[green]✓ CLI access updated.[/] {profile} · {len(scopes)} scopes granted")
|
||||
if args.show_scopes:
|
||||
console.print(f" Scopes: [dim]{escape(' '.join(scopes))}[/]")
|
||||
return http.EXIT_OK
|
||||
|
||||
|
||||
def _string_list(value: Any) -> list[str]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
items = cast("list[Any]", cast("Any", value))
|
||||
return [str(item) for item in items]
|
||||
|
||||
|
||||
def _error(console: Console, error: http.CloudError, *, as_json: bool) -> int:
|
||||
if as_json:
|
||||
raw_payload: Any = error.payload
|
||||
error_payload = cast("dict[str, Any]", raw_payload)
|
||||
payload = dict(error_payload) if isinstance(raw_payload, dict) else {}
|
||||
payload["error"] = str(error)
|
||||
if payload.get("detail") == payload.get("error"):
|
||||
payload.pop("detail", None)
|
||||
emit(console, payload, as_json=True)
|
||||
else:
|
||||
console.print(f"[red]Error:[/] {escape(sanitize_terminal_text(error))}")
|
||||
return error.exit_code
|
||||
|
||||
|
||||
def _print_help(console: Console) -> None:
|
||||
console.print("[bold]strix cloud session[/] commands:")
|
||||
console.print(" show Show the remote CLI session (default).")
|
||||
console.print(" scopes Show granted scopes and consent ceiling.")
|
||||
console.print(" scopes set PROFILE Use minimal, recommended, or full.")
|
||||
console.print(" scopes set --scopes SCOPE… Use a custom set within the ceiling.")
|
||||
403
strix/interface/cloud/source_scan.py
Normal file
403
strix/interface/cloud/source_scan.py
Normal file
@@ -0,0 +1,403 @@
|
||||
"""Local-source approval, upload, and scan-launch lifecycle."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from urllib.parse import quote
|
||||
|
||||
from rich.markup import escape
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
from strix.interface.cloud.render import emit
|
||||
from strix.interface.cloud.source_upload import prepare_source, remove_bundle
|
||||
from strix.interface.terminal_text import sanitize_terminal_text
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import argparse
|
||||
from typing import NoReturn
|
||||
|
||||
from rich.console import Console
|
||||
|
||||
from strix.interface.cloud.source_upload import SourceBundle
|
||||
|
||||
|
||||
_SHA256 = re.compile(r"^[0-9a-fA-F]{64}$")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalSourceScan:
|
||||
"""Own one local bundle and its staged upload through a scan launch."""
|
||||
|
||||
bundle: SourceBundle | None = None
|
||||
upload_id: str | None = None
|
||||
idempotency_key: str | None = None
|
||||
_launch_started: bool = False
|
||||
|
||||
def prepare_and_attach(
|
||||
self,
|
||||
console: Console,
|
||||
args: argparse.Namespace,
|
||||
body: dict[str, Any],
|
||||
*,
|
||||
as_json: bool,
|
||||
token: str | None,
|
||||
) -> bool:
|
||||
"""Prepare source, emit a dry run, or upload and attach it to ``body``.
|
||||
|
||||
Returns ``True`` when a dry run was emitted and request execution should stop.
|
||||
"""
|
||||
self.bundle = prepare_scan_source(console, args, as_json=as_json)
|
||||
if self.bundle is None:
|
||||
return False
|
||||
if getattr(args, "dry_run", False):
|
||||
emit(
|
||||
console,
|
||||
{"source": self.bundle.summary(show_files=getattr(args, "show_files", False))},
|
||||
as_json=as_json,
|
||||
view="source_manifest",
|
||||
)
|
||||
return True
|
||||
self.upload_id = _upload_scan_source(self.bundle, token=token)
|
||||
existing = body.get("upload_ids")
|
||||
body["upload_ids"] = [
|
||||
*(existing if isinstance(existing, list) else []),
|
||||
self.upload_id,
|
||||
]
|
||||
return False
|
||||
|
||||
def mark_launch_started(self) -> None:
|
||||
"""Record that the scan-creation request may have reached the platform."""
|
||||
self._launch_started = self.upload_id is not None
|
||||
|
||||
def handle_request_failure(self, error: BaseException, *, token: str | None) -> None:
|
||||
"""Clean or retain a staged upload according to request ambiguity."""
|
||||
if self.upload_id is None:
|
||||
return
|
||||
if self._launch_started:
|
||||
if isinstance(error, KeyboardInterrupt):
|
||||
raise _interrupted_source_upload_error(
|
||||
self.upload_id, self.idempotency_key
|
||||
) from None
|
||||
if isinstance(error, Exception):
|
||||
raise _retained_source_upload_error(
|
||||
self.upload_id, error, self.idempotency_key
|
||||
) from error
|
||||
return
|
||||
try:
|
||||
_delete_upload(self.upload_id, token=token)
|
||||
except (http.CloudError, KeyboardInterrupt) as cleanup_error:
|
||||
if isinstance(error, Exception):
|
||||
raise _source_cleanup_error(self.upload_id, error, cleanup_error) from error
|
||||
interrupted = http.CloudError("source upload interrupted.", exit_code=130)
|
||||
raise _source_cleanup_error(self.upload_id, interrupted, cleanup_error) from None
|
||||
|
||||
def handle_response_failure(
|
||||
self,
|
||||
error: BaseException,
|
||||
*,
|
||||
definitive: bool,
|
||||
token: str | None,
|
||||
) -> None:
|
||||
"""Clean a rejected upload or retain one whose scan result is ambiguous."""
|
||||
if self.upload_id is None:
|
||||
return
|
||||
if definitive:
|
||||
try:
|
||||
_delete_upload(self.upload_id, token=token)
|
||||
except (http.CloudError, KeyboardInterrupt) as cleanup_error:
|
||||
if isinstance(error, Exception):
|
||||
raise _source_cleanup_error(self.upload_id, error, cleanup_error) from error
|
||||
interrupted = http.CloudError("source upload interrupted.", exit_code=130)
|
||||
raise _source_cleanup_error(self.upload_id, interrupted, cleanup_error) from None
|
||||
return
|
||||
if isinstance(error, Exception):
|
||||
raise _retained_source_upload_error(
|
||||
self.upload_id, error, self.idempotency_key
|
||||
) from error
|
||||
|
||||
def wrap_result(self, result: Any, args: argparse.Namespace) -> Any:
|
||||
"""Attach the approved source manifest to a successful scan response."""
|
||||
if self.bundle is None:
|
||||
return result
|
||||
return {
|
||||
"source": self.bundle.summary(show_files=getattr(args, "show_files", False)),
|
||||
"upload_id": self.upload_id,
|
||||
"scan": result,
|
||||
}
|
||||
|
||||
def close(self) -> None:
|
||||
"""Remove the private temporary bundle, if one was built."""
|
||||
if self.bundle is not None:
|
||||
remove_bundle(self.bundle)
|
||||
|
||||
|
||||
def prepare_scan_source(
|
||||
console: Console, args: argparse.Namespace, *, as_json: bool
|
||||
) -> SourceBundle | None:
|
||||
"""Build and approve the exact local-source snapshot for one invocation."""
|
||||
source = getattr(args, "source", None)
|
||||
source_flags = (
|
||||
"dry_run",
|
||||
"show_files",
|
||||
"include_hidden",
|
||||
"include_sensitive",
|
||||
"include_archives",
|
||||
"approve_sha256",
|
||||
)
|
||||
if source is None:
|
||||
if any(getattr(args, name, False) for name in source_flags) or getattr(args, "exclude", []):
|
||||
raise http.CloudError("source upload options require --source DIRECTORY.")
|
||||
return None
|
||||
bundle = prepare_source(
|
||||
source,
|
||||
include_hidden=bool(getattr(args, "include_hidden", False)),
|
||||
include_sensitive=bool(getattr(args, "include_sensitive", False)),
|
||||
include_archives=bool(getattr(args, "include_archives", False)),
|
||||
exclude=cast("list[str]", getattr(args, "exclude", [])),
|
||||
)
|
||||
keep_bundle = False
|
||||
try:
|
||||
approved_digest = _validate_source_digest_approval(args, bundle)
|
||||
if getattr(args, "dry_run", False):
|
||||
keep_bundle = True
|
||||
return bundle
|
||||
if getattr(args, "yes", False) or approved_digest is not None:
|
||||
keep_bundle = True
|
||||
return bundle
|
||||
if as_json or not (sys.stdin.isatty() and sys.stdout.isatty()):
|
||||
_source_approval_error(
|
||||
"source upload requires explicit approval in non-interactive mode. "
|
||||
"Review with --dry-run --show-files, then rerun with "
|
||||
"--approve-sha256 <reviewed hash>; use --yes only for a deliberate "
|
||||
"one-shot approval of the snapshot built by that invocation."
|
||||
)
|
||||
console.print(
|
||||
"[bold]Local source upload[/]\n"
|
||||
f" {len(bundle.manifest.files):,} file(s), "
|
||||
f"{_format_bytes(bundle.manifest.total_bytes)} "
|
||||
f"({_format_bytes(bundle.archive_bytes)} compressed)\n"
|
||||
f" {sum(bundle.manifest.excluded.values()):,} path(s) excluded\n"
|
||||
" Only the selected files will be sent to Strix Cloud."
|
||||
)
|
||||
if getattr(args, "show_files", False):
|
||||
console.print(f"\n[bold]Selected files ({len(bundle.manifest.files):,})[/]")
|
||||
for selected in bundle.manifest.files:
|
||||
console.print(
|
||||
f" {escape(sanitize_terminal_text(selected.archive_name))}", soft_wrap=True
|
||||
)
|
||||
answer = (
|
||||
console.input("Upload this source and start the scan? [y/N]: ", markup=False)
|
||||
.strip()
|
||||
.lower()
|
||||
)
|
||||
if answer not in ("y", "yes"):
|
||||
_source_approval_error("source upload cancelled.")
|
||||
keep_bundle = True
|
||||
return bundle
|
||||
finally:
|
||||
if not keep_bundle:
|
||||
remove_bundle(bundle)
|
||||
|
||||
|
||||
def _validate_source_digest_approval(args: argparse.Namespace, bundle: SourceBundle) -> str | None:
|
||||
approved_digest = getattr(args, "approve_sha256", None)
|
||||
if approved_digest is None:
|
||||
return None
|
||||
if not isinstance(approved_digest, str) or not _SHA256.fullmatch(approved_digest):
|
||||
_source_approval_error("--approve-sha256 must be exactly 64 hexadecimal characters.")
|
||||
if bundle.archive_sha256 != approved_digest.lower():
|
||||
_source_approval_error(
|
||||
"source archive SHA-256 does not match --approve-sha256; review a fresh "
|
||||
"--dry-run before uploading."
|
||||
)
|
||||
return approved_digest
|
||||
|
||||
|
||||
def _source_approval_error(message: str) -> NoReturn:
|
||||
raise http.CloudError(message)
|
||||
|
||||
|
||||
def _upload_scan_source(bundle: SourceBundle, *, token: str | None) -> str:
|
||||
file_name = f"strix-source-{bundle.archive_sha256[:12]}.zip"
|
||||
requested = http.check(
|
||||
http.request(
|
||||
"POST",
|
||||
"/uploads/request",
|
||||
token=token,
|
||||
body={
|
||||
"file_name": file_name,
|
||||
"file_size": bundle.archive_bytes,
|
||||
"category": "repository",
|
||||
},
|
||||
)
|
||||
)
|
||||
if not isinstance(requested, dict):
|
||||
raise http.CloudError("the platform returned an invalid source upload response.")
|
||||
fields = cast("dict[str, Any]", requested)
|
||||
upload_id = fields.get("upload_id")
|
||||
signed_url = fields.get("signed_url")
|
||||
upload_token = fields.get("token")
|
||||
if not all(isinstance(value, str) and value for value in (upload_id, signed_url, upload_token)):
|
||||
error = http.CloudError("the platform did not return complete source upload credentials.")
|
||||
if isinstance(upload_id, str) and upload_id:
|
||||
try:
|
||||
_delete_upload(upload_id, token=token)
|
||||
except (http.CloudError, KeyboardInterrupt) as cleanup_error:
|
||||
raise _source_cleanup_error(upload_id, error, cleanup_error) from error
|
||||
raise error
|
||||
try:
|
||||
http.upload_file(cast("str", signed_url), cast("str", upload_token), bundle.archive_path)
|
||||
completed = http.check(
|
||||
http.request(
|
||||
"POST",
|
||||
"/uploads/complete",
|
||||
token=token,
|
||||
body={"upload_id": upload_id},
|
||||
)
|
||||
)
|
||||
_validate_completed_upload(completed, expected_id=cast("str", upload_id))
|
||||
except BaseException as error:
|
||||
try:
|
||||
_delete_upload(cast("str", upload_id), token=token)
|
||||
except (http.CloudError, KeyboardInterrupt) as cleanup_error:
|
||||
if isinstance(error, Exception):
|
||||
raise _source_cleanup_error(cast("str", upload_id), error, cleanup_error) from error
|
||||
interrupted = http.CloudError("source upload interrupted.", exit_code=130)
|
||||
raise _source_cleanup_error(
|
||||
cast("str", upload_id), interrupted, cleanup_error
|
||||
) from None
|
||||
raise
|
||||
return cast("str", upload_id)
|
||||
|
||||
|
||||
def _validate_completed_upload(completed: Any, *, expected_id: str) -> None:
|
||||
fields = cast("dict[str, Any]", completed) if isinstance(completed, dict) else {}
|
||||
if fields.get("id") != expected_id:
|
||||
raise http.CloudError("the platform returned an invalid source upload completion response.")
|
||||
|
||||
|
||||
def _delete_upload(upload_id: str, *, token: str | None) -> None:
|
||||
response = http.request("DELETE", f"/uploads/{quote(upload_id, safe='')}", token=token)
|
||||
if response.status_code == 404 or 200 <= response.status_code < 300:
|
||||
return
|
||||
http.check(response)
|
||||
|
||||
|
||||
def _source_cleanup_note(upload_id: str, cleanup_error: BaseException) -> str:
|
||||
return (
|
||||
f"Cleanup of source upload {upload_id} could not be confirmed: {cleanup_error}. "
|
||||
f"Retry with `strix cloud uploads delete {upload_id}`."
|
||||
)
|
||||
|
||||
|
||||
def _source_cleanup_error(
|
||||
upload_id: str, error: Exception, cleanup_error: BaseException
|
||||
) -> http.CloudError:
|
||||
"""Report a staged source object whenever automatic deletion is uncertain."""
|
||||
message = f"{error} {_source_cleanup_note(upload_id, cleanup_error)}"
|
||||
payload: dict[str, Any] = {}
|
||||
exit_code = http.EXIT_ERROR
|
||||
if isinstance(error, http.CloudError):
|
||||
exit_code = error.exit_code
|
||||
raw_payload: Any = error.payload
|
||||
if isinstance(raw_payload, dict):
|
||||
payload.update(cast("dict[str, Any]", raw_payload))
|
||||
elif raw_payload is not None:
|
||||
payload["detail"] = raw_payload
|
||||
payload.update(
|
||||
{
|
||||
"error": message,
|
||||
"upload_id": upload_id,
|
||||
"upload_retained": True,
|
||||
"cleanup_unknown": True,
|
||||
}
|
||||
)
|
||||
return http.CloudError(message, exit_code=exit_code, payload=payload)
|
||||
|
||||
|
||||
def _interrupted_source_upload_error(
|
||||
upload_id: str, idempotency_key: str | None = None
|
||||
) -> http.CloudError:
|
||||
retry_note = _idempotency_retry_note(idempotency_key)
|
||||
message = (
|
||||
"Interrupted while starting the scan. The launch outcome is unknown, so source upload "
|
||||
f"{upload_id} was retained. Check `strix cloud scans list` before retrying; if no scan "
|
||||
f"was created, run `strix cloud uploads delete {upload_id}`.{retry_note}"
|
||||
)
|
||||
payload: dict[str, Any] = {
|
||||
"error": message,
|
||||
"interrupted": True,
|
||||
"upload_id": upload_id,
|
||||
"upload_retained": True,
|
||||
"launch_outcome_unknown": True,
|
||||
}
|
||||
_attach_idempotency_recovery(payload, idempotency_key)
|
||||
return http.CloudError(message, exit_code=130, payload=payload)
|
||||
|
||||
|
||||
def _retained_source_upload_error(
|
||||
upload_id: str,
|
||||
error: Exception,
|
||||
idempotency_key: str | None = None,
|
||||
) -> http.CloudError:
|
||||
"""Preserve source when the platform may already have accepted its scan."""
|
||||
retry_note = _idempotency_retry_note(idempotency_key)
|
||||
message = (
|
||||
f"{error} The scan launch outcome is unknown, so source upload {upload_id} was retained. "
|
||||
"Check `strix cloud scans list` before retrying; if no scan was created, clean it up "
|
||||
f"with `strix cloud uploads delete {upload_id}`. Linked uploads cannot be deleted."
|
||||
f"{retry_note}"
|
||||
)
|
||||
payload: dict[str, Any] = {}
|
||||
exit_code = http.EXIT_ERROR
|
||||
if isinstance(error, http.CloudError):
|
||||
exit_code = error.exit_code
|
||||
raw_payload: Any = error.payload
|
||||
error_payload = cast("dict[str, Any]", raw_payload)
|
||||
if isinstance(raw_payload, dict):
|
||||
payload.update(error_payload)
|
||||
elif raw_payload is not None:
|
||||
payload["detail"] = raw_payload
|
||||
payload.update(
|
||||
{
|
||||
"error": message,
|
||||
"upload_id": upload_id,
|
||||
"upload_retained": True,
|
||||
"launch_outcome_unknown": True,
|
||||
}
|
||||
)
|
||||
_attach_idempotency_recovery(payload, idempotency_key)
|
||||
return http.CloudError(message, exit_code=exit_code, payload=payload)
|
||||
|
||||
|
||||
def _idempotency_retry_note(idempotency_key: str | None) -> str:
|
||||
if not idempotency_key:
|
||||
return ""
|
||||
return (
|
||||
" An exact retry is safe only with the same request body and "
|
||||
f"`--idempotency-key {idempotency_key}`."
|
||||
)
|
||||
|
||||
|
||||
def _attach_idempotency_recovery(payload: dict[str, Any], idempotency_key: str | None) -> None:
|
||||
if not idempotency_key:
|
||||
return
|
||||
payload.update(
|
||||
{
|
||||
"idempotency_key": idempotency_key,
|
||||
"retry_safe": True,
|
||||
"retry_same_request": True,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _format_bytes(value: int) -> str:
|
||||
if value < 1024:
|
||||
return f"{value} B"
|
||||
if value < 1024 * 1024:
|
||||
return f"{value / 1024:.1f} KB"
|
||||
return f"{value / (1024 * 1024):.1f} MB"
|
||||
703
strix/interface/cloud/source_upload.py
Normal file
703
strix/interface/cloud/source_upload.py
Normal file
@@ -0,0 +1,703 @@
|
||||
"""Privacy-conscious local source packaging for managed scans."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import fnmatch
|
||||
import hashlib
|
||||
import os
|
||||
import shutil
|
||||
import stat
|
||||
import subprocess # nosec B404
|
||||
import tempfile
|
||||
import zipfile
|
||||
from collections import Counter
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
from typing import Protocol
|
||||
|
||||
class _ScandirIterator(Iterator[os.DirEntry[str]], Protocol):
|
||||
def close(self) -> None: ...
|
||||
|
||||
|
||||
MAX_FILES = 20_000
|
||||
MAX_FILE_BYTES = 25 * 1024 * 1024
|
||||
MAX_TOTAL_BYTES = 250 * 1024 * 1024
|
||||
MAX_ARCHIVE_BYTES = 50 * 1024 * 1024
|
||||
MAX_CANDIDATE_PATHS = 200_000
|
||||
MAX_IGNORE_BYTES = 64 * 1024
|
||||
MAX_IGNORE_PATTERNS = 1_000
|
||||
MAX_IGNORE_PATTERN_CHARS = 1_024
|
||||
|
||||
_ALWAYS_EXCLUDED_DIRS = frozenset(
|
||||
{
|
||||
".git",
|
||||
".hg",
|
||||
".svn",
|
||||
"node_modules",
|
||||
"vendor",
|
||||
"venv",
|
||||
".venv",
|
||||
"env",
|
||||
"__pycache__",
|
||||
".tox",
|
||||
".pytest_cache",
|
||||
".mypy_cache",
|
||||
".ruff_cache",
|
||||
"dist",
|
||||
"build",
|
||||
"coverage",
|
||||
"target",
|
||||
".next",
|
||||
".nuxt",
|
||||
".gradle",
|
||||
}
|
||||
)
|
||||
_SENSITIVE_NAMES = frozenset(
|
||||
{
|
||||
"id_rsa",
|
||||
"id_dsa",
|
||||
"id_ecdsa",
|
||||
"id_ed25519",
|
||||
"credentials.json",
|
||||
"service-account.json",
|
||||
"service_account.json",
|
||||
".env",
|
||||
".npmrc",
|
||||
".pypirc",
|
||||
".netrc",
|
||||
".git-credentials",
|
||||
"application_default_credentials.json",
|
||||
}
|
||||
)
|
||||
_SENSITIVE_PATTERNS = (
|
||||
"*.pem",
|
||||
"*.key",
|
||||
"*.p12",
|
||||
"*.pfx",
|
||||
"*.keystore",
|
||||
"*.jks",
|
||||
"secrets.*",
|
||||
"secret.*",
|
||||
".env.*",
|
||||
)
|
||||
_SENSITIVE_PATH_SUFFIXES = (
|
||||
(".aws", "credentials"),
|
||||
(".aws", "config"),
|
||||
(".docker", "config.json"),
|
||||
(".config", "gcloud", "credentials.db"),
|
||||
(".azure", "accesstokens.json"),
|
||||
(".azure", "azureprofile.json"),
|
||||
(".kube", "config"),
|
||||
)
|
||||
_ARCHIVE_SUFFIXES = (
|
||||
".zip",
|
||||
".tar",
|
||||
".tgz",
|
||||
".tar.gz",
|
||||
".tar.bz2",
|
||||
".tar.xz",
|
||||
".7z",
|
||||
".rar",
|
||||
".gz",
|
||||
".bz2",
|
||||
".xz",
|
||||
".jar",
|
||||
".war",
|
||||
".whl",
|
||||
".nupkg",
|
||||
".apk",
|
||||
".ipa",
|
||||
)
|
||||
_ARCHIVE_MAGIC_PREFIXES = (
|
||||
b"PK\x03\x04",
|
||||
b"PK\x05\x06",
|
||||
b"PK\x07\x08",
|
||||
b"\x1f\x8b",
|
||||
b"BZh",
|
||||
b"\xfd7zXZ\x00",
|
||||
b"7z\xbc\xaf\x27\x1c",
|
||||
b"Rar!\x1a\x07",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SelectedFile:
|
||||
path: Path
|
||||
archive_name: str
|
||||
size: int
|
||||
device: int
|
||||
inode: int
|
||||
mtime_ns: int
|
||||
ctime_ns: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceManifest:
|
||||
source: Path
|
||||
files: tuple[SelectedFile, ...]
|
||||
excluded: Counter[str]
|
||||
include_hidden: bool
|
||||
include_sensitive: bool
|
||||
include_archives: bool
|
||||
|
||||
@property
|
||||
def total_bytes(self) -> int:
|
||||
return sum(item.size for item in self.files)
|
||||
|
||||
def as_dict(
|
||||
self,
|
||||
*,
|
||||
show_files: bool,
|
||||
archive_bytes: int | None = None,
|
||||
archive_sha256: str | None = None,
|
||||
) -> dict[str, object]:
|
||||
result: dict[str, object] = {
|
||||
"source": str(self.source),
|
||||
"file_count": len(self.files),
|
||||
"uncompressed_bytes": self.total_bytes,
|
||||
"excluded_count": sum(self.excluded.values()),
|
||||
"excluded_by_reason": dict(sorted(self.excluded.items())),
|
||||
"include_hidden": self.include_hidden,
|
||||
"include_sensitive": self.include_sensitive,
|
||||
"include_archives": self.include_archives,
|
||||
}
|
||||
if archive_bytes is not None:
|
||||
result["archive_bytes"] = archive_bytes
|
||||
if archive_sha256 is not None:
|
||||
result["archive_sha256"] = archive_sha256
|
||||
if show_files:
|
||||
result["files"] = [item.archive_name for item in self.files]
|
||||
return result
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceBundle:
|
||||
manifest: SourceManifest
|
||||
archive_path: Path
|
||||
archive_bytes: int
|
||||
archive_sha256: str
|
||||
|
||||
def summary(self, *, show_files: bool) -> dict[str, object]:
|
||||
return self.manifest.as_dict(
|
||||
show_files=show_files,
|
||||
archive_bytes=self.archive_bytes,
|
||||
archive_sha256=self.archive_sha256,
|
||||
)
|
||||
|
||||
|
||||
def prepare_source(
|
||||
value: str,
|
||||
*,
|
||||
include_hidden: bool,
|
||||
include_sensitive: bool,
|
||||
include_archives: bool,
|
||||
exclude: list[str],
|
||||
) -> SourceBundle:
|
||||
"""Select safe source files and build a bounded temporary ZIP archive."""
|
||||
source = Path(value).expanduser().resolve()
|
||||
if not source.is_dir():
|
||||
raise http.CloudError(f"--source must be a directory: {source}")
|
||||
manifest = select_source(
|
||||
source,
|
||||
include_hidden=include_hidden,
|
||||
include_sensitive=include_sensitive,
|
||||
include_archives=include_archives,
|
||||
exclude=exclude,
|
||||
)
|
||||
if not manifest.files:
|
||||
raise http.CloudError("no files remain after applying source upload exclusions.")
|
||||
|
||||
with tempfile.NamedTemporaryFile(prefix="strix-source-", suffix=".zip", delete=False) as handle:
|
||||
archive_path = Path(handle.name)
|
||||
try:
|
||||
_write_archive(archive_path, manifest.files)
|
||||
except BaseException:
|
||||
archive_path.unlink(missing_ok=True)
|
||||
raise
|
||||
archive_bytes = archive_path.stat().st_size
|
||||
if archive_bytes > MAX_ARCHIVE_BYTES:
|
||||
archive_path.unlink(missing_ok=True)
|
||||
raise http.CloudError(
|
||||
"source archive is larger than the 50 MB upload limit; narrow --source or "
|
||||
"add --exclude patterns."
|
||||
)
|
||||
digest = _sha256(archive_path)
|
||||
return SourceBundle(manifest, archive_path, archive_bytes, digest)
|
||||
|
||||
|
||||
def select_source(
|
||||
source: Path,
|
||||
*,
|
||||
include_hidden: bool = False,
|
||||
include_sensitive: bool = False,
|
||||
include_archives: bool = False,
|
||||
exclude: list[str] | None = None,
|
||||
) -> SourceManifest:
|
||||
excluded: Counter[str] = Counter()
|
||||
selected: list[SelectedFile] = []
|
||||
patterns = [*_load_ignore_patterns(source), *(exclude or [])]
|
||||
_validate_patterns(patterns)
|
||||
total_bytes = 0
|
||||
for relative in _candidate_paths(
|
||||
source,
|
||||
include_hidden=include_hidden,
|
||||
patterns=patterns,
|
||||
excluded=excluded,
|
||||
):
|
||||
archive_name = relative.as_posix()
|
||||
reason = _exclusion_reason(
|
||||
relative,
|
||||
include_hidden=include_hidden,
|
||||
include_sensitive=include_sensitive,
|
||||
include_archives=include_archives,
|
||||
patterns=patterns,
|
||||
)
|
||||
if reason:
|
||||
excluded[reason] += 1
|
||||
continue
|
||||
path = source / relative
|
||||
try:
|
||||
info = path.lstat()
|
||||
except OSError:
|
||||
excluded["unreadable"] += 1
|
||||
continue
|
||||
if not stat.S_ISREG(info.st_mode):
|
||||
excluded["symlink_or_non_file"] += 1
|
||||
continue
|
||||
if not include_archives and _has_archive_magic(path):
|
||||
excluded["nested_archive"] += 1
|
||||
continue
|
||||
if info.st_size > MAX_FILE_BYTES:
|
||||
raise http.CloudError(
|
||||
f"{archive_name} is larger than the 25 MB per-file limit; exclude it explicitly."
|
||||
)
|
||||
selected.append(
|
||||
SelectedFile(
|
||||
path=path,
|
||||
archive_name=archive_name,
|
||||
size=info.st_size,
|
||||
device=info.st_dev,
|
||||
inode=info.st_ino,
|
||||
mtime_ns=info.st_mtime_ns,
|
||||
ctime_ns=info.st_ctime_ns,
|
||||
)
|
||||
)
|
||||
total_bytes += info.st_size
|
||||
if len(selected) > MAX_FILES:
|
||||
raise http.CloudError(
|
||||
f"source contains more than {MAX_FILES:,} files; narrow --source or add exclusions."
|
||||
)
|
||||
if total_bytes > MAX_TOTAL_BYTES:
|
||||
raise http.CloudError(
|
||||
"selected source is larger than the 250 MB expanded-size limit; narrow --source "
|
||||
"or add --exclude patterns."
|
||||
)
|
||||
selected.sort(key=lambda item: item.archive_name)
|
||||
return SourceManifest(
|
||||
source,
|
||||
tuple(selected),
|
||||
excluded,
|
||||
include_hidden,
|
||||
include_sensitive,
|
||||
include_archives,
|
||||
)
|
||||
|
||||
|
||||
def remove_bundle(bundle: SourceBundle) -> None:
|
||||
bundle.archive_path.unlink(missing_ok=True)
|
||||
|
||||
|
||||
def _candidate_paths(
|
||||
source: Path,
|
||||
*,
|
||||
include_hidden: bool,
|
||||
patterns: list[str],
|
||||
excluded: Counter[str],
|
||||
) -> Iterator[Path]:
|
||||
git_root = _git_root(source)
|
||||
if git_root is not None:
|
||||
git = shutil.which("git")
|
||||
if git is not None:
|
||||
yield from _git_candidate_paths(git, git_root, source)
|
||||
return
|
||||
yield from _walk_candidate_paths(
|
||||
source,
|
||||
include_hidden=include_hidden,
|
||||
patterns=patterns,
|
||||
excluded=excluded,
|
||||
)
|
||||
|
||||
|
||||
def _git_candidate_paths(git: str, git_root: Path, source: Path) -> Iterator[Path]:
|
||||
"""Stream Git's NUL-delimited manifest without buffering an unbounded repository."""
|
||||
relative_source = source.relative_to(git_root)
|
||||
command = [
|
||||
git,
|
||||
"-C",
|
||||
str(git_root),
|
||||
"ls-files",
|
||||
"-z",
|
||||
"--cached",
|
||||
"--others",
|
||||
"--exclude-standard",
|
||||
"--",
|
||||
]
|
||||
if relative_source != Path():
|
||||
command.append(relative_source.as_posix())
|
||||
try:
|
||||
process = subprocess.Popen( # noqa: S603 # nosec B603
|
||||
command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.DEVNULL,
|
||||
)
|
||||
except OSError as exc:
|
||||
raise http.CloudError(f"could not enumerate Git source files: {exc}") from exc
|
||||
assert process.stdout is not None
|
||||
buffer = b""
|
||||
count = 0
|
||||
try:
|
||||
while chunk := process.stdout.read(64 * 1024):
|
||||
buffer += chunk
|
||||
records = buffer.split(b"\0")
|
||||
buffer = records.pop()
|
||||
for raw in records:
|
||||
relative = _git_relative_path(raw, relative_source)
|
||||
if relative is None:
|
||||
continue
|
||||
count += 1
|
||||
_check_candidate_limit(count)
|
||||
yield relative
|
||||
if buffer:
|
||||
raise http.CloudError("Git returned a malformed source file manifest.")
|
||||
if process.wait() != 0:
|
||||
raise http.CloudError("Git could not enumerate the source directory.")
|
||||
finally:
|
||||
process.stdout.close()
|
||||
if process.poll() is None:
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=1)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
process.wait()
|
||||
|
||||
|
||||
def _git_relative_path(raw: bytes, relative_source: Path) -> Path | None:
|
||||
repo_relative = Path(os.fsdecode(raw))
|
||||
try:
|
||||
relative = repo_relative.relative_to(relative_source)
|
||||
except ValueError:
|
||||
return None
|
||||
if relative.is_absolute() or ".." in relative.parts:
|
||||
raise http.CloudError("Git returned an unsafe source path.")
|
||||
return relative
|
||||
|
||||
|
||||
def _walk_candidate_paths(
|
||||
source: Path,
|
||||
*,
|
||||
include_hidden: bool,
|
||||
patterns: list[str],
|
||||
excluded: Counter[str],
|
||||
) -> Iterator[Path]:
|
||||
"""Walk top-down so excluded dependency, VCS, and hidden trees are never traversed."""
|
||||
count = 0
|
||||
stack: list[tuple[Path, _ScandirIterator]] = []
|
||||
try:
|
||||
stack.append((source, os.scandir(source)))
|
||||
while stack:
|
||||
root_path, entries = stack[-1]
|
||||
try:
|
||||
entry = next(entries)
|
||||
except StopIteration:
|
||||
entries.close()
|
||||
stack.pop()
|
||||
continue
|
||||
count += 1
|
||||
_check_candidate_limit(count)
|
||||
path = root_path / entry.name
|
||||
relative = path.relative_to(source)
|
||||
try:
|
||||
is_directory = entry.is_dir(follow_symlinks=False)
|
||||
is_symlink = entry.is_symlink()
|
||||
except OSError:
|
||||
excluded["unreadable"] += 1
|
||||
continue
|
||||
if is_directory:
|
||||
reason = _pruned_directory_reason(
|
||||
relative,
|
||||
include_hidden=include_hidden,
|
||||
patterns=patterns,
|
||||
)
|
||||
if reason:
|
||||
excluded[reason] += 1
|
||||
continue
|
||||
try:
|
||||
stack.append((path, os.scandir(path)))
|
||||
except OSError:
|
||||
excluded["unreadable"] += 1
|
||||
continue
|
||||
if is_symlink:
|
||||
excluded["symlink_or_non_file"] += 1
|
||||
continue
|
||||
yield relative
|
||||
except OSError as exc:
|
||||
raise http.CloudError(f"could not enumerate source directory {source}: {exc}") from exc
|
||||
finally:
|
||||
for _, entries in stack:
|
||||
entries.close()
|
||||
|
||||
|
||||
def _pruned_directory_reason(
|
||||
relative: Path,
|
||||
*,
|
||||
include_hidden: bool,
|
||||
patterns: list[str],
|
||||
) -> str | None:
|
||||
lower_parts = tuple(part.lower() for part in relative.parts)
|
||||
if any(part == ".git" for part in lower_parts):
|
||||
return "git_metadata"
|
||||
if any(part in _ALWAYS_EXCLUDED_DIRS for part in lower_parts):
|
||||
return "dependency_or_build_output"
|
||||
if not include_hidden and any(part.startswith(".") for part in relative.parts):
|
||||
return "hidden"
|
||||
if any(_matches_user_pattern(relative, pattern) for pattern in patterns):
|
||||
return "user_pattern"
|
||||
return None
|
||||
|
||||
|
||||
def _check_candidate_limit(count: int) -> None:
|
||||
if count > MAX_CANDIDATE_PATHS:
|
||||
raise http.CloudError(
|
||||
f"source enumeration exceeded {MAX_CANDIDATE_PATHS:,} paths before filtering; "
|
||||
"narrow --source or add directory exclusions."
|
||||
)
|
||||
|
||||
|
||||
def _git_root(source: Path) -> Path | None:
|
||||
git = shutil.which("git")
|
||||
if git is None:
|
||||
return None
|
||||
result = subprocess.run( # noqa: S603 # nosec B603
|
||||
[git, "-C", str(source), "rev-parse", "--show-toplevel"],
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
return None
|
||||
try:
|
||||
return Path(result.stdout.strip()).resolve()
|
||||
except OSError:
|
||||
return None
|
||||
|
||||
|
||||
def _exclusion_reason( # noqa: PLR0911
|
||||
relative: Path,
|
||||
*,
|
||||
include_hidden: bool,
|
||||
include_sensitive: bool,
|
||||
include_archives: bool,
|
||||
patterns: list[str],
|
||||
) -> str | None:
|
||||
parts = relative.parts
|
||||
lower_parts = tuple(part.lower() for part in parts)
|
||||
if any(part == ".git" for part in lower_parts):
|
||||
return "git_metadata"
|
||||
if any(part in _ALWAYS_EXCLUDED_DIRS for part in lower_parts[:-1]):
|
||||
return "dependency_or_build_output"
|
||||
if not include_hidden and any(part.startswith(".") for part in parts):
|
||||
return "hidden"
|
||||
if any(_matches_user_pattern(relative, pattern) for pattern in patterns):
|
||||
return "user_pattern"
|
||||
name = relative.name.lower()
|
||||
if not include_sensitive and (
|
||||
name in _SENSITIVE_NAMES
|
||||
or any(fnmatch.fnmatch(name, pattern) for pattern in _SENSITIVE_PATTERNS)
|
||||
or any(
|
||||
lower_parts[-len(suffix) :] == suffix
|
||||
for suffix in _SENSITIVE_PATH_SUFFIXES
|
||||
if len(lower_parts) >= len(suffix)
|
||||
)
|
||||
):
|
||||
return "sensitive_filename"
|
||||
if not include_archives and name.endswith(_ARCHIVE_SUFFIXES):
|
||||
return "nested_archive"
|
||||
return None
|
||||
|
||||
|
||||
def _matches_user_pattern(relative: Path, pattern: str) -> bool:
|
||||
"""Match exclude globs, including intuitive trailing-slash directory rules."""
|
||||
relative_posix = relative.as_posix()
|
||||
posix = PurePosixPath(relative_posix)
|
||||
if pattern.endswith("/"):
|
||||
directory_pattern = pattern.rstrip("/")
|
||||
if not directory_pattern:
|
||||
return False
|
||||
return (
|
||||
posix.match(directory_pattern)
|
||||
or fnmatch.fnmatch(relative_posix, directory_pattern)
|
||||
or any(
|
||||
PurePosixPath(parent.as_posix()).match(directory_pattern)
|
||||
or fnmatch.fnmatch(parent.as_posix(), directory_pattern)
|
||||
for parent in posix.parents
|
||||
if parent != PurePosixPath(".")
|
||||
)
|
||||
)
|
||||
return posix.match(pattern) or fnmatch.fnmatch(relative_posix, pattern)
|
||||
|
||||
|
||||
def _write_archive(destination: Path, files: tuple[SelectedFile, ...]) -> None:
|
||||
with zipfile.ZipFile(
|
||||
destination, "w", compression=zipfile.ZIP_DEFLATED, compresslevel=6
|
||||
) as archive:
|
||||
for item in files:
|
||||
flags = os.O_RDONLY | getattr(os, "O_NOFOLLOW", 0)
|
||||
try:
|
||||
descriptor = os.open(item.path, flags)
|
||||
except OSError as exc:
|
||||
raise http.CloudError(f"could not safely read {item.archive_name}: {exc}") from exc
|
||||
with os.fdopen(descriptor, "rb") as source_file:
|
||||
current = os.fstat(source_file.fileno())
|
||||
if (
|
||||
not stat.S_ISREG(current.st_mode)
|
||||
or current.st_size != item.size
|
||||
or current.st_dev != item.device
|
||||
or current.st_ino != item.inode
|
||||
or current.st_mtime_ns != item.mtime_ns
|
||||
or current.st_ctime_ns != item.ctime_ns
|
||||
):
|
||||
raise http.CloudError(
|
||||
f"{item.archive_name} changed while the source archive was being built; "
|
||||
"retry."
|
||||
)
|
||||
info = zipfile.ZipInfo(item.archive_name)
|
||||
info.compress_type = zipfile.ZIP_DEFLATED
|
||||
info.external_attr = 0o100644 << 16
|
||||
with archive.open(info, "w", force_zip64=True) as target:
|
||||
remaining = item.size
|
||||
while remaining:
|
||||
chunk = source_file.read(min(1024 * 1024, remaining))
|
||||
if not chunk:
|
||||
raise http.CloudError(
|
||||
f"{item.archive_name} changed while the source archive was being "
|
||||
"built; retry."
|
||||
)
|
||||
target.write(chunk)
|
||||
remaining -= len(chunk)
|
||||
final = os.fstat(source_file.fileno())
|
||||
if (
|
||||
source_file.read(1)
|
||||
or not stat.S_ISREG(final.st_mode)
|
||||
or final.st_size != item.size
|
||||
or final.st_dev != item.device
|
||||
or final.st_ino != item.inode
|
||||
or final.st_mtime_ns != item.mtime_ns
|
||||
or final.st_ctime_ns != item.ctime_ns
|
||||
):
|
||||
raise http.CloudError(
|
||||
f"{item.archive_name} changed while the source archive was being "
|
||||
"built; retry."
|
||||
)
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as stream:
|
||||
for chunk in iter(lambda: stream.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _has_archive_magic(path: Path) -> bool:
|
||||
"""Recognize common archive containers even when their suffix is disguised."""
|
||||
flags = os.O_RDONLY | getattr(os, "O_NOFOLLOW", 0)
|
||||
try:
|
||||
descriptor = os.open(path, flags)
|
||||
with os.fdopen(descriptor, "rb") as stream:
|
||||
header = stream.read(512)
|
||||
except OSError:
|
||||
return False
|
||||
return header.startswith(_ARCHIVE_MAGIC_PREFIXES) or header[257:262] == b"ustar"
|
||||
|
||||
|
||||
def _load_ignore_patterns(source: Path) -> list[str]:
|
||||
path = source / ".strixignore"
|
||||
raw_text = _read_ignore_file(path)
|
||||
if raw_text is None:
|
||||
return []
|
||||
if len(raw_text) > MAX_IGNORE_BYTES:
|
||||
raise http.CloudError(f"{path} is larger than the {MAX_IGNORE_BYTES:,}-byte limit.")
|
||||
try:
|
||||
lines = raw_text.decode("utf-8").splitlines()
|
||||
except UnicodeDecodeError as exc:
|
||||
raise http.CloudError(f"{path} must be UTF-8 text.") from exc
|
||||
patterns: list[str] = []
|
||||
for line_number, raw in enumerate(lines, start=1):
|
||||
value = raw.strip()
|
||||
if not value or value.startswith("#"):
|
||||
continue
|
||||
if value.startswith("!"):
|
||||
raise http.CloudError(
|
||||
f"{path}:{line_number}: negated patterns are not supported; use exclude-only globs."
|
||||
)
|
||||
patterns.append(value)
|
||||
if len(patterns) > MAX_IGNORE_PATTERNS:
|
||||
raise http.CloudError(
|
||||
f"{path} contains more than {MAX_IGNORE_PATTERNS:,} exclusion patterns."
|
||||
)
|
||||
return patterns
|
||||
|
||||
|
||||
def _read_ignore_file(path: Path) -> bytes | None:
|
||||
"""Read a bounded regular ignore file without blocking on a FIFO or device."""
|
||||
try:
|
||||
descriptor = os.open(
|
||||
path,
|
||||
os.O_RDONLY | getattr(os, "O_NOFOLLOW", 0) | getattr(os, "O_NONBLOCK", 0),
|
||||
)
|
||||
except FileNotFoundError:
|
||||
return None
|
||||
except OSError as exc:
|
||||
raise http.CloudError(f"could not read {path}: {exc}") from exc
|
||||
try:
|
||||
info = os.fstat(descriptor)
|
||||
except OSError as exc:
|
||||
os.close(descriptor)
|
||||
raise http.CloudError(f"could not inspect {path}: {exc}") from exc
|
||||
if not stat.S_ISREG(info.st_mode):
|
||||
os.close(descriptor)
|
||||
raise http.CloudError(f"{path} must be a regular file.")
|
||||
try:
|
||||
stream = os.fdopen(descriptor, "rb")
|
||||
except OSError as exc:
|
||||
os.close(descriptor)
|
||||
raise http.CloudError(f"could not read {path}: {exc}") from exc
|
||||
try:
|
||||
return stream.read(MAX_IGNORE_BYTES + 1)
|
||||
except OSError as exc:
|
||||
raise http.CloudError(f"could not read {path}: {exc}") from exc
|
||||
finally:
|
||||
stream.close()
|
||||
|
||||
|
||||
def _validate_patterns(patterns: list[str]) -> None:
|
||||
if len(patterns) > MAX_IGNORE_PATTERNS:
|
||||
raise http.CloudError(
|
||||
f"source upload accepts at most {MAX_IGNORE_PATTERNS:,} exclusion patterns."
|
||||
)
|
||||
for pattern in patterns:
|
||||
if len(pattern) > MAX_IGNORE_PATTERN_CHARS:
|
||||
raise http.CloudError(
|
||||
"source exclusion patterns must be at most "
|
||||
f"{MAX_IGNORE_PATTERN_CHARS:,} characters each."
|
||||
)
|
||||
if "\x00" in pattern:
|
||||
raise http.CloudError("source exclusion patterns cannot contain NUL bytes.")
|
||||
1147
strix/interface/cloud/spec.py
Normal file
1147
strix/interface/cloud/spec.py
Normal file
File diff suppressed because it is too large
Load Diff
291
strix/interface/cloud/workspaces.py
Normal file
291
strix/interface/cloud/workspaces.py
Normal file
@@ -0,0 +1,291 @@
|
||||
"""`strix cloud workspaces use` — switch the stored token to another workspace.
|
||||
|
||||
The command lists the workspaces of the account, finds the requested one by
|
||||
ID or by exact name, asks the platform to rotate that token in place, and
|
||||
stores the returned workspace metadata. The bearer secret and expiry stay the
|
||||
same; the account's role in the target workspace limits the granted scopes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
|
||||
import strix.interface.cloud.http as http # noqa: PLR0402
|
||||
from strix.interface.cloud.arguments import CloudArgumentParser
|
||||
from strix.interface.cloud.render import emit, json_mode
|
||||
from strix.interface.platform_cli import AUTH_PATH, read_record, save_record
|
||||
from strix.interface.platform_identity import read_or_create_identity
|
||||
from strix.interface.terminal_text import sanitize_terminal_text
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import argparse
|
||||
|
||||
|
||||
def run_workspace_use(argv: list[str]) -> int:
|
||||
"""Entry point for ``strix cloud workspaces use``. Returns an exit code."""
|
||||
console = Console()
|
||||
parser = CloudArgumentParser(
|
||||
prog="strix cloud workspaces use",
|
||||
description="Switch the stored API token to another workspace.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"workspace",
|
||||
metavar="WORKSPACE",
|
||||
help="Workspace number from `workspaces list`, ID, or exact name.",
|
||||
)
|
||||
scope_mode = parser.add_mutually_exclusive_group()
|
||||
scope_mode.add_argument(
|
||||
"--scopes",
|
||||
nargs="+",
|
||||
metavar="SCOPE",
|
||||
default=None,
|
||||
help=(
|
||||
"Use a custom scope set within the login-approved ceiling. "
|
||||
"Without this option, preserve the server-side scope preference."
|
||||
),
|
||||
)
|
||||
scope_mode.add_argument(
|
||||
"--scope-profile",
|
||||
choices=("minimal", "recommended", "full"),
|
||||
default=None,
|
||||
help="Change to a profile within the authority approved at login.",
|
||||
)
|
||||
parser.add_argument("--show-scopes", action="store_true", help="Print every granted scope.")
|
||||
parser.add_argument("--json", action="store_true", help="Print the raw JSON response.")
|
||||
parser.add_argument("--token", default=None, help="API token override.")
|
||||
parser.add_argument(
|
||||
"--workspace-id",
|
||||
default=None,
|
||||
metavar="ORG_ID",
|
||||
help="Expected workspace for an override CLI token.",
|
||||
)
|
||||
parser.add_argument("--app-url", default=None, metavar="URL", help="Platform URL override.")
|
||||
parser.add_argument(
|
||||
"--timeout", default=None, type=float, metavar="SECONDS", help="Request timeout in seconds."
|
||||
)
|
||||
as_json = json_mode(flag="--json" in argv)
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
except SystemExit as exc:
|
||||
return exc.code if isinstance(exc.code, int) else 2
|
||||
except http.CloudError as exc:
|
||||
_emit_cloud_error(console, exc, as_json=as_json)
|
||||
return exc.exit_code
|
||||
|
||||
as_json = json_mode(flag=bool(args.json))
|
||||
try:
|
||||
http.configure(
|
||||
base_url=args.app_url,
|
||||
timeout=args.timeout,
|
||||
token_override=bool(args.token),
|
||||
workspace_id=args.workspace_id,
|
||||
)
|
||||
return _use(console, args, as_json=as_json)
|
||||
except http.CloudError as exc:
|
||||
_emit_cloud_error(console, exc, as_json=as_json)
|
||||
return exc.exit_code
|
||||
|
||||
|
||||
def _use( # noqa: PLR0912, PLR0915
|
||||
console: Console, args: argparse.Namespace, *, as_json: bool
|
||||
) -> int:
|
||||
workspace = _find_workspace(args.workspace, token=args.token)
|
||||
stored_record: dict[str, Any] = read_record() or {}
|
||||
# An override token may belong to a different account. Never mix its new
|
||||
# workspace state with identity or scope preferences from the stored sign-in.
|
||||
external_token = args.token is not None or bool(os.environ.get("STRIX_API_TOKEN", "").strip())
|
||||
record: dict[str, Any] = {} if external_token else dict(stored_record)
|
||||
body: dict[str, Any] = {}
|
||||
if args.scopes:
|
||||
body["scopes"] = args.scopes
|
||||
body["scope_profile"] = "custom"
|
||||
elif args.scope_profile:
|
||||
body["scope_profile"] = args.scope_profile
|
||||
if not external_token:
|
||||
try:
|
||||
body.update(read_or_create_identity())
|
||||
except (OSError, ValueError) as exc:
|
||||
raise http.CloudError(f"could not load the CLI device identity: {exc}") from exc
|
||||
switched = _switch_workspace_token(
|
||||
str(workspace["id"]),
|
||||
token=args.token,
|
||||
body=body or None,
|
||||
)
|
||||
if not isinstance(switched, dict):
|
||||
raise _workspace_switch_unknown("the platform returned an invalid response")
|
||||
switched_record = cast("dict[str, Any]", switched)
|
||||
switched_token = switched_record.get("api_token")
|
||||
if not isinstance(switched_token, str) or not switched_token.strip():
|
||||
raise _workspace_switch_unknown("the platform response omitted the token")
|
||||
switched_scopes = switched_record.get("scopes")
|
||||
switched_scope_items = cast("list[Any]", cast("Any", switched_scopes))
|
||||
if not isinstance(switched_scopes, list) or not all(
|
||||
isinstance(scope, str) for scope in switched_scope_items
|
||||
):
|
||||
raise _workspace_switch_unknown("the platform response contained invalid scopes")
|
||||
validated_scopes = cast("list[str]", switched_scope_items)
|
||||
|
||||
record.update(
|
||||
{
|
||||
"api_token": switched_token,
|
||||
"organization_id": switched_record.get("organization_id", workspace["id"]),
|
||||
"organization_name": switched_record.get(
|
||||
"organization_name", workspace.get("name", "")
|
||||
),
|
||||
"expires_at": switched_record.get("expires_at") or stored_record.get("expires_at"),
|
||||
"scopes": validated_scopes,
|
||||
"requested_scopes": switched_record.get("requested_scopes", validated_scopes),
|
||||
"scope_ceiling": switched_record.get("scope_ceiling", []),
|
||||
"scope_profile": switched_record.get("scope_profile", "custom"),
|
||||
"token_id": switched_record.get("token_id"),
|
||||
"credential_source": switched_record.get("credential_source", "api"),
|
||||
"device_name": switched_record.get("device_name"),
|
||||
"app_url": http.app_url(),
|
||||
}
|
||||
)
|
||||
if switched_record.get("email"):
|
||||
record["email"] = switched_record["email"]
|
||||
if not external_token:
|
||||
try:
|
||||
save_record(record)
|
||||
except OSError as exc:
|
||||
raise http.CloudError(
|
||||
"the platform switched the token, but the local workspace metadata could not be "
|
||||
f"stored in {AUTH_PATH}: {exc}. The bearer is still valid; fix the file and safely "
|
||||
"rerun the same workspace use command.",
|
||||
payload={
|
||||
"workspace_switched": True,
|
||||
"local_record_updated": False,
|
||||
"retry_safe": True,
|
||||
},
|
||||
) from exc
|
||||
|
||||
result = {
|
||||
"workspace_id": record["organization_id"],
|
||||
"workspace_name": record["organization_name"],
|
||||
"scopes": record["scopes"],
|
||||
"requested_scopes": record.get("requested_scopes", record["scopes"]),
|
||||
"scope_ceiling": record.get("scope_ceiling", []),
|
||||
"scope_profile": record.get("scope_profile", "custom"),
|
||||
"expires_at": record.get("expires_at"),
|
||||
"token_id": record.get("token_id"),
|
||||
"credential_source": record.get("credential_source", "api"),
|
||||
"device_name": record.get("device_name"),
|
||||
"stored": not external_token,
|
||||
}
|
||||
if as_json:
|
||||
emit(console, result, as_json=True)
|
||||
return http.EXIT_OK
|
||||
workspace_name = escape(sanitize_terminal_text(record["organization_name"]))
|
||||
console.print(f"[green]✓ Switched to workspace [bold]{workspace_name}[/].[/]")
|
||||
scopes = record.get("scopes")
|
||||
if isinstance(scopes, list) and scopes:
|
||||
scope_items = cast("list[Any]", cast("Any", scopes))
|
||||
scope_names = [scope for scope in scope_items if isinstance(scope, str)]
|
||||
if scope_names and args.show_scopes:
|
||||
rendered_scopes = escape(sanitize_terminal_text(" ".join(scope_names)))
|
||||
console.print(f" Scopes: [dim]{rendered_scopes}[/]")
|
||||
elif scope_names:
|
||||
profile = str(record.get("scope_profile") or "custom").title()
|
||||
console.print(f" Access: [dim]{profile} · {len(scope_names)} scopes granted[/]")
|
||||
if external_token:
|
||||
console.print(" Token: [dim]override used for this command only; not stored[/]")
|
||||
else:
|
||||
console.print(f" Token: stored in [dim]{escape(sanitize_terminal_text(AUTH_PATH))}[/]")
|
||||
return http.EXIT_OK
|
||||
|
||||
|
||||
def _switch_workspace_token(
|
||||
workspace_id: str,
|
||||
*,
|
||||
token: str | None,
|
||||
body: dict[str, Any] | None,
|
||||
) -> Any:
|
||||
"""Switch in place, distinguishing definitive rejections from lost outcomes."""
|
||||
try:
|
||||
response = http.request(
|
||||
"POST",
|
||||
f"/workspaces/{workspace_id}/token",
|
||||
token=token,
|
||||
body=body,
|
||||
)
|
||||
except http.CloudError as exc:
|
||||
raise _workspace_switch_unknown(str(exc)) from exc
|
||||
|
||||
# Client/auth/conflict responses prove the rotation did not return success.
|
||||
# A 5xx or malformed success may arrive after the database commit, but the
|
||||
# server preserves the bearer so replaying this exact command is safe.
|
||||
if response.status_code in {400, 401, 403, 404, 409, 422}:
|
||||
return http.check(response)
|
||||
try:
|
||||
return http.check(response)
|
||||
except http.CloudError as exc:
|
||||
raise _workspace_switch_unknown(str(exc)) from exc
|
||||
|
||||
|
||||
def _workspace_switch_unknown(detail: str) -> http.CloudError:
|
||||
return http.CloudError(
|
||||
"workspace switch outcome is unknown: "
|
||||
f"{sanitize_terminal_text(detail)}. The bearer secret is unchanged; safely rerun the "
|
||||
"same workspace use command, or list workspaces to check the current one.",
|
||||
payload={
|
||||
"switch_outcome_unknown": True,
|
||||
"retry_safe": True,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _emit_cloud_error(console: Console, error: http.CloudError, *, as_json: bool) -> None:
|
||||
if as_json:
|
||||
raw_payload: Any = error.payload
|
||||
error_payload = cast("dict[str, Any]", raw_payload)
|
||||
payload = dict(error_payload) if isinstance(raw_payload, dict) else {}
|
||||
payload["error"] = str(error)
|
||||
emit(console, payload, as_json=True)
|
||||
return
|
||||
console.print(f"[red]Error:[/] {escape(sanitize_terminal_text(error))}")
|
||||
|
||||
|
||||
def _find_workspace(selector: str, *, token: str | None) -> dict[str, Any]:
|
||||
listed = http.check(http.request("GET", "/workspaces", token=token))
|
||||
listed_record = cast("dict[str, Any]", listed) if isinstance(listed, dict) else {}
|
||||
items = listed_record.get("workspaces")
|
||||
item_values = cast("list[Any]", cast("Any", items)) if isinstance(items, list) else []
|
||||
workspaces = [
|
||||
cast("dict[str, Any]", cast("Any", item)) for item in item_values if isinstance(item, dict)
|
||||
]
|
||||
if not workspaces:
|
||||
raise http.CloudError("no workspaces found for this account.")
|
||||
wanted = selector.strip()
|
||||
if wanted.isdigit():
|
||||
index = int(wanted)
|
||||
if 1 <= index <= len(workspaces):
|
||||
return workspaces[index - 1]
|
||||
raise http.CloudError(
|
||||
f"workspace number must be between 1 and {len(workspaces)}. "
|
||||
"Run `strix cloud workspaces` to see the numbered list."
|
||||
)
|
||||
by_id = [w for w in workspaces if w.get("id") == wanted]
|
||||
if by_id:
|
||||
return by_id[0]
|
||||
by_name = [w for w in workspaces if str(w.get("name", "")).casefold() == wanted.casefold()]
|
||||
if len(by_name) == 1:
|
||||
return by_name[0]
|
||||
if len(by_name) > 1:
|
||||
numbers = ", ".join(
|
||||
str(index)
|
||||
for index, workspace in enumerate(workspaces, start=1)
|
||||
if workspace in by_name
|
||||
)
|
||||
raise http.CloudError(
|
||||
f"multiple workspaces are named {wanted!r}. Use its list number: {numbers}"
|
||||
)
|
||||
names = ", ".join(
|
||||
f"{index}: {workspace.get('name')}" for index, workspace in enumerate(workspaces, start=1)
|
||||
)
|
||||
raise http.CloudError(f"no workspace matches {wanted!r}. Your workspaces: {names}")
|
||||
373
strix/interface/completions.py
Normal file
373
strix/interface/completions.py
Normal file
@@ -0,0 +1,373 @@
|
||||
"""Shell completion scripts and candidates for the Strix CLI."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from strix.interface.cloud.spec import DEFAULT_VERBS, SPEC, Cmd
|
||||
from strix.interface.terminal_text import has_terminal_control, sanitize_terminal_text
|
||||
|
||||
|
||||
_ROOT_COMMANDS = ("cloud", "auth", "view", "completions", "completion")
|
||||
_SESSION_COMMANDS = ("login", "logout", "whoami", "session", "credits")
|
||||
_COMMON_FLAGS = (
|
||||
"--json",
|
||||
"--token",
|
||||
"--workspace-id",
|
||||
"--app-url",
|
||||
"--timeout",
|
||||
"-h",
|
||||
"--help",
|
||||
)
|
||||
_COMMON_VALUE_FLAGS = frozenset({"--token", "--workspace-id", "--app-url", "--timeout"})
|
||||
_WORKSPACE_USE_FLAGS = (*_COMMON_FLAGS, "--scopes", "--scope-profile", "--show-scopes")
|
||||
|
||||
|
||||
def run_completions(argv: list[str]) -> int:
|
||||
"""Print a shell integration script or hidden completion candidates."""
|
||||
if argv and argv[0] == "--candidates":
|
||||
for candidate in completion_candidates(argv[1:]):
|
||||
sys.stdout.write(candidate + "\n")
|
||||
return 0
|
||||
if not argv or argv[0] in ("-h", "--help", "help"):
|
||||
sys.stdout.write(
|
||||
"Usage: strix completions <zsh|bash|fish>\n\n"
|
||||
"Enable tab completion for the current shell:\n"
|
||||
" zsh: source <(strix completions zsh)\n"
|
||||
" bash: source <(strix completions bash)\n"
|
||||
" fish: strix completions fish | source\n"
|
||||
)
|
||||
return 0
|
||||
shell = argv[0].lower()
|
||||
scripts = {"zsh": _zsh_script, "bash": _bash_script, "fish": _fish_script}
|
||||
generator = scripts.get(shell)
|
||||
if generator is None:
|
||||
sys.stderr.write(
|
||||
f"Unknown shell: {sanitize_terminal_text(shell)}. Choose zsh, bash, or fish.\n"
|
||||
)
|
||||
return 2
|
||||
sys.stdout.write(generator())
|
||||
return 0
|
||||
|
||||
|
||||
def completion_candidates(words: list[str]) -> list[str]:
|
||||
"""Return candidates for words after the ``strix`` executable."""
|
||||
prior, current = _split_cursor(words)
|
||||
if not prior:
|
||||
candidates = _matching(_ROOT_COMMANDS, current)
|
||||
elif prior[0] != "cloud":
|
||||
candidates = []
|
||||
else:
|
||||
candidates = _cloud_candidates(prior[1:], current)
|
||||
# The line-oriented shell protocol cannot represent these names safely.
|
||||
# Omitting them is preferable to returning a sanitized path that does not exist.
|
||||
return [candidate for candidate in candidates if not has_terminal_control(candidate)]
|
||||
|
||||
|
||||
def _split_cursor(words: list[str]) -> tuple[list[str], str]:
|
||||
if not words:
|
||||
return [], ""
|
||||
return words[:-1], words[-1]
|
||||
|
||||
|
||||
def _cloud_candidates(prior: list[str], current: str) -> list[str]: # noqa: PLR0911
|
||||
groups = (*_SESSION_COMMANDS, *SPEC, "workspace")
|
||||
if not prior:
|
||||
return _matching(groups, current)
|
||||
group = "workspaces" if prior[0] == "workspace" else prior[0]
|
||||
rest = prior[1:]
|
||||
if group in _SESSION_COMMANDS:
|
||||
return _session_candidates(group, rest, current)
|
||||
commands = SPEC.get(group)
|
||||
if commands is None:
|
||||
return _matching(groups, current)
|
||||
default_verb = DEFAULT_VERBS.get(group)
|
||||
default_is_active = (rest and rest[0].startswith("-")) or (not rest and current.startswith("-"))
|
||||
if default_verb is not None and default_is_active:
|
||||
return _command_candidates(commands[default_verb], rest, current)
|
||||
|
||||
command_paths = sorted(
|
||||
((verb.split(), cmd) for verb, cmd in commands.items()),
|
||||
key=lambda item: len(item[0]),
|
||||
reverse=True,
|
||||
)
|
||||
for path, cmd in command_paths:
|
||||
if rest[: len(path)] == path:
|
||||
command_candidates = _command_candidates(cmd, rest[len(path) :], current)
|
||||
if rest == path:
|
||||
nested_words = {
|
||||
candidate_path[len(path)]
|
||||
for candidate_path, _candidate_cmd in command_paths
|
||||
if len(candidate_path) > len(path) and candidate_path[: len(path)] == path
|
||||
}
|
||||
return sorted({*command_candidates, *_matching(nested_words, current)})
|
||||
return command_candidates
|
||||
if group == "workspaces" and rest[:1] == ["use"]:
|
||||
return _flag_candidates(
|
||||
_WORKSPACE_USE_FLAGS,
|
||||
rest[1:],
|
||||
current,
|
||||
value_flags=_COMMON_VALUE_FLAGS | {"--scopes"},
|
||||
)
|
||||
|
||||
verb_paths = [path for path, _cmd in command_paths]
|
||||
if group == "workspaces":
|
||||
verb_paths.append(["use"])
|
||||
matching_paths = [path for path in verb_paths if path[: len(rest)] == rest]
|
||||
if not matching_paths:
|
||||
return []
|
||||
next_words = sorted({path[len(rest)] for path in matching_paths if len(path) > len(rest)})
|
||||
return _matching(next_words, current)
|
||||
|
||||
|
||||
def _session_candidates(group: str, prior: list[str], current: str) -> list[str]:
|
||||
if group == "session":
|
||||
if not prior:
|
||||
return _matching(("show", "scopes", "help", *_COMMON_FLAGS, "--show-scopes"), current)
|
||||
if prior[:1] == ["scopes"] and len(prior) == 1:
|
||||
return _matching(("set", *_COMMON_FLAGS, "--show-scopes"), current)
|
||||
if prior[:2] == ["scopes", "set"]:
|
||||
return _matching(
|
||||
("minimal", "recommended", "full", "--scopes", *_COMMON_FLAGS, "--show-scopes"),
|
||||
current,
|
||||
)
|
||||
return _flag_candidates(
|
||||
(*_COMMON_FLAGS, "--show-scopes"),
|
||||
prior,
|
||||
current,
|
||||
value_flags=_COMMON_VALUE_FLAGS | {"--scopes"},
|
||||
)
|
||||
flags = _session_flags(group)
|
||||
value_flags: frozenset[str] = frozenset()
|
||||
if group == "login":
|
||||
value_flags = frozenset({"--scopes", "--scope-profile", "--workspace", "--device-name"})
|
||||
elif group == "credits":
|
||||
value_flags = _COMMON_VALUE_FLAGS
|
||||
return _flag_candidates(flags, prior, current, value_flags=value_flags)
|
||||
|
||||
|
||||
def _session_flags(group: str) -> tuple[str, ...]:
|
||||
if group == "login":
|
||||
return (
|
||||
"--no-browser",
|
||||
"--scopes",
|
||||
"--scope-profile",
|
||||
"--workspace",
|
||||
"--device-name",
|
||||
"-h",
|
||||
"--help",
|
||||
)
|
||||
if group == "whoami":
|
||||
return ("--json", "--show-scopes", "-h", "--help")
|
||||
if group == "logout":
|
||||
return ("--json", "--local-only", "-h", "--help")
|
||||
if group == "credits":
|
||||
return _COMMON_FLAGS
|
||||
return ("-h", "--help")
|
||||
|
||||
|
||||
def _command_candidates(cmd: Cmd, prior: list[str], current: str) -> list[str]:
|
||||
filesystem = _filesystem_candidates(cmd, prior, current)
|
||||
if filesystem is not None:
|
||||
return filesystem
|
||||
return _flag_candidates(
|
||||
_command_flags(cmd),
|
||||
prior,
|
||||
current,
|
||||
value_flags=_command_value_flags(cmd),
|
||||
)
|
||||
|
||||
|
||||
def _flag_candidates(
|
||||
flags: tuple[str, ...],
|
||||
prior: list[str],
|
||||
current: str,
|
||||
*,
|
||||
value_flags: frozenset[str],
|
||||
) -> list[str]:
|
||||
if prior and prior[-1] in value_flags and not current.startswith("-"):
|
||||
return []
|
||||
return _matching(flags, current)
|
||||
|
||||
|
||||
def _command_flags(cmd: Cmd) -> tuple[str, ...]:
|
||||
flags: list[str] = list(_COMMON_FLAGS)
|
||||
for param in cmd.query + cmd.body:
|
||||
flag = "--" + (param.flag or _kebab(param.name))
|
||||
flags.append(flag)
|
||||
if param.kind == "bool":
|
||||
flags.append("--no-" + flag.removeprefix("--"))
|
||||
if cmd.method in ("POST", "PUT", "PATCH"):
|
||||
flags.append("--data")
|
||||
if cmd.idempotent:
|
||||
flags.append("--idempotency-key")
|
||||
if cmd.binary or cmd.path == "/audit":
|
||||
flags.extend(("--output", "--force"))
|
||||
if cmd.link:
|
||||
flags.append("--no-browser")
|
||||
if cmd.wait_path or cmd.wait_self:
|
||||
flags.extend(("--wait", "--wait-timeout"))
|
||||
if cmd.path == "/billing/topup":
|
||||
flags.extend(("--yes", "--no-pay", "--payment-method"))
|
||||
if cmd.path == "/scans" and cmd.method == "POST":
|
||||
flags.extend(
|
||||
(
|
||||
"--source",
|
||||
"--approve-sha256",
|
||||
"--dry-run",
|
||||
"--yes",
|
||||
"--show-files",
|
||||
"--exclude",
|
||||
"--include-hidden",
|
||||
"--include-sensitive",
|
||||
"--include-archives",
|
||||
)
|
||||
)
|
||||
if cmd.path == "/billing/auto-topup" and cmd.method == "PUT":
|
||||
flags.append("--no-monthly-cap")
|
||||
return tuple(dict.fromkeys(flags))
|
||||
|
||||
|
||||
def _command_value_flags(cmd: Cmd) -> frozenset[str]:
|
||||
flags = set(_COMMON_VALUE_FLAGS)
|
||||
for param in cmd.query + cmd.body:
|
||||
if param.kind != "bool":
|
||||
flags.add("--" + (param.flag or _kebab(param.name)))
|
||||
if cmd.method in ("POST", "PUT", "PATCH"):
|
||||
flags.add("--data")
|
||||
if cmd.idempotent:
|
||||
flags.add("--idempotency-key")
|
||||
if cmd.binary or cmd.path == "/audit":
|
||||
flags.add("--output")
|
||||
if cmd.wait_path or cmd.wait_self:
|
||||
flags.add("--wait-timeout")
|
||||
if cmd.path == "/billing/topup":
|
||||
flags.add("--payment-method")
|
||||
if cmd.path == "/scans" and cmd.method == "POST":
|
||||
flags.update(("--source", "--approve-sha256", "--exclude"))
|
||||
return frozenset(flags)
|
||||
|
||||
|
||||
def _filesystem_candidates( # noqa: PLR0911
|
||||
cmd: Cmd, prior: list[str], current: str
|
||||
) -> list[str] | None:
|
||||
inline = (
|
||||
("--source=", True, ""),
|
||||
("--output=", False, ""),
|
||||
("--data=@", False, "@"),
|
||||
)
|
||||
for option, directories_only, marker in inline:
|
||||
if current.startswith(option):
|
||||
value = current.removeprefix(option)
|
||||
return [
|
||||
option + candidate.removeprefix(marker)
|
||||
for candidate in _path_candidates(
|
||||
marker + value,
|
||||
directories_only=directories_only,
|
||||
marker=marker,
|
||||
)
|
||||
]
|
||||
|
||||
if not prior or current.startswith("-"):
|
||||
return None
|
||||
option = prior[-1]
|
||||
if option == "--source" and cmd.path == "/scans" and cmd.method == "POST":
|
||||
return _path_candidates(current, directories_only=True)
|
||||
if option == "--output" and (cmd.binary or cmd.path == "/audit"):
|
||||
return _path_candidates(current)
|
||||
if option == "--data" and cmd.method in ("POST", "PUT", "PATCH"):
|
||||
if not current:
|
||||
return ["@"]
|
||||
if current.startswith("@"):
|
||||
return _path_candidates(current, marker="@")
|
||||
return []
|
||||
return None
|
||||
|
||||
|
||||
def _path_candidates(
|
||||
value: str,
|
||||
*,
|
||||
directories_only: bool = False,
|
||||
marker: str = "",
|
||||
) -> list[str]:
|
||||
raw = value.removeprefix(marker) if marker else value
|
||||
ends_with_separator = raw.endswith(("/", "\\"))
|
||||
expanded = Path(raw or ".").expanduser()
|
||||
directory = expanded if ends_with_separator else expanded.parent
|
||||
name_prefix = "" if ends_with_separator else expanded.name
|
||||
raw_base = raw if ends_with_separator else raw[: len(raw) - len(name_prefix)]
|
||||
try:
|
||||
entries = directory.iterdir()
|
||||
matches = [
|
||||
entry
|
||||
for entry in entries
|
||||
if entry.name.startswith(name_prefix) and (not directories_only or entry.is_dir())
|
||||
]
|
||||
except OSError:
|
||||
return []
|
||||
|
||||
candidates: list[str] = []
|
||||
for entry in sorted(matches, key=lambda item: item.name.casefold()):
|
||||
candidate = marker + raw_base + entry.name
|
||||
if entry.is_dir():
|
||||
candidate += "/"
|
||||
candidates.append(candidate)
|
||||
return candidates
|
||||
|
||||
|
||||
def _kebab(value: str) -> str:
|
||||
output: list[str] = []
|
||||
for char in value:
|
||||
if char.isupper():
|
||||
output.extend(("-", char.lower()))
|
||||
else:
|
||||
output.append("-" if char == "_" else char)
|
||||
return "".join(output)
|
||||
|
||||
|
||||
def _matching(candidates: Any, prefix: str) -> list[str]:
|
||||
return sorted({str(candidate) for candidate in candidates if str(candidate).startswith(prefix)})
|
||||
|
||||
|
||||
def _zsh_script() -> str:
|
||||
return r"""#compdef strix
|
||||
_strix() {
|
||||
local -a candidates
|
||||
candidates=("${(@f)$($words[1] completions --candidates "${words[@]:2}")}")
|
||||
_describe 'strix' candidates
|
||||
}
|
||||
compdef _strix strix
|
||||
"""
|
||||
|
||||
|
||||
def _bash_script() -> str:
|
||||
return r"""_strix_completion() {
|
||||
local -a candidates
|
||||
local candidate
|
||||
while IFS= read -r candidate; do
|
||||
candidates+=("$candidate")
|
||||
done < <(strix completions --candidates "${COMP_WORDS[@]:1:$COMP_CWORD}")
|
||||
COMPREPLY=("${candidates[@]}")
|
||||
for candidate in "${COMPREPLY[@]}"; do
|
||||
if [[ $candidate == */ ]]; then
|
||||
if type compopt >/dev/null 2>&1; then
|
||||
compopt -o nospace
|
||||
fi
|
||||
break
|
||||
fi
|
||||
done
|
||||
}
|
||||
complete -F _strix_completion strix
|
||||
"""
|
||||
|
||||
|
||||
def _fish_script() -> str:
|
||||
return r"""function __strix_candidates
|
||||
set -l words (commandline -opc)
|
||||
set -e words[1]
|
||||
command strix completions --candidates $words (commandline -ct)
|
||||
end
|
||||
complete -c strix -f -a '(__strix_candidates)'
|
||||
"""
|
||||
@@ -62,6 +62,14 @@ import logging # noqa: E402
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ROOT_SUBCOMMAND_HELP = """
|
||||
Additional commands:
|
||||
strix cloud ... Use the managed Strix platform
|
||||
strix auth ... Manage model-subscription sign-in
|
||||
strix view [RUN] View a completed or running scan
|
||||
strix completions SHELL Generate zsh, bash, or fish tab completion
|
||||
"""
|
||||
|
||||
|
||||
def _exception_messages(exc: BaseException) -> tuple[str, ...]:
|
||||
messages: list[str] = []
|
||||
@@ -127,12 +135,10 @@ def _subscription_error_hint(exc: BaseException) -> str | None:
|
||||
|
||||
|
||||
async def warm_up_llm(show_model_warning: bool = True) -> None:
|
||||
from agents.model_settings import ModelSettings
|
||||
from agents.models.interface import ModelTracing
|
||||
|
||||
from strix.config.models import (
|
||||
RECOMMENDED_MODEL_NAMES,
|
||||
StrixProvider,
|
||||
configure_sdk_model_defaults,
|
||||
is_known_openai_bare_model,
|
||||
is_recommended_or_frontier_model,
|
||||
@@ -209,12 +215,11 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
|
||||
logger.info("LLM warm-up succeeded for model %s", (llm.model or "").strip())
|
||||
|
||||
if settings.dedupe.model:
|
||||
from strix.report.dedupe import _dedupe_extra_args
|
||||
from strix.report.dedupe import resolve_dedupe_model
|
||||
|
||||
dedupe_model = settings.dedupe.model.strip()
|
||||
raw_model = dedupe_model
|
||||
deduper = StrixProvider().get_model(dedupe_model)
|
||||
deduper_extra = _dedupe_extra_args(settings.dedupe)
|
||||
deduper = resolve_dedupe_model(settings.dedupe, dedupe_model)
|
||||
# A dedicated dedupe model may route to another provider, which must
|
||||
# never receive the main endpoint's headers; it has its own
|
||||
# DEDUPE_LLM_EXTRA_HEADERS.
|
||||
@@ -226,9 +231,6 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
|
||||
extra_headers=settings.dedupe.extra_headers,
|
||||
has_tools=False,
|
||||
)
|
||||
if deduper_extra:
|
||||
merged = {**(deduper_settings.extra_args or {}), **deduper_extra}
|
||||
deduper_settings = deduper_settings.resolve(ModelSettings(extra_args=merged))
|
||||
await asyncio.wait_for(
|
||||
deduper.get_response(
|
||||
system_instructions="You are a helpful assistant.",
|
||||
@@ -416,6 +418,13 @@ def main() -> None:
|
||||
if sys.platform == "win32":
|
||||
asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy())
|
||||
|
||||
if len(sys.argv) == 2 and sys.argv[1] in ("-h", "--help"):
|
||||
try:
|
||||
parse_arguments()
|
||||
except SystemExit as exc:
|
||||
Console().print(_ROOT_SUBCOMMAND_HELP.strip(), markup=False)
|
||||
raise SystemExit(exc.code) from None
|
||||
|
||||
# `strix view [<run>]` is a viewer-only subcommand, dispatched before the
|
||||
# scan argument parser (which requires a target) and before any scan setup.
|
||||
if len(sys.argv) > 1 and sys.argv[1] == "view":
|
||||
@@ -431,6 +440,19 @@ def main() -> None:
|
||||
|
||||
sys.exit(run_auth(sys.argv[2:]))
|
||||
|
||||
# Generate native shell completion scripts before scan argument parsing.
|
||||
if len(sys.argv) > 1 and sys.argv[1] in ("completion", "completions"):
|
||||
from strix.interface.completions import run_completions
|
||||
|
||||
sys.exit(run_completions(sys.argv[2:]))
|
||||
|
||||
# `strix cloud …` drives the managed platform (app.strix.ai) and exits;
|
||||
# it needs no target, Docker, or scan setup.
|
||||
if len(sys.argv) > 1 and sys.argv[1] == "cloud":
|
||||
from strix.interface.cloud import run_cloud
|
||||
|
||||
sys.exit(run_cloud(sys.argv[2:]))
|
||||
|
||||
from strix.llm.warmup import start_import_warmup
|
||||
|
||||
start_import_warmup()
|
||||
|
||||
798
strix/interface/platform_cli.py
Normal file
798
strix/interface/platform_cli.py
Normal file
@@ -0,0 +1,798 @@
|
||||
"""`strix cloud login` — managed platform sign-in (app.strix.ai).
|
||||
|
||||
Signing in runs an OAuth 2.0 device authorization flow in the browser, creates
|
||||
the Strix account and workspace when they do not exist yet, and stores a
|
||||
personal API token in ``~/.strix/platform-auth.json``. The token drives the
|
||||
managed REST API (scans, credits, top-ups) without a dashboard visit.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import json
|
||||
import sys
|
||||
import time
|
||||
import webbrowser
|
||||
from pathlib import Path
|
||||
from typing import Any, NoReturn, cast
|
||||
from urllib.parse import urlparse, urlsplit, urlunsplit
|
||||
|
||||
import requests
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
|
||||
from strix.config import load_settings
|
||||
from strix.interface.platform_identity import read_or_create_identity
|
||||
from strix.interface.terminal_text import sanitize_terminal_text
|
||||
from strix.interface.url_safety import is_safe_web_url
|
||||
from strix.utils.secret_files import write_secret_text
|
||||
|
||||
|
||||
AUTH_PATH = Path.home() / ".strix" / "platform-auth.json"
|
||||
|
||||
_HTTP_TIMEOUT_S = 30
|
||||
_DEFAULT_POLL_INTERVAL_S = 5
|
||||
_MAX_POLL_INTERVAL_S = 60
|
||||
_MAX_EXPIRES_IN_S = 30 * 60
|
||||
|
||||
_ROLE_RANK = {"viewer": 0, "analyst": 1, "admin": 2}
|
||||
|
||||
|
||||
class PlatformAuthError(Exception):
|
||||
"""Raised when the device authorization flow fails."""
|
||||
|
||||
|
||||
class _SessionUsageError(Exception):
|
||||
"""A session subcommand received invalid arguments."""
|
||||
|
||||
|
||||
class _SessionArgumentParser(argparse.ArgumentParser):
|
||||
def error(self, message: str) -> NoReturn:
|
||||
raise _SessionUsageError(f"invalid arguments for {self.prog}: {message}")
|
||||
|
||||
|
||||
def _terminal_markup(value: object) -> str:
|
||||
return escape(sanitize_terminal_text(value))
|
||||
|
||||
|
||||
def _app_url() -> str:
|
||||
return load_settings().viewer.app_url.rstrip("/")
|
||||
|
||||
|
||||
def read_record() -> dict[str, Any] | None:
|
||||
try:
|
||||
data = json.loads(AUTH_PATH.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return None
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
record = cast("dict[str, Any]", data)
|
||||
if not record.get("api_token"):
|
||||
return None
|
||||
return record
|
||||
|
||||
|
||||
def save_record(record: dict[str, Any]) -> None:
|
||||
write_secret_text(AUTH_PATH, json.dumps(record, indent=2))
|
||||
|
||||
|
||||
def logout() -> bool:
|
||||
try:
|
||||
AUTH_PATH.unlink()
|
||||
except FileNotFoundError:
|
||||
return True
|
||||
except OSError:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def run_login(argv: list[str]) -> int:
|
||||
"""Entry point for ``strix cloud login``. Returns a process exit code."""
|
||||
console = Console()
|
||||
subcommand = argv[0] if argv else None
|
||||
|
||||
if subcommand == "status":
|
||||
return _status(console, argv[1:])
|
||||
if subcommand == "logout":
|
||||
return _logout(console, argv[1:])
|
||||
return _login(console, argv)
|
||||
|
||||
|
||||
def _login(console: Console, argv: list[str]) -> int:
|
||||
parser = argparse.ArgumentParser(prog="strix cloud login", add_help=True)
|
||||
parser.add_argument(
|
||||
"--no-browser",
|
||||
action="store_true",
|
||||
help="Do not open the browser. Print the verification URL instead.",
|
||||
)
|
||||
scope_mode = parser.add_mutually_exclusive_group()
|
||||
scope_mode.add_argument(
|
||||
"--scopes",
|
||||
nargs="+",
|
||||
metavar="SCOPE",
|
||||
default=None,
|
||||
help=(
|
||||
"API scopes for the token, for example scans:read billing:write. "
|
||||
"The server always includes a minimum scope set. "
|
||||
"Without this option, an interactive picker opens after the browser step."
|
||||
),
|
||||
)
|
||||
scope_mode.add_argument(
|
||||
"--scope-profile",
|
||||
choices=("minimal", "recommended", "full"),
|
||||
default=None,
|
||||
help="Scope profile to approve. Defaults to an interactive choice in a TTY.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--device-name",
|
||||
default=None,
|
||||
metavar="NAME",
|
||||
help="Privacy-safe label shown for this CLI session in the dashboard.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--workspace",
|
||||
metavar="WORKSPACE",
|
||||
default=None,
|
||||
help=(
|
||||
"Workspace that receives the token, by ID or by exact name. "
|
||||
"Without this option, an interactive picker opens when you have "
|
||||
"more than one workspace."
|
||||
),
|
||||
)
|
||||
previous_record = read_record()
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
except SystemExit as exc: # argparse already printed the message
|
||||
return exc.code if isinstance(exc.code, int) else 2
|
||||
|
||||
console.print()
|
||||
host = urlparse(_app_url()).netloc or _app_url()
|
||||
console.print(f"[bold]Signing in to the Strix platform[/] [dim]({_terminal_markup(host)})[/]")
|
||||
console.print(
|
||||
"[dim]This creates your account and workspace when needed, and stores an API token.[/]"
|
||||
)
|
||||
console.print()
|
||||
|
||||
try:
|
||||
record = _run_device_flow(
|
||||
console,
|
||||
open_browser=not args.no_browser,
|
||||
scopes=args.scopes,
|
||||
scope_profile=args.scope_profile,
|
||||
workspace=args.workspace,
|
||||
device_name=args.device_name,
|
||||
)
|
||||
except PlatformAuthError as exc:
|
||||
console.print(f"[red]Sign-in failed:[/] {_terminal_markup(exc)}")
|
||||
return 1
|
||||
except KeyboardInterrupt:
|
||||
console.print("\n[yellow]Sign-in cancelled.[/]")
|
||||
return 130
|
||||
|
||||
try:
|
||||
save_record(record)
|
||||
except OSError as exc:
|
||||
console.print(
|
||||
f"[red]Sign-in succeeded, but the token could not be stored:[/] {_terminal_markup(exc)}"
|
||||
)
|
||||
console.print(
|
||||
f"[dim]Check that {_terminal_markup(AUTH_PATH.parent)} is writable, "
|
||||
"then run `strix cloud login` again.[/]"
|
||||
)
|
||||
return 1
|
||||
_revoke_replaced_legacy_session(previous_record, record)
|
||||
_print_success(console, record)
|
||||
return 0
|
||||
|
||||
|
||||
def _run_device_flow( # noqa: PLR0912, PLR0915
|
||||
console: Console,
|
||||
*,
|
||||
open_browser: bool,
|
||||
scopes: list[str] | None = None,
|
||||
scope_profile: str | None = None,
|
||||
workspace: str | None = None,
|
||||
device_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
app_url = _app_url()
|
||||
interactive = workspace is not None or (
|
||||
sys.stdin.isatty() and scopes is None and scope_profile is None
|
||||
)
|
||||
try:
|
||||
identity = read_or_create_identity(device_name=device_name)
|
||||
except (OSError, ValueError) as exc:
|
||||
raise PlatformAuthError(f"could not prepare the CLI device identity: {exc}") from exc
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{app_url}/api/v1/cli/login",
|
||||
timeout=_HTTP_TIMEOUT_S,
|
||||
allow_redirects=False,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
raise PlatformAuthError(f"could not reach {app_url}: {exc}") from exc
|
||||
if not 200 <= response.status_code < 300:
|
||||
raise PlatformAuthError(_error_detail(response))
|
||||
authorization = _json_object(response)
|
||||
|
||||
user_code = str(authorization.get("user_code") or "")
|
||||
verification_uri = str(
|
||||
authorization.get("verification_uri_complete")
|
||||
or authorization.get("verification_uri")
|
||||
or ""
|
||||
)
|
||||
device_code = str(authorization.get("device_code") or "")
|
||||
expires_in = _as_positive_int(
|
||||
authorization.get("expires_in"), default=300, maximum=_MAX_EXPIRES_IN_S
|
||||
)
|
||||
interval = _as_positive_int(
|
||||
authorization.get("interval"),
|
||||
default=_DEFAULT_POLL_INTERVAL_S,
|
||||
maximum=_MAX_POLL_INTERVAL_S,
|
||||
)
|
||||
if not device_code or not verification_uri:
|
||||
raise PlatformAuthError("the server returned an incomplete device authorization")
|
||||
if not is_safe_web_url(verification_uri, trusted_origin=app_url):
|
||||
raise PlatformAuthError("the server returned an invalid verification URL")
|
||||
|
||||
console.print(
|
||||
Panel.fit(
|
||||
Text.assemble(
|
||||
("Confirmation code: ", "dim"),
|
||||
(sanitize_terminal_text(user_code), "bold cyan"),
|
||||
),
|
||||
title="Verify this device",
|
||||
)
|
||||
)
|
||||
console.print("Open this URL in your browser and confirm the code:")
|
||||
console.print(sanitize_terminal_text(verification_uri), markup=False, soft_wrap=True)
|
||||
|
||||
if open_browser:
|
||||
with contextlib.suppress(Exception):
|
||||
webbrowser.open(verification_uri)
|
||||
|
||||
console.print("[dim]Waiting for browser confirmation…[/]")
|
||||
|
||||
poll_body: dict[str, Any] = {"device_code": device_code, **identity}
|
||||
if interactive:
|
||||
poll_body["interactive"] = True
|
||||
elif scopes:
|
||||
poll_body["scopes"] = scopes
|
||||
elif scope_profile:
|
||||
poll_body["scope_profile"] = scope_profile
|
||||
|
||||
deadline = time.monotonic() + expires_in
|
||||
while time.monotonic() < deadline:
|
||||
remaining = deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
break
|
||||
time.sleep(min(interval, remaining))
|
||||
try:
|
||||
poll = requests.post(
|
||||
f"{app_url}/api/v1/cli/login/poll",
|
||||
json=poll_body,
|
||||
timeout=_HTTP_TIMEOUT_S,
|
||||
allow_redirects=False,
|
||||
)
|
||||
except requests.RequestException:
|
||||
continue
|
||||
if 200 <= poll.status_code < 300:
|
||||
return _finish_login(
|
||||
console,
|
||||
app_url,
|
||||
poll,
|
||||
scopes=scopes,
|
||||
scope_profile=scope_profile,
|
||||
workspace=workspace,
|
||||
)
|
||||
delta = _handle_poll_error(poll)
|
||||
if delta is None:
|
||||
break
|
||||
interval = min(interval + delta, _MAX_POLL_INTERVAL_S)
|
||||
|
||||
raise PlatformAuthError("the sign-in request expired. Run `strix cloud login` again.")
|
||||
|
||||
|
||||
def _handle_poll_error(poll: requests.Response) -> int | None:
|
||||
"""Return the interval increase, or None when the device code expired."""
|
||||
error = ""
|
||||
with contextlib.suppress(ValueError, AttributeError):
|
||||
error = str(poll.json().get("error", ""))
|
||||
if error == "authorization_pending":
|
||||
return 0
|
||||
if error == "slow_down":
|
||||
return 5
|
||||
if error == "access_denied":
|
||||
raise PlatformAuthError("the sign-in request was denied in the browser")
|
||||
if error == "expired_token":
|
||||
return None
|
||||
raise PlatformAuthError(_error_detail(poll))
|
||||
|
||||
|
||||
def _finish_login(
|
||||
console: Console,
|
||||
app_url: str,
|
||||
poll: requests.Response,
|
||||
*,
|
||||
scopes: list[str] | None,
|
||||
scope_profile: str | None,
|
||||
workspace: str | None,
|
||||
) -> dict[str, Any]:
|
||||
result = _json_object(poll)
|
||||
if result.get("selection_required"):
|
||||
return _complete_selection(
|
||||
console,
|
||||
app_url,
|
||||
result,
|
||||
scopes=scopes,
|
||||
scope_profile=scope_profile,
|
||||
workspace=workspace,
|
||||
)
|
||||
return _bind_login_record(_require_api_token(result), app_url)
|
||||
|
||||
|
||||
def _signed_in_record(
|
||||
response: requests.Response,
|
||||
*,
|
||||
app_url: str,
|
||||
) -> dict[str, Any]:
|
||||
return _bind_login_record(
|
||||
_require_api_token(_json_object(response)),
|
||||
app_url,
|
||||
)
|
||||
|
||||
|
||||
def _require_api_token(record: dict[str, Any]) -> dict[str, Any]:
|
||||
api_token = record.get("api_token")
|
||||
if not isinstance(api_token, str) or not api_token.strip():
|
||||
raise PlatformAuthError("the server returned a sign-in response without an API token")
|
||||
return record
|
||||
|
||||
|
||||
def _bind_login_record(record: dict[str, Any], app_url: str) -> dict[str, Any]:
|
||||
"""Bind a stored credential to its issuer and preserve its scope preference."""
|
||||
parsed = urlsplit(app_url)
|
||||
if (
|
||||
parsed.scheme not in {"http", "https"}
|
||||
or not parsed.netloc
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or parsed.query
|
||||
or parsed.fragment
|
||||
or "\\" in app_url
|
||||
or any(character.isspace() for character in app_url)
|
||||
or "%" in parsed.netloc
|
||||
):
|
||||
raise PlatformAuthError("the configured platform URL is invalid")
|
||||
bound = dict(record)
|
||||
bound["app_url"] = urlunsplit(
|
||||
(parsed.scheme.lower(), parsed.netloc.lower(), parsed.path.rstrip("/"), "", "")
|
||||
)
|
||||
preference: Any = record.get("requested_scopes", record.get("scopes"))
|
||||
preference_items = cast("list[Any]", preference)
|
||||
if isinstance(preference, list) and all(isinstance(scope, str) for scope in preference_items):
|
||||
bound["requested_scopes"] = list(dict.fromkeys(cast("list[str]", preference_items)))
|
||||
return bound
|
||||
|
||||
|
||||
def _complete_selection(
|
||||
console: Console,
|
||||
app_url: str,
|
||||
selection: dict[str, Any],
|
||||
*,
|
||||
scopes: list[str] | None,
|
||||
scope_profile: str | None,
|
||||
workspace: str | None,
|
||||
) -> dict[str, Any]:
|
||||
organizations = _dict_items(selection.get("organizations"))
|
||||
catalog = _dict_items(selection.get("scopes"))
|
||||
selection_token = str(selection.get("selection_token") or "")
|
||||
if not selection_token or not organizations:
|
||||
raise PlatformAuthError("the server returned an incomplete selection response")
|
||||
|
||||
chosen_org = _choose_workspace(console, organizations, workspace)
|
||||
role = str(chosen_org.get("role") or "admin")
|
||||
chosen_scopes = scopes
|
||||
chosen_profile = scope_profile
|
||||
if chosen_scopes is None and chosen_profile is None and sys.stdin.isatty():
|
||||
chosen_profile, chosen_scopes = _choose_scopes(console, catalog, role)
|
||||
|
||||
body: dict[str, Any] = {
|
||||
"selection_token": selection_token,
|
||||
"organization_id": chosen_org.get("id"),
|
||||
}
|
||||
if chosen_scopes is not None:
|
||||
body["scopes"] = chosen_scopes
|
||||
body["scope_profile"] = "custom"
|
||||
elif chosen_profile is not None:
|
||||
body["scope_profile"] = chosen_profile
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{app_url}/api/v1/cli/login/complete",
|
||||
json=body,
|
||||
timeout=_HTTP_TIMEOUT_S,
|
||||
allow_redirects=False,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
raise PlatformAuthError(f"could not reach {app_url}: {exc}") from exc
|
||||
if not 200 <= response.status_code < 300:
|
||||
raise PlatformAuthError(_error_detail(response))
|
||||
return _signed_in_record(
|
||||
response,
|
||||
app_url=app_url,
|
||||
)
|
||||
|
||||
|
||||
def _dict_items(value: Any) -> list[dict[str, Any]]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
items = cast("list[Any]", cast("Any", value))
|
||||
return [cast("dict[str, Any]", cast("Any", item)) for item in items if isinstance(item, dict)]
|
||||
|
||||
|
||||
def _choose_workspace(
|
||||
console: Console, organizations: list[dict[str, Any]], workspace: str | None
|
||||
) -> dict[str, Any]:
|
||||
if workspace is not None:
|
||||
wanted = workspace.strip().casefold()
|
||||
by_id = [org for org in organizations if str(org.get("id", "")).casefold() == wanted]
|
||||
if by_id:
|
||||
return by_id[0]
|
||||
by_name = [
|
||||
org for org in organizations if str(org.get("name", "")).strip().casefold() == wanted
|
||||
]
|
||||
if len(by_name) == 1:
|
||||
return by_name[0]
|
||||
if len(by_name) > 1:
|
||||
matching_ids = ", ".join(str(org.get("id", "")) for org in by_name)
|
||||
raise PlatformAuthError(
|
||||
f"multiple workspaces are named {workspace!r}; use an exact workspace ID: "
|
||||
f"{matching_ids}"
|
||||
)
|
||||
names = ", ".join(str(org.get("name", "")) for org in organizations)
|
||||
raise PlatformAuthError(f"no workspace matches {workspace!r}. Your workspaces: {names}")
|
||||
if len(organizations) == 1:
|
||||
return organizations[0]
|
||||
if not sys.stdin.isatty():
|
||||
choices = ", ".join(f"{org.get('name', '')} ({org.get('id', '')})" for org in organizations)
|
||||
raise PlatformAuthError(
|
||||
"more than one workspace is available; rerun with --workspace NAME_OR_ID. "
|
||||
f"Available workspaces: {choices}"
|
||||
)
|
||||
|
||||
console.print()
|
||||
console.print("[bold]Select a workspace for the API token:[/]")
|
||||
for index, org in enumerate(organizations, start=1):
|
||||
name = _terminal_markup(org.get("name", ""))
|
||||
org_role = _terminal_markup(org.get("role", ""))
|
||||
console.print(f" [cyan]{index}[/]. {name} [dim]({org_role})[/]")
|
||||
while True:
|
||||
answer = console.input(f"Workspace [1-{len(organizations)}] (1): ").strip() or "1"
|
||||
if answer.isdigit() and 1 <= int(answer) <= len(organizations):
|
||||
return organizations[int(answer) - 1]
|
||||
console.print("[yellow]Enter a number from the list.[/]")
|
||||
|
||||
|
||||
def _choose_scopes(
|
||||
console: Console, catalog: list[dict[str, Any]], role: str
|
||||
) -> tuple[str, list[str] | None]:
|
||||
"""Prompt for a named scope profile or a custom scope list."""
|
||||
rank = _ROLE_RANK.get(role, 2)
|
||||
allowed = [
|
||||
item for item in catalog if _ROLE_RANK.get(str(item.get("min_role", "viewer")), 0) <= rank
|
||||
]
|
||||
if not allowed:
|
||||
return "recommended", None
|
||||
|
||||
console.print()
|
||||
console.print("[bold]Select token scopes:[/]")
|
||||
console.print(
|
||||
" [cyan]1[/]. Recommended [dim](scans, findings, schedules, assets, uploads, "
|
||||
"workspace switching, billing/top-ups; no token creation)[/]"
|
||||
)
|
||||
console.print(" [cyan]2[/]. Full access [dim](every scope your role allows)[/]")
|
||||
console.print(" [cyan]3[/]. Minimal [dim](scan read/write and billing read)[/]")
|
||||
console.print(" [cyan]4[/]. Custom [dim](pick individual scopes)[/]")
|
||||
while True:
|
||||
answer = console.input("Scopes [1-4] (1): ").strip() or "1"
|
||||
if answer == "1":
|
||||
return "recommended", None
|
||||
if answer == "2":
|
||||
return "full", None
|
||||
if answer == "3":
|
||||
return "minimal", None
|
||||
if answer == "4":
|
||||
return "custom", _choose_custom_scopes(console, allowed)
|
||||
console.print("[yellow]Enter a number from 1 to 4.[/]")
|
||||
|
||||
|
||||
def _choose_custom_scopes(console: Console, allowed: list[dict[str, Any]]) -> list[str]:
|
||||
selected = {
|
||||
str(item["scope"])
|
||||
for item in allowed
|
||||
if item.get("scope") and (item.get("default") or item.get("minimum"))
|
||||
}
|
||||
while True:
|
||||
console.print()
|
||||
for index, item in enumerate(allowed, start=1):
|
||||
scope = str(item.get("scope", ""))
|
||||
mark = "[green]x[/]" if scope in selected else " "
|
||||
required = " [dim](always included)[/]" if item.get("minimum") else ""
|
||||
rendered_scope = _terminal_markup(scope)
|
||||
description = _terminal_markup(item.get("description", ""))
|
||||
console.print(
|
||||
f" [{mark}] [cyan]{index:>2}[/]. {rendered_scope}{required}"
|
||||
f"\n [dim]{description}[/]"
|
||||
)
|
||||
answer = console.input(
|
||||
"Toggle scopes by number (comma separated), or press Enter to confirm: "
|
||||
).strip()
|
||||
if not answer:
|
||||
return sorted(selected)
|
||||
for part in answer.replace(",", " ").split():
|
||||
if not part.isdigit() or not 1 <= int(part) <= len(allowed):
|
||||
console.print(
|
||||
f"[yellow]Ignored {_terminal_markup(part)!r}: not a number from the list.[/]"
|
||||
)
|
||||
continue
|
||||
item = allowed[int(part) - 1]
|
||||
scope = str(item.get("scope", ""))
|
||||
if item.get("minimum"):
|
||||
console.print(f"[yellow]{_terminal_markup(scope)} is always included.[/]")
|
||||
continue
|
||||
if scope in selected:
|
||||
selected.discard(scope)
|
||||
else:
|
||||
selected.add(scope)
|
||||
|
||||
|
||||
def _json_object(response: requests.Response) -> dict[str, Any]:
|
||||
try:
|
||||
data = response.json()
|
||||
except ValueError as exc:
|
||||
raise PlatformAuthError("the server returned a response that is not JSON") from exc
|
||||
if not isinstance(data, dict):
|
||||
raise PlatformAuthError("the server returned an unexpected response shape")
|
||||
return cast("dict[str, Any]", data)
|
||||
|
||||
|
||||
def _as_positive_int(value: Any, *, default: int, maximum: int) -> int:
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
return default
|
||||
if parsed <= 0:
|
||||
return default
|
||||
return min(parsed, maximum)
|
||||
|
||||
|
||||
def _error_detail(response: requests.Response) -> str:
|
||||
with contextlib.suppress(ValueError, AttributeError):
|
||||
detail = response.json().get("detail")
|
||||
if detail:
|
||||
return str(detail)
|
||||
return f"HTTP {response.status_code}"
|
||||
|
||||
|
||||
def _session_headers(record: dict[str, Any]) -> dict[str, str]:
|
||||
headers = {"Authorization": f"Bearer {record['api_token']}"}
|
||||
workspace_id = record.get("organization_id")
|
||||
if isinstance(workspace_id, str) and workspace_id:
|
||||
headers["X-Strix-Workspace"] = workspace_id
|
||||
return headers
|
||||
|
||||
|
||||
def _revoke_stored_session(record: dict[str, Any]) -> tuple[bool, str | None]:
|
||||
"""Revoke one server session; return (definitively_inactive, error)."""
|
||||
app_url = record.get("app_url")
|
||||
if not isinstance(app_url, str) or not app_url:
|
||||
return False, (
|
||||
"the stored sign-in has no trusted platform URL; use --local-only to remove it"
|
||||
)
|
||||
try:
|
||||
response = requests.delete(
|
||||
f"{app_url.rstrip('/')}/api/v1/cli/session",
|
||||
headers=_session_headers(record),
|
||||
timeout=_HTTP_TIMEOUT_S,
|
||||
allow_redirects=False,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
return False, f"could not revoke the remote CLI session: {exc}"
|
||||
if response.status_code in {200, 204, 401}:
|
||||
return True, None
|
||||
return False, f"could not revoke the remote CLI session: {_error_detail(response)}"
|
||||
|
||||
|
||||
def _print_logout_failure(console: Console, message: str, *, as_json: bool) -> int:
|
||||
if as_json:
|
||||
sys.stdout.write(json.dumps({"error": message, "removed": False}) + "\n")
|
||||
else:
|
||||
console.print(f"[red]Sign-out failed:[/] {_terminal_markup(message)}")
|
||||
console.print("[dim]The local token was kept so you can safely retry.[/]")
|
||||
return 1
|
||||
|
||||
|
||||
def _revoke_replaced_legacy_session(
|
||||
previous: dict[str, Any] | None, current: dict[str, Any]
|
||||
) -> None:
|
||||
"""Best-effort cleanup when the first device-aware login replaces a legacy token."""
|
||||
if not previous or previous.get("api_token") == current.get("api_token"):
|
||||
return
|
||||
if previous.get("app_url") != current.get("app_url"):
|
||||
return
|
||||
with contextlib.suppress(KeyError, requests.RequestException):
|
||||
requests.delete(
|
||||
f"{previous['app_url']}/api/v1/cli/session",
|
||||
headers=_session_headers(previous),
|
||||
timeout=_HTTP_TIMEOUT_S,
|
||||
allow_redirects=False,
|
||||
)
|
||||
|
||||
|
||||
def _print_success(console: Console, record: dict[str, Any]) -> None:
|
||||
email = record.get("email", "")
|
||||
organization = record.get("organization_name") or record.get("organization_id", "")
|
||||
console.print()
|
||||
console.print("[green]✓ Signed in to the Strix platform.[/]")
|
||||
if email:
|
||||
console.print(f" Account: [bold]{_terminal_markup(email)}[/]")
|
||||
if organization:
|
||||
console.print(f" Workspace: [bold]{_terminal_markup(organization)}[/]")
|
||||
scopes = record.get("scopes")
|
||||
if isinstance(scopes, list) and scopes:
|
||||
console.print(f" Access: [dim]{_terminal_markup(_scope_summary(record))}[/]")
|
||||
console.print(f" Token: stored in [dim]{_terminal_markup(AUTH_PATH)}[/]")
|
||||
console.print()
|
||||
console.print(
|
||||
"[dim]The managed platform is ready. Run `strix cloud` to list the commands. "
|
||||
"See https://docs.app.strix.ai for the API reference.[/]"
|
||||
)
|
||||
|
||||
|
||||
def _status(console: Console, argv: list[str]) -> int: # noqa: PLR0912
|
||||
parser = _SessionArgumentParser(
|
||||
prog="strix cloud whoami",
|
||||
description="Show the stored managed-platform account, workspace, scopes, and expiry.",
|
||||
)
|
||||
parser.add_argument("--json", action="store_true", help="Print the session as JSON.")
|
||||
parser.add_argument("--show-scopes", action="store_true", help="Print every granted scope.")
|
||||
as_json = "--json" in argv or not sys.stdout.isatty()
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
except _SessionUsageError as exc:
|
||||
if as_json:
|
||||
sys.stdout.write(json.dumps({"error": str(exc)}) + "\n")
|
||||
else:
|
||||
console.print(f"[red]Error:[/] {_terminal_markup(exc)}")
|
||||
return 2
|
||||
except SystemExit as exc:
|
||||
return exc.code if isinstance(exc.code, int) else 2
|
||||
|
||||
as_json = bool(args.json) or not sys.stdout.isatty()
|
||||
record = read_record()
|
||||
if record is None:
|
||||
if as_json:
|
||||
sys.stdout.write(json.dumps({"signed_in": False, "error": "Not signed in"}) + "\n")
|
||||
return 1
|
||||
console.print("[yellow]Not signed in.[/] Run [bold]strix cloud login[/] to sign in.")
|
||||
return 1
|
||||
email = record.get("email", "unknown")
|
||||
organization = record.get("organization_name") or record.get("organization_id", "")
|
||||
expires_at = record.get("expires_at", "")
|
||||
if as_json:
|
||||
payload = {
|
||||
"signed_in": True,
|
||||
"email": email,
|
||||
"organization_id": record.get("organization_id"),
|
||||
"organization_name": record.get("organization_name"),
|
||||
"scopes": record.get("scopes", []),
|
||||
"expires_at": expires_at or None,
|
||||
**({"app_url": record["app_url"]} if record.get("app_url") else {}),
|
||||
}
|
||||
sys.stdout.write(json.dumps(payload, indent=2, default=str) + "\n")
|
||||
return 0
|
||||
console.print(f"[green]Signed in[/] as [bold]{_terminal_markup(email)}[/]")
|
||||
if organization:
|
||||
console.print(f" Workspace: {_terminal_markup(organization)}")
|
||||
if expires_at:
|
||||
console.print(f" Token expires: {_terminal_markup(expires_at)}")
|
||||
if record.get("app_url"):
|
||||
console.print(f" Platform: {_terminal_markup(record['app_url'])}")
|
||||
scopes = record.get("scopes")
|
||||
if isinstance(scopes, list) and scopes:
|
||||
scope_items = cast("list[Any]", cast("Any", scopes))
|
||||
if args.show_scopes:
|
||||
console.print(
|
||||
f" Scopes: {_terminal_markup(' '.join(str(scope) for scope in scope_items))}"
|
||||
)
|
||||
else:
|
||||
console.print(f" Access: {_terminal_markup(_scope_summary(record))}")
|
||||
return 0
|
||||
|
||||
|
||||
def _scope_summary(record: dict[str, Any]) -> str:
|
||||
scopes = record.get("scopes")
|
||||
scope_items = cast("list[Any]", cast("Any", scopes)) if isinstance(scopes, list) else []
|
||||
count = len(scope_items)
|
||||
profile = str(record.get("scope_profile") or "custom").replace("_", " ").title()
|
||||
return f"{profile} · {count} scope{'s' if count != 1 else ''} granted"
|
||||
|
||||
|
||||
def _logout(console: Console, argv: list[str]) -> int: # noqa: PLR0911, PLR0912
|
||||
parser = _SessionArgumentParser(
|
||||
prog="strix cloud logout",
|
||||
description="Revoke this CLI session and remove its token from this machine.",
|
||||
)
|
||||
parser.add_argument("--json", action="store_true", help="Print the result as JSON.")
|
||||
parser.add_argument(
|
||||
"--local-only",
|
||||
action="store_true",
|
||||
help="Remove only the local token, leaving the remote session active.",
|
||||
)
|
||||
as_json = "--json" in argv or not sys.stdout.isatty()
|
||||
try:
|
||||
args = parser.parse_args(argv)
|
||||
except _SessionUsageError as exc:
|
||||
if as_json:
|
||||
sys.stdout.write(json.dumps({"error": str(exc)}) + "\n")
|
||||
else:
|
||||
console.print(f"[red]Error:[/] {_terminal_markup(exc)}")
|
||||
return 2
|
||||
except SystemExit as exc:
|
||||
return exc.code if isinstance(exc.code, int) else 2
|
||||
as_json = bool(args.json) or not sys.stdout.isatty()
|
||||
if read_record() is None and not AUTH_PATH.exists():
|
||||
if as_json:
|
||||
sys.stdout.write(json.dumps({"signed_in": False, "removed": False}) + "\n")
|
||||
return 0
|
||||
console.print("[yellow]Not signed in.[/]")
|
||||
return 0
|
||||
record = read_record()
|
||||
remotely_revoked = False
|
||||
if record is not None and not args.local_only:
|
||||
remotely_revoked, revoke_error = _revoke_stored_session(record)
|
||||
if revoke_error:
|
||||
return _print_logout_failure(console, revoke_error, as_json=as_json)
|
||||
|
||||
if not logout():
|
||||
if as_json:
|
||||
sys.stdout.write(
|
||||
json.dumps(
|
||||
{
|
||||
"error": "Could not remove the stored API token",
|
||||
"signed_in": True,
|
||||
"removed": False,
|
||||
}
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
return 1
|
||||
console.print(
|
||||
f"[red]Could not remove the stored API token.[/] Delete "
|
||||
f"{_terminal_markup(AUTH_PATH)} manually."
|
||||
)
|
||||
return 1
|
||||
if as_json:
|
||||
sys.stdout.write(
|
||||
json.dumps(
|
||||
{
|
||||
"signed_in": False,
|
||||
"removed": True,
|
||||
"remotely_revoked": remotely_revoked,
|
||||
"local_only": bool(args.local_only),
|
||||
}
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
return 0
|
||||
if args.local_only:
|
||||
console.print(
|
||||
"[yellow]Local sign-out only.[/] The remote CLI session is still active; "
|
||||
"revoke it from API Access if needed."
|
||||
)
|
||||
else:
|
||||
console.print("[green]Signed out.[/] The CLI session was revoked and removed locally.")
|
||||
return 0
|
||||
46
strix/interface/platform_identity.py
Normal file
46
strix/interface/platform_identity.py
Normal file
@@ -0,0 +1,46 @@
|
||||
"""Stable, privacy-safe identity for this Strix CLI installation."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import platform
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
from uuid import uuid4
|
||||
|
||||
from strix.utils.secret_files import write_secret_text
|
||||
|
||||
|
||||
IDENTITY_PATH = Path.home() / ".strix" / "cli-identity.json"
|
||||
|
||||
|
||||
def _default_device_name(instance_id: str) -> str:
|
||||
system = {"Darwin": "macOS", "Windows": "Windows", "Linux": "Linux"}.get(
|
||||
platform.system(), "Computer"
|
||||
)
|
||||
return f"{system} CLI · {instance_id[:8]}"
|
||||
|
||||
|
||||
def read_or_create_identity(*, device_name: str | None = None) -> dict[str, str]:
|
||||
"""Return one installation ID, optionally updating its user-facing label."""
|
||||
record: dict[str, Any] = {}
|
||||
try:
|
||||
raw = json.loads(IDENTITY_PATH.read_text(encoding="utf-8"))
|
||||
if isinstance(raw, dict):
|
||||
record = cast("dict[str, Any]", raw)
|
||||
except (OSError, json.JSONDecodeError):
|
||||
pass
|
||||
|
||||
instance_id = record.get("client_instance_id")
|
||||
if not isinstance(instance_id, str) or len(instance_id) < 8:
|
||||
instance_id = str(uuid4())
|
||||
label = device_name.strip() if device_name is not None else record.get("device_name")
|
||||
if not isinstance(label, str) or not label.strip():
|
||||
label = _default_device_name(instance_id)
|
||||
label = " ".join(label.split())
|
||||
if not 1 <= len(label) <= 80:
|
||||
raise ValueError("device name must be 1-80 printable characters")
|
||||
|
||||
identity = {"client_instance_id": instance_id, "device_name": label}
|
||||
write_secret_text(IDENTITY_PATH, json.dumps(identity, indent=2))
|
||||
return identity
|
||||
21
strix/interface/terminal_text.py
Normal file
21
strix/interface/terminal_text.py
Normal file
@@ -0,0 +1,21 @@
|
||||
"""Safe rendering of untrusted text in a terminal."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
|
||||
_TERMINAL_CONTROL = re.compile(r"[\x00-\x1f\x7f-\x9f]")
|
||||
|
||||
|
||||
def has_terminal_control(value: object) -> bool:
|
||||
"""Return whether text contains bytes that can alter terminal state/protocols."""
|
||||
return _TERMINAL_CONTROL.search(str(value)) is not None
|
||||
|
||||
|
||||
def sanitize_terminal_text(value: object) -> str:
|
||||
"""Make C0/C1 control bytes visible so they cannot operate a terminal."""
|
||||
return _TERMINAL_CONTROL.sub(
|
||||
lambda match: f"\\x{ord(match.group()):02x}",
|
||||
str(value),
|
||||
)
|
||||
@@ -103,6 +103,11 @@ class TuiController:
|
||||
self.messages: list[dict[str, str]] = []
|
||||
self._next_message_id = 1
|
||||
self.error: str | None = None
|
||||
# The run's MCP connection roster (name / tool_count / dead), pushed by
|
||||
# the engine via the mcp_status_sink once the connections are established
|
||||
# and again each time one dies. Empty for a run with no MCP connections,
|
||||
# so the Go sidebar simply omits the panel. Non-secret by construction.
|
||||
self.mcp_connections: list[dict[str, Any]] = []
|
||||
self.viewer_status = "idle"
|
||||
self.viewer_url: str | None = None
|
||||
self._viewer_httpd: Any = None
|
||||
@@ -128,6 +133,24 @@ class TuiController:
|
||||
if scan_loop is not None:
|
||||
self.scan_loop = scan_loop
|
||||
|
||||
def set_mcp_connections(self, roster: list[dict[str, Any]]) -> None:
|
||||
"""Store the run's MCP connection roster and repaint.
|
||||
|
||||
``roster`` is the engine's non-secret status snapshot: one entry per
|
||||
connection carrying ``name``, ``tool_count``, and ``dead``. Called once
|
||||
when the connections are established (all healthy) and again whenever a
|
||||
connection dies (the same whole-roster snapshot, with that one now dead)."""
|
||||
self.mcp_connections = [
|
||||
{
|
||||
"name": str(entry.get("name", "")),
|
||||
"tool_count": int(entry.get("tool_count", 0) or 0),
|
||||
"dead": bool(entry.get("dead", False)),
|
||||
}
|
||||
for entry in roster
|
||||
if isinstance(entry, dict) and entry.get("name")
|
||||
]
|
||||
self.notify_changed()
|
||||
|
||||
def begin_preparation(self) -> None:
|
||||
"""Mark a directly-launched run as preparing behind the live TUI."""
|
||||
self.scan_state = "preparing"
|
||||
@@ -200,6 +223,14 @@ class TuiController:
|
||||
],
|
||||
"usage": terminal_projection(usage, max_string=256, max_items=20),
|
||||
"subscription": subscription,
|
||||
"connections": [
|
||||
{
|
||||
"name": terminal_projection(entry["name"], max_string=64),
|
||||
"tool_count": entry["tool_count"],
|
||||
"dead": entry["dead"],
|
||||
}
|
||||
for entry in self.mcp_connections[:32]
|
||||
],
|
||||
"viewer_status": self.viewer_status,
|
||||
"viewer_url": terminal_projection(self.viewer_url, max_string=1024),
|
||||
"error": terminal_projection(self.error, max_string=2 * 1024),
|
||||
@@ -380,6 +411,7 @@ class TuiController:
|
||||
delivered = await asyncio.wrap_future(future)
|
||||
if not delivered:
|
||||
raise RuntimeError("Message could not be delivered")
|
||||
self.live_view.upsert_agent(agent_id, status="waiting", error_message=None)
|
||||
return {"sent": True}
|
||||
|
||||
async def _stop_agent(self, payload: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
@@ -60,6 +60,9 @@ class TuiLiveView(BaseLiveView):
|
||||
if error_message and current.get("error_message") != error_message:
|
||||
current["error_message"] = error_message
|
||||
changed = True
|
||||
elif error_message is None and "error_message" in current:
|
||||
current.pop("error_message", None)
|
||||
changed = True
|
||||
if changed:
|
||||
current["updated_at"] = now
|
||||
return changed
|
||||
|
||||
@@ -164,19 +164,21 @@ def bounded_state_projection(state: dict[str, Any]) -> dict[str, Any]:
|
||||
"scan_state": state["scan_state"],
|
||||
"targets": state["targets"][:4],
|
||||
"target_count": state["target_count"],
|
||||
"working_dir": terminal_projection(state.get("working_dir", ""), max_string=256),
|
||||
"pending_mount": terminal_projection(state.get("pending_mount", ""), max_string=256),
|
||||
"instruction": terminal_projection(state["instruction"], max_string=128),
|
||||
"scan_mode": state["scan_mode"],
|
||||
"max_budget_usd": state["max_budget_usd"],
|
||||
"max_turns": state["max_turns"],
|
||||
"scope_mode": state["scope_mode"],
|
||||
"diff_base": state["diff_base"],
|
||||
"provider": state["provider"],
|
||||
"model": state["model"],
|
||||
"model_warning": "",
|
||||
"caido_url": None,
|
||||
"messages": [],
|
||||
"usage": state["usage"],
|
||||
"subscription": state["subscription"],
|
||||
"connections": state.get("connections", [])[:32],
|
||||
"viewer_status": state["viewer_status"],
|
||||
"viewer_url": None,
|
||||
"error": terminal_projection(state["error"], max_string=256),
|
||||
|
||||
@@ -209,7 +209,7 @@ func (m *Model) ensureAgentVisible() {
|
||||
m.agentOffset = 0
|
||||
return
|
||||
}
|
||||
_, _, agentHeight := m.sidebarHeights()
|
||||
_, _, _, agentHeight := m.sidebarHeights()
|
||||
rows := max(1, agentHeight-4)
|
||||
row := selectedAgentRow(entries, m.selectedAgent)
|
||||
if row < m.agentOffset {
|
||||
@@ -221,7 +221,7 @@ func (m *Model) ensureAgentVisible() {
|
||||
}
|
||||
|
||||
func (m Model) agentPageSize() int {
|
||||
_, _, agentHeight := m.sidebarHeights()
|
||||
_, _, _, agentHeight := m.sidebarHeights()
|
||||
return max(1, agentHeight-4)
|
||||
}
|
||||
|
||||
|
||||
105
strix/interface/tui/internal/app/mcp_test.go
Normal file
105
strix/interface/tui/internal/app/mcp_test.go
Normal file
@@ -0,0 +1,105 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/charmbracelet/x/ansi"
|
||||
"github.com/usestrix/strix/tui/internal/protocol"
|
||||
)
|
||||
|
||||
func mcpModel(t *testing.T) Model {
|
||||
t.Helper()
|
||||
m := New(nil)
|
||||
m.width, m.height = 130, 40
|
||||
m.showSplash = false
|
||||
m.handleEnvelope(stateEnvelope(t, 1, protocol.Snapshot{
|
||||
ScanState: "running",
|
||||
Connections: []protocol.Connection{
|
||||
{Name: "supabase", ToolCount: 3, Dead: false},
|
||||
{Name: "vercel", ToolCount: 1, Dead: true},
|
||||
},
|
||||
}))
|
||||
return m
|
||||
}
|
||||
|
||||
func TestMcpPanelShowsHealthyAndOffline(t *testing.T) {
|
||||
m := mcpModel(t)
|
||||
out := ansi.Strip(m.mcpConnectionsView(40, 6))
|
||||
for _, want := range []string{"MCP Connections (2)", "supabase", "3 tools", "vercel", "offline"} {
|
||||
if !strings.Contains(out, want) {
|
||||
t.Fatalf("panel missing %q:\n%s", want, out)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A roster longer than the panel height shows a window of rows rather than every
|
||||
// connection, while the header keeps the full count.
|
||||
func TestMcpPanelWindowsLargeRosterAndCountsAll(t *testing.T) {
|
||||
m := New(nil)
|
||||
m.width, m.height = 130, 40
|
||||
m.showSplash = false
|
||||
conns := make([]protocol.Connection, 0, 12)
|
||||
for i := 0; i < 12; i++ {
|
||||
conns = append(conns, protocol.Connection{Name: fmt.Sprintf("conn-%02d", i), ToolCount: 2})
|
||||
}
|
||||
m.snapshot.Connections = conns
|
||||
|
||||
// rows = 6 → one header line + five roster rows.
|
||||
out := ansi.Strip(m.mcpConnectionsView(40, 6))
|
||||
if !strings.Contains(out, "MCP Connections (12)") {
|
||||
t.Fatalf("header did not carry the full connection count:\n%s", out)
|
||||
}
|
||||
if !strings.Contains(out, "conn-00") {
|
||||
t.Fatalf("top of the roster was not rendered:\n%s", out)
|
||||
}
|
||||
if strings.Contains(out, "conn-11") {
|
||||
t.Fatalf("a roster past the panel height should be windowed, not fully drawn:\n%s", out)
|
||||
}
|
||||
if got := strings.Count(out, "\n") + 1; got != 6 {
|
||||
t.Fatalf("panel rendered %d lines, want 6 (header + five rows)", got)
|
||||
}
|
||||
|
||||
// Scrolling the roster brings the tail into view while the header count holds.
|
||||
m.mcpOffset = 7
|
||||
scrolled := ansi.Strip(m.mcpConnectionsView(40, 6))
|
||||
if !strings.Contains(scrolled, "conn-11") || !strings.Contains(scrolled, "MCP Connections (12)") {
|
||||
t.Fatalf("scrolled window did not reveal the tail with the count intact:\n%s", scrolled)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMcpPanelHeightReservedFromAgentBudget(t *testing.T) {
|
||||
m := mcpModel(t)
|
||||
_, _, mcpHeight, _ := m.sidebarHeights()
|
||||
if mcpHeight <= 0 {
|
||||
t.Fatalf("connections present but no panel height was reserved: %d", mcpHeight)
|
||||
}
|
||||
|
||||
empty := New(nil)
|
||||
empty.width, empty.height = 130, 40
|
||||
empty.showSplash = false
|
||||
empty.handleEnvelope(stateEnvelope(t, 1, protocol.Snapshot{ScanState: "running"}))
|
||||
if _, _, emptyHeight, _ := empty.sidebarHeights(); emptyHeight != 0 {
|
||||
t.Fatalf("no connections should leave the panel absent, got height %d", emptyHeight)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMcpInUseReadsRunningConnectionTaggedCalls(t *testing.T) {
|
||||
m := mcpModel(t)
|
||||
m.handleEnvelope(bootstrapEnvelope(t, "events", 1,
|
||||
protocol.Event{ID: "e1", Type: "tool", AgentID: "a1", Data: map[string]any{
|
||||
"tool_name": "call_mcp", "mcp_connection": "supabase", "status": "running",
|
||||
}},
|
||||
protocol.Event{ID: "e2", Type: "tool", AgentID: "a1", Data: map[string]any{
|
||||
"tool_name": "call_mcp", "mcp_connection": "vercel", "status": "completed",
|
||||
}},
|
||||
))
|
||||
inUse := m.mcpInUse()
|
||||
if !inUse["supabase"] {
|
||||
t.Fatalf("a running connection-tagged call should mark the connection in use")
|
||||
}
|
||||
if inUse["vercel"] {
|
||||
t.Fatalf("a completed call must not mark the connection in use")
|
||||
}
|
||||
}
|
||||
@@ -73,6 +73,7 @@ const (
|
||||
focusChat
|
||||
focusAgents
|
||||
focusVulnerabilities
|
||||
focusMcp
|
||||
)
|
||||
|
||||
type scrollbarTarget int
|
||||
@@ -82,6 +83,7 @@ const (
|
||||
scrollbarTrace
|
||||
scrollbarAgents
|
||||
scrollbarFindings
|
||||
scrollbarMcp
|
||||
)
|
||||
|
||||
type Model struct {
|
||||
@@ -109,6 +111,7 @@ type Model struct {
|
||||
selectedVuln int
|
||||
agentOffset int
|
||||
vulnOffset int
|
||||
mcpOffset int
|
||||
modalChoice int
|
||||
reportFocus string
|
||||
ready bool
|
||||
@@ -356,7 +359,9 @@ func (m Model) Update(msg tea.Msg) (tea.Model, tea.Cmd) {
|
||||
m.resyncRequested[msg.collection] = false
|
||||
}
|
||||
} else if msg.command == "collection.resync" && msg.requestID != "" && msg.collection != "" {
|
||||
m.resyncRequests[msg.requestID] = msg.collection
|
||||
if m.resyncRequested[msg.collection] {
|
||||
m.resyncRequests[msg.requestID] = msg.collection
|
||||
}
|
||||
}
|
||||
case selectionCopiedMsg:
|
||||
text := "Copied to clipboard"
|
||||
|
||||
@@ -103,6 +103,21 @@ func bootstrapEnvelope(t *testing.T, collection string, revision int, items ...a
|
||||
return protocol.Envelope{Version: protocol.Version, Type: "collection_bootstrap", Payload: rawJSON(t, payload)}
|
||||
}
|
||||
|
||||
func TestStateSnapshotClearsNilError(t *testing.T) {
|
||||
model := New(nil)
|
||||
errText := "provider rejected"
|
||||
model.handleEnvelope(stateEnvelope(t, 1, protocol.Snapshot{ScanState: "failed", Error: &errText}))
|
||||
if model.errorText != errText {
|
||||
t.Fatalf("error was not installed: %q", model.errorText)
|
||||
}
|
||||
|
||||
model.handleEnvelope(stateEnvelope(t, 2, protocol.Snapshot{ScanState: "running"}))
|
||||
|
||||
if model.errorText != "" {
|
||||
t.Fatalf("nil snapshot error did not clear errorText: %q", model.errorText)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBackendDisconnectBecomesFatalUnlessUserIsQuitting(t *testing.T) {
|
||||
model := New(nil)
|
||||
updated, cmd := model.Update(wireErrMsg{err: fmt.Errorf("socket closed")})
|
||||
@@ -160,6 +175,27 @@ func TestCollectionBootstrapChunksAndVersionedDelta(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentCollectionDeltaClearsErrorMessage(t *testing.T) {
|
||||
model := New(nil)
|
||||
failed := protocol.Agent{ID: "root", Name: "Strix", Status: "failed", ErrorMessage: "provider rejected"}
|
||||
model.handleEnvelope(bootstrapEnvelope(t, "agents", 1, failed))
|
||||
|
||||
resumed := protocol.Agent{ID: "root", Name: "Strix", Status: "waiting"}
|
||||
delta := protocol.CollectionDelta{
|
||||
Collection: "agents", BaseRevision: 1, Revision: 2, Cursor: 0, NextCursor: 1, Done: true,
|
||||
Operations: []protocol.CollectionOperation{{Op: "upsert", Item: rawJSON(t, resumed)}},
|
||||
}
|
||||
model.handleEnvelope(protocol.Envelope{Version: protocol.Version, Type: "collection_delta", Payload: rawJSON(t, delta)})
|
||||
|
||||
if len(model.snapshot.Agents) != 1 {
|
||||
t.Fatalf("agents were not retained: %#v", model.snapshot.Agents)
|
||||
}
|
||||
agent := model.snapshot.Agents[0]
|
||||
if agent.Status != "waiting" || agent.ErrorMessage != "" {
|
||||
t.Fatalf("agent error was not cleared: %#v", agent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectionMismatchRequestsOneResync(t *testing.T) {
|
||||
connection := &recordingConn{}
|
||||
model := New(newClient(connection))
|
||||
@@ -184,6 +220,42 @@ func TestCollectionMismatchRequestsOneResync(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFailedResyncResultBeforeSentMsgRearmsResync(t *testing.T) {
|
||||
connection := &recordingConn{}
|
||||
model := New(newClient(connection))
|
||||
model.collectionRevisions["events"] = 4
|
||||
bad := protocol.CollectionDelta{
|
||||
Collection: "events", BaseRevision: 2, Revision: 3, Cursor: 0, NextCursor: 0, Done: true,
|
||||
}
|
||||
|
||||
cmd := model.handleEnvelope(protocol.Envelope{Version: protocol.Version, Type: "collection_delta", Payload: rawJSON(t, bad)})
|
||||
if cmd == nil {
|
||||
t.Fatal("revision mismatch did not request a resync")
|
||||
}
|
||||
sent, ok := cmd().(sentMsg)
|
||||
if !ok || sent.err != nil || sent.requestID == "" {
|
||||
t.Fatalf("resync send = %#v", sent)
|
||||
}
|
||||
|
||||
failed := protocol.CommandResult{
|
||||
OK: false,
|
||||
Command: "collection.resync",
|
||||
Error: &protocol.CommandError{Code: "command_failed", Message: "resync failed"},
|
||||
}
|
||||
model.handleEnvelope(protocol.Envelope{
|
||||
Version: protocol.Version, Type: "command_result", RequestID: sent.requestID, Payload: rawJSON(t, failed),
|
||||
})
|
||||
updated, _ := model.Update(sent)
|
||||
model = updated.(Model)
|
||||
|
||||
if model.resyncRequested["events"] {
|
||||
t.Fatal("failed resync result left resync suppressed")
|
||||
}
|
||||
if retry := model.handleEnvelope(protocol.Envelope{Version: protocol.Version, Type: "collection_delta", Payload: rawJSON(t, bad)}); retry == nil {
|
||||
t.Fatal("resync was not rearmed after failure")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentsCollectionPreservesSelectedIDAcrossUpsertsAndDeletes(t *testing.T) {
|
||||
model := New(nil)
|
||||
model.handleEnvelope(bootstrapEnvelope(t, "agents", 1,
|
||||
@@ -599,7 +671,7 @@ func TestVulnerabilityListSupportsWheelAndPageNavigation(t *testing.T) {
|
||||
})
|
||||
}
|
||||
_, _, chatWidth, _ := model.layout()
|
||||
_, _, agentHeight := model.sidebarHeights()
|
||||
_, _, _, agentHeight := model.sidebarHeights()
|
||||
pageItems := model.vulnerabilityPageItems()
|
||||
|
||||
updated, _ := model.updateMouse(tea.MouseMsg{
|
||||
@@ -865,6 +937,20 @@ func TestPanelPaddingResetsLeakingLineBackground(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestFillBackgroundRestoresBaseForegroundAfterReset(t *testing.T) {
|
||||
const textFG = "\x1b[38;2;212;212;212m"
|
||||
view := "\x1b[38;2;167;139;250m◈ \x1b[0m\x1b[2mspawning\x1b[0m"
|
||||
filled := fillBackground(view)
|
||||
baseStyle := blackBG + textFG
|
||||
|
||||
if !strings.HasPrefix(filled, baseStyle) {
|
||||
t.Fatalf("frame does not set its base colors: %q", filled)
|
||||
}
|
||||
if got, want := strings.Count(filled, "\x1b[0m"+baseStyle), 2; got != want {
|
||||
t.Fatalf("base colors restored after %d resets, want %d: %q", got, want, filled)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMainTraceTreeAndFindingsRenderScrollbars(t *testing.T) {
|
||||
model := New(nil)
|
||||
model.width, model.height = 150, 35
|
||||
@@ -909,7 +995,7 @@ func TestMainScrollbarsSupportClickAndDrag(t *testing.T) {
|
||||
model.viewport.SetContent(model.viewportContent)
|
||||
showSidebar, _, chatWidth, chatHeight := model.layout()
|
||||
viewerHeight := model.viewerHeight()
|
||||
_, vulnHeight, agentHeight := model.sidebarHeights()
|
||||
_, vulnHeight, _, agentHeight := model.sidebarHeights()
|
||||
if !showSidebar {
|
||||
t.Fatal("test requires sidebar")
|
||||
}
|
||||
@@ -953,6 +1039,61 @@ func TestMainScrollbarsSupportClickAndDrag(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMcpRosterScrollsByKeyWheelAndScrollbar(t *testing.T) {
|
||||
model := New(nil)
|
||||
model.width, model.height = 150, 35
|
||||
model.ready = true
|
||||
conns := make([]protocol.Connection, 0, 12)
|
||||
for i := 0; i < 12; i++ {
|
||||
conns = append(conns, protocol.Connection{Name: fmt.Sprintf("conn-%02d", i), ToolCount: 2})
|
||||
}
|
||||
model.snapshot.Connections = conns
|
||||
|
||||
showSidebar, _, chatWidth, _ := model.layout()
|
||||
if !showSidebar {
|
||||
t.Fatal("test requires sidebar")
|
||||
}
|
||||
viewerHeight := model.viewerHeight()
|
||||
_, vulnHeight, mcpHeight, agentHeight := model.sidebarHeights()
|
||||
mcpTop := viewerHeight + agentHeight + vulnHeight
|
||||
bottom := model.clampMcpOffset(1 << 30)
|
||||
if bottom == 0 {
|
||||
t.Fatalf("a roster of %d should overflow the panel", len(conns))
|
||||
}
|
||||
|
||||
// Wheel over the panel focuses it and advances the window.
|
||||
updated, _ := model.updateMouse(tea.MouseMsg{
|
||||
X: chatWidth + 2, Y: mcpTop + 1, Button: tea.MouseButtonWheelDown,
|
||||
})
|
||||
model = updated.(Model)
|
||||
if model.focus != focusMcp || model.mcpOffset != 3 {
|
||||
t.Fatalf("wheel scroll did not focus and advance roster: focus=%v offset=%d", model.focus, model.mcpOffset)
|
||||
}
|
||||
|
||||
// Page down pins to the bottom; up steps back one.
|
||||
updated, _ = model.updateMain(tea.KeyMsg{Type: tea.KeyPgDown})
|
||||
model = updated.(Model)
|
||||
if model.mcpOffset != bottom {
|
||||
t.Fatalf("page down did not reach the roster bottom: offset=%d want=%d", model.mcpOffset, bottom)
|
||||
}
|
||||
updated, _ = model.updateMain(tea.KeyMsg{Type: tea.KeyUp})
|
||||
model = updated.(Model)
|
||||
if model.mcpOffset != bottom-1 {
|
||||
t.Fatalf("up did not step the roster back one: offset=%d want=%d", model.mcpOffset, bottom-1)
|
||||
}
|
||||
|
||||
// Clicking the scrollbar thumb captures it and moves the window.
|
||||
model.mcpOffset = 0
|
||||
updated, _ = model.updateMouse(tea.MouseMsg{
|
||||
X: model.width - 3, Y: mcpTop + mcpHeight - 2,
|
||||
Button: tea.MouseButtonLeft, Action: tea.MouseActionPress,
|
||||
})
|
||||
model = updated.(Model)
|
||||
if model.draggingScrollbar != scrollbarMcp || model.mcpOffset == 0 {
|
||||
t.Fatalf("mcp scrollbar click failed: drag=%v offset=%d", model.draggingScrollbar, model.mcpOffset)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTerminalSnapshotWithoutAgentsDoesNotKeepLoading(t *testing.T) {
|
||||
tests := []struct {
|
||||
state string
|
||||
|
||||
@@ -59,6 +59,14 @@ func (m Model) updateMain(key tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
m.ensureVulnerabilityVisible()
|
||||
return m, nil
|
||||
}
|
||||
if m.focus == focusMcp && len(m.snapshot.Connections) > 0 {
|
||||
delta := 1
|
||||
if key.String() == "up" {
|
||||
delta = -1
|
||||
}
|
||||
m.mcpOffset = m.clampMcpOffset(m.mcpOffset + delta)
|
||||
return m, nil
|
||||
}
|
||||
case "enter", " ":
|
||||
if m.focus == focusVulnerabilities && len(m.snapshot.Vulnerabilities) > 0 {
|
||||
if key.String() == "enter" {
|
||||
@@ -94,6 +102,10 @@ func (m Model) updateMain(key tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
m.ensureVulnerabilityVisible()
|
||||
return m, nil
|
||||
}
|
||||
if m.focus == focusMcp && len(m.snapshot.Connections) > 0 {
|
||||
m.mcpOffset = m.clampMcpOffset(m.mcpOffset - m.mcpPageSize())
|
||||
return m, nil
|
||||
}
|
||||
m.focus = focusChat
|
||||
m.input.Blur()
|
||||
m.followOutput = false
|
||||
@@ -105,6 +117,10 @@ func (m Model) updateMain(key tea.KeyMsg) (tea.Model, tea.Cmd) {
|
||||
m.ensureVulnerabilityVisible()
|
||||
return m, nil
|
||||
}
|
||||
if m.focus == focusMcp && len(m.snapshot.Connections) > 0 {
|
||||
m.mcpOffset = m.clampMcpOffset(m.mcpOffset + m.mcpPageSize())
|
||||
return m, nil
|
||||
}
|
||||
m.focus = focusChat
|
||||
m.input.Blur()
|
||||
m.viewport.HalfViewDown()
|
||||
@@ -147,10 +163,10 @@ func (m Model) updateMouse(msg tea.MouseMsg) (tea.Model, tea.Cmd) {
|
||||
}
|
||||
showSidebar, _, chatWidth, chatHeight := m.layout()
|
||||
viewerHeight := m.viewerHeight()
|
||||
_, vulnHeight, agentHeight := m.sidebarHeights()
|
||||
_, vulnHeight, mcpHeight, agentHeight := m.sidebarHeights()
|
||||
x, y := msg.X, msg.Y
|
||||
if m.updateMainScrollbarMouse(
|
||||
msg, showSidebar, chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight,
|
||||
msg, showSidebar, chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight, mcpHeight,
|
||||
) {
|
||||
return m, nil
|
||||
}
|
||||
@@ -196,6 +212,10 @@ func (m Model) updateMouse(msg tea.MouseMsg) (tea.Model, tea.Cmd) {
|
||||
m.input.Blur()
|
||||
m.vulnOffset = max(0, m.vulnOffset-3)
|
||||
m.keepVulnerabilitySelectionInWindow()
|
||||
case mcpHeight > 0 && y < viewerHeight+agentHeight+vulnHeight+mcpHeight:
|
||||
m.focus = focusMcp
|
||||
m.input.Blur()
|
||||
m.mcpOffset = m.clampMcpOffset(m.mcpOffset - 3)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
@@ -222,6 +242,10 @@ func (m Model) updateMouse(msg tea.MouseMsg) (tea.Model, tea.Cmd) {
|
||||
totalRows, _ := m.vulnerabilityScrollRows()
|
||||
m.vulnOffset = min(max(0, totalRows-m.vulnerabilityPageSize()), m.vulnOffset+3)
|
||||
m.keepVulnerabilitySelectionInWindow()
|
||||
case mcpHeight > 0 && y < viewerHeight+agentHeight+vulnHeight+mcpHeight:
|
||||
m.focus = focusMcp
|
||||
m.input.Blur()
|
||||
m.mcpOffset = m.clampMcpOffset(m.mcpOffset + 3)
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
@@ -303,7 +327,7 @@ func (m Model) updateMouse(msg tea.MouseMsg) (tea.Model, tea.Cmd) {
|
||||
func (m *Model) updateMainScrollbarMouse(
|
||||
msg tea.MouseMsg,
|
||||
showSidebar bool,
|
||||
chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight int,
|
||||
chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight, mcpHeight int,
|
||||
) bool {
|
||||
if msg.Action == tea.MouseActionRelease {
|
||||
if m.draggingScrollbar == scrollbarNone {
|
||||
@@ -313,18 +337,18 @@ func (m *Model) updateMainScrollbarMouse(
|
||||
return true
|
||||
}
|
||||
if msg.Action == tea.MouseActionMotion && m.draggingScrollbar != scrollbarNone {
|
||||
m.scrollFromMouse(m.draggingScrollbar, msg.Y, chatHeight, viewerHeight, agentHeight)
|
||||
m.scrollFromMouse(m.draggingScrollbar, msg.Y, chatHeight, viewerHeight, agentHeight, vulnHeight)
|
||||
return true
|
||||
}
|
||||
if msg.Action != tea.MouseActionPress || msg.Button != tea.MouseButtonLeft {
|
||||
return false
|
||||
}
|
||||
target := m.scrollbarAt(msg, showSidebar, chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight)
|
||||
target := m.scrollbarAt(msg, showSidebar, chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight, mcpHeight)
|
||||
if target == scrollbarNone {
|
||||
return false
|
||||
}
|
||||
m.draggingScrollbar = target
|
||||
m.scrollFromMouse(target, msg.Y, chatHeight, viewerHeight, agentHeight)
|
||||
m.scrollFromMouse(target, msg.Y, chatHeight, viewerHeight, agentHeight, vulnHeight)
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -341,8 +365,9 @@ func nearColumn(x, column int) bool {
|
||||
func (m Model) scrollbarAt(
|
||||
msg tea.MouseMsg,
|
||||
showSidebar bool,
|
||||
chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight int,
|
||||
chatWidth, chatHeight, viewerHeight, agentHeight, vulnHeight, mcpHeight int,
|
||||
) scrollbarTarget {
|
||||
mcpTop := viewerHeight + agentHeight + vulnHeight
|
||||
switch {
|
||||
case nearColumn(msg.X, chatWidth-2) && msg.Y >= 1 && msg.Y < chatHeight-1 &&
|
||||
m.viewport.TotalLineCount() > m.viewport.VisibleLineCount():
|
||||
@@ -358,13 +383,20 @@ func (m Model) scrollbarAt(
|
||||
if totalRows > m.vulnerabilityPageSize() {
|
||||
return scrollbarFindings
|
||||
}
|
||||
// The roster scrolls below a fixed header, so its bar starts two rows into
|
||||
// the panel (border then header) rather than one.
|
||||
case showSidebar && mcpHeight > 0 && nearColumn(msg.X, m.width-3) &&
|
||||
msg.Y >= mcpTop+2 && msg.Y < mcpTop+mcpHeight-1:
|
||||
if len(m.snapshot.Connections) > m.mcpPageSize() {
|
||||
return scrollbarMcp
|
||||
}
|
||||
}
|
||||
return scrollbarNone
|
||||
}
|
||||
|
||||
func (m *Model) scrollFromMouse(
|
||||
target scrollbarTarget,
|
||||
y, chatHeight, viewerHeight, agentHeight int,
|
||||
y, chatHeight, viewerHeight, agentHeight, vulnHeight int,
|
||||
) {
|
||||
switch target {
|
||||
case scrollbarTrace:
|
||||
@@ -390,6 +422,13 @@ func (m *Model) scrollFromMouse(
|
||||
// The offset is a row, so dragging moves the list continuously.
|
||||
m.vulnOffset = scrollbarOffset(y-viewerHeight-agentHeight-1, height, totalRows, height)
|
||||
m.keepVulnerabilitySelectionInWindow()
|
||||
case scrollbarMcp:
|
||||
height := m.mcpPageSize()
|
||||
total := len(m.snapshot.Connections)
|
||||
m.focus = focusMcp
|
||||
m.input.Blur()
|
||||
// The bar starts two rows into the panel (border then the fixed header).
|
||||
m.mcpOffset = scrollbarOffset(y-viewerHeight-agentHeight-vulnHeight-2, height, total, height)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -545,6 +584,9 @@ func (m *Model) cycleFocus(delta int) {
|
||||
if len(m.snapshot.Vulnerabilities) > 0 {
|
||||
available = append(available, focusVulnerabilities)
|
||||
}
|
||||
if len(m.snapshot.Connections) > 0 {
|
||||
available = append(available, focusMcp)
|
||||
}
|
||||
}
|
||||
idx := 0
|
||||
for i, focus := range available {
|
||||
|
||||
@@ -352,21 +352,28 @@ func (m Model) toastOverlay(view string) string {
|
||||
return strings.Join(bg, "\n")
|
||||
}
|
||||
|
||||
// blackBG is the SGR that selects a solid black background.
|
||||
const blackBG = "\x1b[48;2;0;0;0m"
|
||||
// Base frame colors are reapplied after full SGR resets so the TUI does not
|
||||
// inherit an unreadable foreground from the user's terminal profile.
|
||||
const (
|
||||
blackBG = "\x1b[48;2;0;0;0m"
|
||||
textFG = "\x1b[38;2;212;212;212m"
|
||||
baseFrameColors = blackBG + textFG
|
||||
)
|
||||
|
||||
// fillBackground paints the whole frame black like Textual's Screen background.
|
||||
// Bubble Tea has no screen compositor, so any cell the view does not explicitly
|
||||
// color shows the terminal's default background. lipgloss emits a full reset
|
||||
// (\x1b[0m) at the end of every styled span, which also clears the background, so
|
||||
// we reassert black after each reset (and at the start). Spans that set their own
|
||||
// background — inline code, selected rows, buttons — keep it, because their color
|
||||
// is emitted before the reset.
|
||||
// (\x1b[0m) at the end of every styled span, which clears both foreground and
|
||||
// background. Reasserting only black made uncolored and faint text inherit the
|
||||
// terminal profile's foreground; light profiles therefore rendered that text
|
||||
// black-on-black. Reapply both base colors after each reset (and at the start).
|
||||
// Spans that set their own colors — inline code, selected rows, buttons — keep
|
||||
// them, because their color is emitted after the base style.
|
||||
func fillBackground(view string) string {
|
||||
if view == "" {
|
||||
return view
|
||||
}
|
||||
return blackBG + strings.ReplaceAll(view, "\x1b[0m", "\x1b[0m"+blackBG)
|
||||
return baseFrameColors + strings.ReplaceAll(view, "\x1b[0m", "\x1b[0m"+baseFrameColors)
|
||||
}
|
||||
|
||||
func (m Model) splashView() string {
|
||||
@@ -500,7 +507,7 @@ func (m Model) mainView() string {
|
||||
func (m Model) sidebarView(width, height int) string {
|
||||
// Stats box height fits its content (auto, max 15); vulns panel max-height 12.
|
||||
statsBody := m.statsView()
|
||||
statsHeight, vulnHeight, agentHeight := m.sidebarHeights()
|
||||
statsHeight, vulnHeight, mcpHeight, agentHeight := m.sidebarHeights()
|
||||
agentBorder := dark
|
||||
if m.focus == focusAgents {
|
||||
agentBorder = green
|
||||
@@ -539,11 +546,19 @@ func (m Model) sidebarView(width, height int) string {
|
||||
)
|
||||
parts = append(parts, lipgloss.NewStyle().Width(width-2).Height(vulnRows).Border(lipgloss.RoundedBorder()).BorderForeground(vulnBorder).Padding(0, 1).Render(findings))
|
||||
}
|
||||
if mcpHeight > 0 {
|
||||
mcpBorder := dark
|
||||
if m.focus == focusMcp {
|
||||
mcpBorder = green
|
||||
}
|
||||
mcpRows := max(1, mcpHeight-2)
|
||||
parts = append(parts, lipgloss.NewStyle().Width(width-2).Height(mcpRows).Border(lipgloss.RoundedBorder()).BorderForeground(mcpBorder).Padding(0, 1).Render(m.mcpConnectionsView(width-4, mcpRows)))
|
||||
}
|
||||
parts = append(parts, lipgloss.NewStyle().Width(width-2).Height(statsHeight-2).Border(lipgloss.RoundedBorder()).BorderForeground(dark).Padding(0, 1).Render(statsBody))
|
||||
return strings.Join(parts, "\n")
|
||||
}
|
||||
|
||||
func (m Model) sidebarHeights() (statsHeight, vulnHeight, agentHeight int) {
|
||||
func (m Model) sidebarHeights() (statsHeight, vulnHeight, mcpHeight, agentHeight int) {
|
||||
// Measure the stats panel the way its box will render it: a long model name
|
||||
// wraps inside the sidebar, and counting only its newlines would size the
|
||||
// box short and push the whole frame past the bottom of the terminal.
|
||||
@@ -552,7 +567,13 @@ func (m Model) sidebarHeights() (statsHeight, vulnHeight, agentHeight int) {
|
||||
if len(m.snapshot.Vulnerabilities) > 0 {
|
||||
vulnHeight = min(12, len(m.vulnerabilityRows(m.vulnerabilityListWidth()))+2)
|
||||
}
|
||||
agentHeight = max(3, m.height-m.viewerHeight()-statsHeight-vulnHeight)
|
||||
// One header line + one line per connection + the box border (2). Capped so a
|
||||
// long roster cannot crowd out the agent tree; a roster past the cap scrolls
|
||||
// inside the panel. Absent entirely when the run has no MCP connections.
|
||||
if len(m.snapshot.Connections) > 0 {
|
||||
mcpHeight = min(9, len(m.snapshot.Connections)+3)
|
||||
}
|
||||
agentHeight = max(3, m.height-m.viewerHeight()-statsHeight-vulnHeight-mcpHeight)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -621,6 +642,111 @@ func (m Model) statsView() string {
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// mcpConnectionsView renders the sidebar MCP panel: a header carrying the total
|
||||
// connection count, then one row per connection with a status glyph and its tool
|
||||
// count (or "offline").
|
||||
// - a solid green dot marks an attached, idle connection;
|
||||
// - a green cycling quarter-circle (◐ ◓ ◑ ◒) marks a call running against it;
|
||||
// - a red dot plus "offline" marks a connection whose live session has died.
|
||||
//
|
||||
// The header stays fixed while the roster below it scrolls: when there are more
|
||||
// connections than the panel can show, the visible window is chosen by
|
||||
// m.mcpOffset and withVerticalScrollbar draws a thumb in the reserved last
|
||||
// column, exactly as the agent tree and findings list scroll.
|
||||
//
|
||||
// "In use" is derived from the connection-tagged tool-call events in the stream,
|
||||
// not carried on the connection roster, so a call in flight shows motion without
|
||||
// any extra backend signal. The quarter-circle rides the shared sweepFrame tick.
|
||||
func (m Model) mcpConnectionsView(width, rows int) string {
|
||||
conns := m.snapshot.Connections
|
||||
header := truncate(lipgloss.NewStyle().Foreground(dim).Render(
|
||||
fmt.Sprintf("MCP Connections (%d)", len(conns))), width)
|
||||
bodyRows := max(0, rows-1)
|
||||
if bodyRows == 0 {
|
||||
return header
|
||||
}
|
||||
inUse := m.mcpInUse()
|
||||
frames := []rune{'◐', '◓', '◑', '◒'}
|
||||
// Reserve the scrollbar column whether or not the bar is showing, so the
|
||||
// roster does not shift sideways as it grows past the panel.
|
||||
rosterWidth := max(1, width-1)
|
||||
start := windowStart(m.mcpOffset, len(conns), bodyRows)
|
||||
end := min(len(conns), start+bodyRows)
|
||||
lines := make([]string, 0, max(0, end-start))
|
||||
for i := start; i < end; i++ {
|
||||
conn := conns[i]
|
||||
var glyph, right string
|
||||
switch {
|
||||
case conn.Dead:
|
||||
glyph = lipgloss.NewStyle().Foreground(red).Render("●")
|
||||
right = lipgloss.NewStyle().Foreground(red).Render("offline")
|
||||
case inUse[conn.Name]:
|
||||
glyph = lipgloss.NewStyle().Foreground(green).Render(string(frames[m.sweepFrame%len(frames)]))
|
||||
right = lipgloss.NewStyle().Foreground(dim).Render(toolsLabel(conn.ToolCount))
|
||||
default:
|
||||
glyph = lipgloss.NewStyle().Foreground(green).Render("●")
|
||||
right = lipgloss.NewStyle().Foreground(dim).Render(toolsLabel(conn.ToolCount))
|
||||
}
|
||||
rightWidth := lipgloss.Width(right)
|
||||
name := truncate(lipgloss.NewStyle().Foreground(textColor).Render(conn.Name), max(1, rosterWidth-2-rightWidth-1))
|
||||
gap := max(1, rosterWidth-2-lipgloss.Width(name)-rightWidth)
|
||||
lines = append(lines, glyph+" "+name+strings.Repeat(" ", gap)+right)
|
||||
}
|
||||
roster := withVerticalScrollbar(
|
||||
strings.Join(lines, "\n"),
|
||||
width,
|
||||
bodyRows,
|
||||
len(conns),
|
||||
bodyRows,
|
||||
m.mcpOffset,
|
||||
m.scrollbarThumb(scrollbarMcp),
|
||||
)
|
||||
return header + "\n" + roster
|
||||
}
|
||||
|
||||
// mcpPageSize is how many connection rows the roster shows at once, below its
|
||||
// fixed header line.
|
||||
func (m Model) mcpPageSize() int {
|
||||
_, _, mcpHeight, _ := m.sidebarHeights()
|
||||
// mcpHeight = 2 (border) + header (1) + roster rows.
|
||||
return max(1, mcpHeight-3)
|
||||
}
|
||||
|
||||
// clampMcpOffset keeps the roster offset within the range that still shows a
|
||||
// full page of connections at the bottom.
|
||||
func (m Model) clampMcpOffset(offset int) int {
|
||||
return min(max(0, offset), max(0, len(m.snapshot.Connections)-m.mcpPageSize()))
|
||||
}
|
||||
|
||||
// mcpInUse is the set of MCP connections with a tool call currently running,
|
||||
// read off the connection-tagged tool events the model already holds. Each MCP
|
||||
// dispatch event carries the connection name (mcp_connection) and a status that
|
||||
// moves running -> completed as its own event is upserted, so a connection is
|
||||
// "in use" exactly while one of its events is still running.
|
||||
func (m Model) mcpInUse() map[string]bool {
|
||||
inUse := map[string]bool{}
|
||||
for _, event := range m.snapshot.Events {
|
||||
if event.Type != "tool" {
|
||||
continue
|
||||
}
|
||||
connection := render.StringValue(event.Data["mcp_connection"])
|
||||
if connection == "" {
|
||||
continue
|
||||
}
|
||||
if render.StringValue(event.Data["status"]) == "running" {
|
||||
inUse[connection] = true
|
||||
}
|
||||
}
|
||||
return inUse
|
||||
}
|
||||
|
||||
func toolsLabel(count int) string {
|
||||
if count == 1 {
|
||||
return "1 tool"
|
||||
}
|
||||
return fmt.Sprintf("%d tools", count)
|
||||
}
|
||||
|
||||
func numberValue(value any) int64 {
|
||||
switch v := value.(type) {
|
||||
case float64:
|
||||
|
||||
@@ -134,7 +134,7 @@ func clampVulnerabilityOffset(offset, total, height int) int {
|
||||
}
|
||||
|
||||
func (m Model) vulnerabilityPageSize() int {
|
||||
_, vulnHeight, _ := m.sidebarHeights()
|
||||
_, vulnHeight, _, _ := m.sidebarHeights()
|
||||
return max(1, vulnHeight-2)
|
||||
}
|
||||
|
||||
|
||||
@@ -33,6 +33,8 @@ func (m *Model) handleEnvelope(envelope protocol.Envelope) tea.Cmd {
|
||||
m.stateRevision = update.Revision
|
||||
if m.snapshot.Error != nil {
|
||||
m.errorText = *m.snapshot.Error
|
||||
} else {
|
||||
m.errorText = ""
|
||||
}
|
||||
if m.snapshot.SetupMode {
|
||||
// The start screen is its own landing page; never sit on the
|
||||
@@ -79,6 +81,10 @@ func (m *Model) handleEnvelope(envelope protocol.Envelope) tea.Cmd {
|
||||
if collection := m.resyncRequests[envelope.RequestID]; collection != "" {
|
||||
m.resyncRequested[collection] = false
|
||||
delete(m.resyncRequests, envelope.RequestID)
|
||||
} else {
|
||||
for collection := range m.resyncRequested {
|
||||
m.resyncRequested[collection] = false
|
||||
}
|
||||
}
|
||||
}
|
||||
message := "Command failed"
|
||||
|
||||
@@ -32,6 +32,17 @@ type Agent struct {
|
||||
ErrorMessage string `json:"error_message"`
|
||||
}
|
||||
|
||||
// Connection is one MCP connection the run may reach, as the backend projects
|
||||
// it for the sidebar's MCP panel. Non-secret by construction: only the display
|
||||
// name, how many tools the connection offers, and whether its live session has
|
||||
// died (its reconnect-retry gave up). "In use" is not carried here; the client
|
||||
// derives it from the connection-tagged tool-call events in the event stream.
|
||||
type Connection struct {
|
||||
Name string `json:"name"`
|
||||
ToolCount int `json:"tool_count"`
|
||||
Dead bool `json:"dead"`
|
||||
}
|
||||
|
||||
type Event struct {
|
||||
ID string `json:"id"`
|
||||
Type string `json:"type"`
|
||||
@@ -68,6 +79,7 @@ type Snapshot struct {
|
||||
Vulnerabilities []map[string]any `json:"-"`
|
||||
Usage map[string]any `json:"usage"`
|
||||
Subscription bool `json:"subscription"`
|
||||
Connections []Connection `json:"connections"`
|
||||
ViewerStatus string `json:"viewer_status"`
|
||||
ViewerURL *string `json:"viewer_url"`
|
||||
Error *string `json:"error"`
|
||||
|
||||
@@ -96,14 +96,12 @@ func TestCoverageDuplicateRejectionSurfacesTheError(t *testing.T) {
|
||||
requireContains(t, out, "/login", "already has coverage entry a1b2c3")
|
||||
}
|
||||
|
||||
func TestGetThreatModelRendersStalenessAndAmendments(t *testing.T) {
|
||||
func TestGetThreatModelRendersAmendments(t *testing.T) {
|
||||
out := ansi.Strip(Tool(tool("get_threat_model",
|
||||
map[string]any{"target": "https://app.example.com"},
|
||||
map[string]any{
|
||||
"success": true,
|
||||
"found": true,
|
||||
"stale": true,
|
||||
"cached_revision": "0123456789abcdef",
|
||||
"success": true,
|
||||
"found": true,
|
||||
"content": "# Overview\nMulti-tenant billing app.\n\n" +
|
||||
"## Trust Boundaries and Assumptions\n\n## Attack Surface\n",
|
||||
"amendments": []any{
|
||||
@@ -116,7 +114,6 @@ func TestGetThreatModelRendersStalenessAndAmendments(t *testing.T) {
|
||||
"completed")))
|
||||
requireContains(t, out,
|
||||
"Threat Model", "https://app.example.com",
|
||||
"stale", "01234567",
|
||||
"1 amendment(s)", "ReconAgent", "staging host shares the production database",
|
||||
"Multi-tenant billing app.", "Overview", "Trust Boundaries and Assumptions",
|
||||
)
|
||||
@@ -126,7 +123,7 @@ func TestGetThreatModelMissingModelIsExplicit(t *testing.T) {
|
||||
out := ansi.Strip(Tool(tool("get_threat_model",
|
||||
map[string]any{"target": "10.0.0.5"},
|
||||
map[string]any{"success": true, "found": false}, "completed")))
|
||||
requireContains(t, out, "No model cached for this target yet")
|
||||
requireContains(t, out, "No model derived for this target yet")
|
||||
}
|
||||
|
||||
func TestSaveThreatModelWarnsWhenAmendmentsAreCleared(t *testing.T) {
|
||||
@@ -134,15 +131,10 @@ func TestSaveThreatModelWarnsWhenAmendmentsAreCleared(t *testing.T) {
|
||||
map[string]any{"target": "app.example.com", "content": "# Overview\nA thing.\n"},
|
||||
map[string]any{
|
||||
"success": true,
|
||||
"revision": "unversioned",
|
||||
"amendments_cleared": 2,
|
||||
},
|
||||
"completed")))
|
||||
requireContains(t, out, "Threat Model Saved", "saved", "cleared 2 amendment(s)")
|
||||
// An unversioned target has no revision worth printing.
|
||||
if strings.Contains(out, "unversioned") {
|
||||
t.Fatalf("unversioned revision should not be rendered:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAmendThreatModelRendersAddendum(t *testing.T) {
|
||||
@@ -202,3 +194,28 @@ func TestVulnerabilityReportRendersCalibrationFields(t *testing.T) {
|
||||
"Fix Verification", "bypass review reasoned only",
|
||||
)
|
||||
}
|
||||
|
||||
func TestVulnerabilityReportUpdateRendersReportAndReason(t *testing.T) {
|
||||
out := ansi.Strip(Tool(tool("update_vulnerability_report",
|
||||
map[string]any{
|
||||
"report_id": "vuln-0009",
|
||||
"update_reason": "built a working unauthenticated file write against the endpoint",
|
||||
"poc_script_code": "curl -X PATCH https://target/files/uuid",
|
||||
},
|
||||
map[string]any{
|
||||
"success": true,
|
||||
"action": "updated",
|
||||
"report_id": "vuln-0009",
|
||||
"severity": "critical",
|
||||
"cvss_score": 9.3,
|
||||
"updated_fields": []any{"poc_script_code"},
|
||||
},
|
||||
"completed")))
|
||||
requireContains(t, out,
|
||||
"Vulnerability Report Updated",
|
||||
"vuln-0009",
|
||||
"built a working unauthenticated file write",
|
||||
"CRITICAL",
|
||||
"9.3",
|
||||
)
|
||||
}
|
||||
|
||||
@@ -45,3 +45,51 @@ func renderMcpInspect(connection, status string) string {
|
||||
b.WriteString(style.Render(icon))
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// renderMcpList renders list_mcps: the inventory of connections the run may
|
||||
// reach, not a call to any of them, so no connection leads and the event
|
||||
// carries no connection tag. Unlike the other MCP results, the names are worth
|
||||
// showing: Strix assembled them itself from the run's registered connections,
|
||||
// so they are short and never an outside server's payload.
|
||||
func renderMcpList(result any, status string) string {
|
||||
var b strings.Builder
|
||||
b.WriteString(mcpIcon + Dim().Render("Listing MCP servers") + "\n")
|
||||
for _, conn := range mcpConnectionEntries(result) {
|
||||
b.WriteString(" " + Col(Slate).Render(conn.name))
|
||||
if conn.dead {
|
||||
b.WriteString(Dim().Render(" · ") + Col(Red).Render("offline"))
|
||||
}
|
||||
b.WriteString("\n")
|
||||
}
|
||||
icon, style := statusIcon(status)
|
||||
b.WriteString(style.Render(icon))
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// mcpListEntry is one connection read out of a list_mcps result: its display
|
||||
// name and whether its live session has died.
|
||||
type mcpListEntry struct {
|
||||
name string
|
||||
dead bool
|
||||
}
|
||||
|
||||
// mcpConnectionEntries reads the connections out of a list_mcps result, which is
|
||||
// {"connections": [{"name": ..., "dead": ...}, ...]}. Anything else (still
|
||||
// running, or a result bounded down to a string) yields no entries, and the
|
||||
// header plus status stand alone.
|
||||
func mcpConnectionEntries(result any) []mcpListEntry {
|
||||
resultMap, _ := result.(map[string]any)
|
||||
connections, _ := resultMap["connections"].([]any)
|
||||
var entries []mcpListEntry
|
||||
for _, raw := range connections {
|
||||
entry, ok := raw.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if name := strings.TrimSpace(StringValue(entry["name"])); name != "" {
|
||||
dead, _ := entry["dead"].(bool)
|
||||
entries = append(entries, mcpListEntry{name: name, dead: dead})
|
||||
}
|
||||
}
|
||||
return entries
|
||||
}
|
||||
|
||||
@@ -70,6 +70,10 @@ func Tool(data map[string]any) string {
|
||||
}
|
||||
|
||||
switch name {
|
||||
// list_mcps inventories every connection rather than touching one, so it is
|
||||
// the one MCP tool with no connection tag and routes by name like a built-in.
|
||||
case "list_mcps":
|
||||
return renderMcpList(result, status)
|
||||
case "exec_command":
|
||||
return renderExecCommand(args, result, status)
|
||||
case "write_stdin":
|
||||
@@ -80,6 +84,8 @@ func Tool(data map[string]any) string {
|
||||
return renderViewImage(args, result)
|
||||
case "create_vulnerability_report":
|
||||
return renderVulnerabilityReport(args, result)
|
||||
case "update_vulnerability_report":
|
||||
return renderVulnerabilityReportUpdate(args, result)
|
||||
case "create_dependency_report":
|
||||
return renderDependencyReport(args, result)
|
||||
case "list_reports":
|
||||
|
||||
@@ -268,6 +268,24 @@ func TestMcpDescribeInspectsConnection(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMcpListMarksDeadConnectionsOffline(t *testing.T) {
|
||||
// list_mcps carries a per-connection dead flag; a dead connection reads as
|
||||
// offline in the inventory while a live one shows normally.
|
||||
result := map[string]any{
|
||||
"connections": []any{
|
||||
map[string]any{"name": "supabase", "tool_count": float64(3), "dead": false},
|
||||
map[string]any{"name": "vercel", "tool_count": float64(1), "dead": true},
|
||||
},
|
||||
}
|
||||
data := tool("list_mcps", nil, result, "completed")
|
||||
|
||||
out := ansi.Strip(Tool(data))
|
||||
requireContains(t, out, "Listing MCP servers", "supabase", "vercel", "offline")
|
||||
if strings.Count(out, "offline") != 1 {
|
||||
t.Fatalf("only the dead connection should read offline:\n%s", out)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollapseToolShellPreviewAndExpand(t *testing.T) {
|
||||
lines := make([]string, 16)
|
||||
for i := range lines {
|
||||
|
||||
@@ -12,15 +12,27 @@ import (
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func renderVulnerabilityReport(args map[string]any, result any) string {
|
||||
return renderReport(args, result, "Vulnerability Report", "Creating report...")
|
||||
}
|
||||
|
||||
// A revision names the report it changes and carries only the fields it
|
||||
// replaces, so it renders the same sections with the ones it left alone absent.
|
||||
func renderVulnerabilityReportUpdate(args map[string]any, result any) string {
|
||||
return renderReport(args, result, "Vulnerability Report Updated", "Updating report...")
|
||||
}
|
||||
|
||||
func renderReport(args map[string]any, result any, heading, pending string) string {
|
||||
resultMap, _ := result.(map[string]any)
|
||||
var b strings.Builder
|
||||
b.WriteString("🐞 " + Bold(ReportHdr).Render("Vulnerability Report"))
|
||||
b.WriteString("🐞 " + Bold(ReportHdr).Render(heading))
|
||||
|
||||
field := func(label, value string) {
|
||||
if value != "" {
|
||||
b.WriteString("\n\n" + Bold(Field).Render(label+": ") + value)
|
||||
}
|
||||
}
|
||||
reportID := StringValue(args["report_id"])
|
||||
field("Report", reportID)
|
||||
title := StringValue(args["title"])
|
||||
field("Title", title)
|
||||
|
||||
@@ -59,6 +71,7 @@ func renderVulnerabilityReport(args map[string]any, result any) string {
|
||||
}
|
||||
}
|
||||
|
||||
section("Reason", StringValue(args["update_reason"]))
|
||||
section("Description", StringValue(args["description"]))
|
||||
section("Impact", StringValue(args["impact"]))
|
||||
section("Technical Analysis", StringValue(args["technical_analysis"]))
|
||||
@@ -76,8 +89,8 @@ func renderVulnerabilityReport(args map[string]any, result any) string {
|
||||
// was verified belongs next to it rather than in the artifact alone.
|
||||
section("Fix Verification", StringValue(args["fix_verification"]))
|
||||
|
||||
if title == "" {
|
||||
b.WriteString("\n " + Dim().Render("Creating report..."))
|
||||
if title == "" && reportID == "" {
|
||||
b.WriteString("\n " + Dim().Render(pending))
|
||||
}
|
||||
return "\n\n" + b.String() + "\n\n"
|
||||
}
|
||||
|
||||
@@ -56,9 +56,6 @@ func renderThreatModel(name string, args map[string]any, result any) string {
|
||||
threatModelBody(&b, StringValue(args["addendum"]))
|
||||
default:
|
||||
b.WriteString("\n " + Col(Green).Render("✓ saved"))
|
||||
if revision := shortRevision(StringValue(m["revision"])); revision != "" {
|
||||
b.WriteString(Dim().Render(" at " + revision))
|
||||
}
|
||||
// Saving folds amendments away, so the count that vanished is worth
|
||||
// stating: it is the one destructive thing this tool does.
|
||||
if cleared, ok := NumericValue(m["amendments_cleared"]); ok && cleared > 0 {
|
||||
@@ -72,15 +69,9 @@ func renderThreatModel(name string, args map[string]any, result any) string {
|
||||
|
||||
func threatModelReadBody(b *strings.Builder, result map[string]any) {
|
||||
if !truthy(result["found"]) {
|
||||
b.WriteString("\n " + Dim().Render("No model cached for this target yet"))
|
||||
b.WriteString("\n " + Dim().Render("No model derived for this target yet"))
|
||||
return
|
||||
}
|
||||
if truthy(result["stale"]) {
|
||||
b.WriteString("\n " + Col(AmberY).Render("⚠ stale"))
|
||||
if cached := shortRevision(StringValue(result["cached_revision"])); cached != "" {
|
||||
b.WriteString(Dim().Render(" (written at " + cached + ")"))
|
||||
}
|
||||
}
|
||||
if amendments, ok := result["amendments"].([]any); ok && len(amendments) > 0 {
|
||||
b.WriteString("\n " + Col(Gold).Render("+ "+strconv.Itoa(len(amendments))+
|
||||
" amendment(s)") + Dim().Render(" — later statements win"))
|
||||
@@ -126,13 +117,3 @@ func threatModelBody(b *strings.Builder, content string) {
|
||||
b.WriteString("\n " + Dim().Render(strings.Join(headings, " · ")))
|
||||
}
|
||||
}
|
||||
|
||||
// shortRevision abbreviates a git sha; "unversioned" targets have no revision
|
||||
// worth showing.
|
||||
func shortRevision(revision string) string {
|
||||
revision = strings.TrimSpace(revision)
|
||||
if revision == "" || revision == "unversioned" {
|
||||
return ""
|
||||
}
|
||||
return firstN(revision, 8)
|
||||
}
|
||||
|
||||
@@ -103,6 +103,7 @@ class TuiLiveView:
|
||||
statuses = agents_data.get("statuses") or {}
|
||||
names = agents_data.get("names") or {}
|
||||
parent_of = agents_data.get("parent_of") or {}
|
||||
errors = agents_data.get("errors") or {}
|
||||
if not isinstance(statuses, dict):
|
||||
return
|
||||
for agent_id, status in statuses.items():
|
||||
@@ -113,6 +114,7 @@ class TuiLiveView:
|
||||
name=names.get(agent_id, agent_id) if isinstance(names, dict) else agent_id,
|
||||
parent_id=parent_of.get(agent_id) if isinstance(parent_of, dict) else None,
|
||||
status=str(status),
|
||||
error_message=errors.get(agent_id) if isinstance(errors, dict) else None,
|
||||
)
|
||||
# Ahead of the replayed history, so it opens the transcript.
|
||||
self.flush_user_instruction()
|
||||
|
||||
@@ -102,6 +102,9 @@ class GoTuiRuntime:
|
||||
self.report_state.vulnerability_found_callback = lambda _report: (
|
||||
self.controller.notify_changed()
|
||||
)
|
||||
self.report_state.vulnerability_updated_callback = lambda _report: (
|
||||
self.controller.notify_changed()
|
||||
)
|
||||
self.controller.notify_changed()
|
||||
|
||||
async def start_from_setup(self, verify: bool = True) -> None:
|
||||
@@ -185,6 +188,7 @@ class GoTuiRuntime:
|
||||
max_turns=self.args.max_turns,
|
||||
max_budget_usd=self.args.max_budget_usd,
|
||||
event_sink=self.capture_event,
|
||||
mcp_status_sink=self.capture_mcp_status,
|
||||
)
|
||||
await self._sync_agent_state()
|
||||
if self.controller.scan_state == "running":
|
||||
@@ -210,6 +214,15 @@ class GoTuiRuntime:
|
||||
self.live_view.ingest_sdk_event(agent_id, event)
|
||||
self.controller.notify_changed()
|
||||
|
||||
def capture_mcp_status(self, roster: list[dict[str, Any]]) -> None:
|
||||
"""Receive the engine's MCP connection roster and hand it to the controller.
|
||||
|
||||
Runs on the scan's event loop (called from the runner at establishment
|
||||
and from a session's on-dead callback), the same loop that drives
|
||||
``capture_event``, so updating the controller and repainting here is
|
||||
safe. The controller renders it as the sidebar MCP connections panel."""
|
||||
self.controller.set_mcp_connections(roster)
|
||||
|
||||
async def _sync_agent_state(self) -> bool:
|
||||
parent_of, statuses, names, errors = await self.coordinator.graph_snapshot()
|
||||
changed = False
|
||||
@@ -248,6 +261,9 @@ class GoTuiRuntime:
|
||||
scan_state = "failed"
|
||||
if root_id is not None and errors.get(root_id):
|
||||
self.controller.error = errors[root_id]
|
||||
elif scan_state == "failed" and root_status in {"running", "waiting", "budget_paused"}:
|
||||
scan_state = "running"
|
||||
self.controller.error = None
|
||||
elif scan_state != "failed":
|
||||
if report_status == "completed":
|
||||
scan_state = "completed"
|
||||
|
||||
85
strix/interface/url_safety.py
Normal file
85
strix/interface/url_safety.py
Normal file
@@ -0,0 +1,85 @@
|
||||
"""Validation for URLs printed or opened on behalf of a remote service."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ipaddress
|
||||
from urllib.parse import SplitResult, urlsplit
|
||||
|
||||
from strix.interface.terminal_text import has_terminal_control
|
||||
|
||||
|
||||
def is_safe_web_url(
|
||||
value: object,
|
||||
*,
|
||||
trusted_origin: str | None = None,
|
||||
require_trusted_origin: bool = False,
|
||||
) -> bool:
|
||||
"""Accept a strict HTTP(S) URL, optionally only on a pre-trusted origin."""
|
||||
parsed = _parse(value)
|
||||
if parsed is None:
|
||||
return False
|
||||
trusted = _parse(trusted_origin) if trusted_origin is not None else None
|
||||
same_origin = trusted is not None and _origin(parsed) == _origin(trusted)
|
||||
if require_trusted_origin:
|
||||
return same_origin
|
||||
if same_origin:
|
||||
return True
|
||||
return _is_safe_external_https(parsed)
|
||||
|
||||
|
||||
def _is_safe_external_https(parsed: SplitResult) -> bool:
|
||||
"""Reject local, numeric-looking, or otherwise ambiguous external hosts."""
|
||||
hostname = (parsed.hostname or "").lower().rstrip(".")
|
||||
if (
|
||||
parsed.scheme != "https"
|
||||
or hostname == "localhost"
|
||||
or hostname.endswith((".localhost", ".local"))
|
||||
):
|
||||
return False
|
||||
try:
|
||||
return ipaddress.ip_address(hostname).is_global
|
||||
except ValueError:
|
||||
pass
|
||||
labels = hostname.split(".")
|
||||
return len(labels) >= 2 and not all(_looks_numeric(label) for label in labels)
|
||||
|
||||
|
||||
def _parse(value: object) -> SplitResult | None:
|
||||
if not isinstance(value, str) or not value or has_terminal_control(value):
|
||||
return None
|
||||
if "\\" in value or any(character.isspace() for character in value):
|
||||
return None
|
||||
try:
|
||||
parsed = urlsplit(value)
|
||||
port = parsed.port
|
||||
except ValueError:
|
||||
return None
|
||||
hostname = parsed.hostname
|
||||
if (
|
||||
parsed.scheme not in {"http", "https"}
|
||||
or not hostname
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or parsed.fragment
|
||||
or "%" in parsed.netloc
|
||||
):
|
||||
return None
|
||||
try:
|
||||
hostname.encode("ascii")
|
||||
except UnicodeEncodeError:
|
||||
return None
|
||||
return parsed if port is None or 1 <= port <= 65535 else None
|
||||
|
||||
|
||||
def _origin(parsed: SplitResult) -> tuple[str, str, int]:
|
||||
default_port = 443 if parsed.scheme == "https" else 80
|
||||
return parsed.scheme, (parsed.hostname or "").lower().rstrip("."), parsed.port or default_port
|
||||
|
||||
|
||||
def _looks_numeric(label: str) -> bool:
|
||||
lowered = label.lower()
|
||||
if lowered.startswith("0x"):
|
||||
return len(lowered) > 2 and all(
|
||||
character in "0123456789abcdef" for character in lowered[2:]
|
||||
)
|
||||
return bool(lowered) and all(character.isdigit() for character in lowered)
|
||||
@@ -30,6 +30,7 @@ import {
|
||||
fetchTranscript,
|
||||
fetchVulnerabilities,
|
||||
forgetAuth,
|
||||
parseMcpConnectionStatus,
|
||||
type AuthStatus,
|
||||
type LoadedRun,
|
||||
type RunsPayload,
|
||||
@@ -169,6 +170,27 @@ export default function App() {
|
||||
const agentCount = run?.transcript.agents.length ?? 0;
|
||||
const verified = auth?.verified === true;
|
||||
|
||||
// The run's persisted MCP roster (from run.json via /api/run), plus the set of
|
||||
// connections with a tool call currently in flight. "In use" is derived here
|
||||
// from the connection-tagged tool events rather than carried on the roster:
|
||||
// an MCP dispatch event carries its connection name and a status that moves
|
||||
// running -> completed, so a connection is in use while one of its events is
|
||||
// still running. This mirrors the terminal UI's MCP panel exactly.
|
||||
const mcpConnections = useMemo(
|
||||
() => (run ? parseMcpConnectionStatus(run.raw) : []),
|
||||
[run]
|
||||
);
|
||||
const mcpInUse = useMemo(() => {
|
||||
const inUse = new Set<string>();
|
||||
for (const event of run?.transcript.events ?? []) {
|
||||
if (event.type !== "tool") continue;
|
||||
const connection = event.data?.mcp_connection;
|
||||
if (typeof connection !== "string" || !connection) continue;
|
||||
if (event.data?.status === "running") inUse.add(connection);
|
||||
}
|
||||
return inUse;
|
||||
}, [run]);
|
||||
|
||||
// Per-run guard for the default view: land on Agents while a scan is live,
|
||||
// Overview once it finishes. Applied at most once per run and never once the
|
||||
// user has navigated manually (userSetView flips the guard).
|
||||
@@ -251,6 +273,8 @@ export default function App() {
|
||||
}}
|
||||
issuesCount={run?.vulnerabilities.length ?? 0}
|
||||
agentCount={agentCount}
|
||||
mcpConnections={mcpConnections}
|
||||
mcpInUse={mcpInUse}
|
||||
runCount={runs?.count ?? 0}
|
||||
finished={run?.finished ?? false}
|
||||
verified={verified}
|
||||
|
||||
@@ -14,6 +14,7 @@ import { IoChatbubblesOutline } from "react-icons/io5";
|
||||
import { cn } from "@/lib/utils";
|
||||
import { ctaUrl, trackCta } from "@/lib/cta";
|
||||
import { UpgradeModal } from "@/components/UpgradeModal";
|
||||
import type { McpConnectionStatus } from "@/data/serverSource";
|
||||
import type { View } from "@/App";
|
||||
|
||||
/**
|
||||
@@ -37,6 +38,8 @@ interface SidebarProps {
|
||||
onSelectView: (view: View) => void;
|
||||
issuesCount: number;
|
||||
agentCount: number;
|
||||
mcpConnections: McpConnectionStatus[];
|
||||
mcpInUse: Set<string>;
|
||||
runCount: number;
|
||||
finished: boolean;
|
||||
verified: boolean;
|
||||
@@ -61,6 +64,8 @@ export default function Sidebar({
|
||||
onSelectView,
|
||||
issuesCount,
|
||||
agentCount,
|
||||
mcpConnections,
|
||||
mcpInUse,
|
||||
runCount,
|
||||
finished,
|
||||
verified,
|
||||
@@ -246,6 +251,9 @@ export default function Sidebar({
|
||||
onClick={() => onSelectView("agents")}
|
||||
/>
|
||||
)}
|
||||
{mcpConnections.length > 0 && (
|
||||
<McpConnectionsPanel connections={mcpConnections} inUse={mcpInUse} />
|
||||
)}
|
||||
<NavItem
|
||||
icon={<History className="h-4 w-4" />}
|
||||
label="Past runs"
|
||||
@@ -421,6 +429,84 @@ function NavItem({ icon, label, active, onClick, count }: NavItemProps) {
|
||||
);
|
||||
}
|
||||
|
||||
// The quarter-circle sweep frames the terminal UI cycles for an in-use
|
||||
// connection, and the sub-second tick that advances them.
|
||||
const SWEEP_FRAMES = ["◐", "◓", "◑", "◒"] as const;
|
||||
const SWEEP_MS = 220;
|
||||
|
||||
/**
|
||||
* The MCP connections panel: a compact roster of the run's connected MCP
|
||||
* servers, matching the terminal UI's sidebar panel. A header carries the
|
||||
* total count; each row shows a status glyph, the connection name, and its
|
||||
* tool count (or "offline"):
|
||||
* - solid green dot: attached and idle;
|
||||
* - green cycling quarter-circle (◐◓◑◒): a tool call is running against it;
|
||||
* - red dot + "offline": the connection's live session has died.
|
||||
*
|
||||
* "In use" is derived by the caller from the connection-tagged tool events, not
|
||||
* carried on the roster, so a call in flight shows motion with no extra signal.
|
||||
* The roster scrolls within a bounded height so a long list never blows out the
|
||||
* rail, mirroring how the nav above it scrolls.
|
||||
*/
|
||||
function McpConnectionsPanel({
|
||||
connections,
|
||||
inUse,
|
||||
}: {
|
||||
connections: McpConnectionStatus[];
|
||||
inUse: Set<string>;
|
||||
}) {
|
||||
const anyInUse = connections.some((c) => !c.dead && inUse.has(c.name));
|
||||
const [frame, setFrame] = useState(0);
|
||||
|
||||
// Advance the sweep only while at least one connection is in use, so an idle
|
||||
// panel does no work.
|
||||
useEffect(() => {
|
||||
if (!anyInUse) return;
|
||||
const id = setInterval(() => setFrame((f) => (f + 1) % SWEEP_FRAMES.length), SWEEP_MS);
|
||||
return () => clearInterval(id);
|
||||
}, [anyInUse]);
|
||||
|
||||
return (
|
||||
<div className="mt-1">
|
||||
<div className="flex h-7 items-center px-2 text-[11px] font-medium text-[#666]">
|
||||
MCP Connections ({connections.length})
|
||||
</div>
|
||||
<div className="max-h-48 overflow-y-auto overflow-x-clip scrollbar-thin">
|
||||
{connections.map((conn) => {
|
||||
const busy = !conn.dead && inUse.has(conn.name);
|
||||
return (
|
||||
<div
|
||||
key={conn.name}
|
||||
className="flex h-7 items-center gap-2 rounded-md px-2"
|
||||
title={conn.provider ? `${conn.name} · ${conn.provider}` : conn.name}
|
||||
>
|
||||
<span
|
||||
className={cn(
|
||||
"w-3 flex-none text-center text-[11px] leading-none",
|
||||
conn.dead ? "text-red-400" : "text-emerald-400"
|
||||
)}
|
||||
aria-hidden="true"
|
||||
>
|
||||
{conn.dead ? "●" : busy ? SWEEP_FRAMES[frame] : "●"}
|
||||
</span>
|
||||
<span className="min-w-0 flex-1 truncate text-[13px] font-medium text-[#ededed]">
|
||||
{conn.name}
|
||||
</span>
|
||||
{conn.dead ? (
|
||||
<span className="flex-none text-[11px] text-red-400">offline</span>
|
||||
) : (
|
||||
<span className="flex-none text-[11px] tabular-nums text-[#666]">
|
||||
{conn.toolCount} {conn.toolCount === 1 ? "tool" : "tools"}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
// Overview icon: a dashboard grid glyph (16x16 viewBox).
|
||||
function ProjectsIcon() {
|
||||
return (
|
||||
|
||||
@@ -14,6 +14,10 @@ import type { ToolRendererProps } from "@/types/events";
|
||||
* never as markdown, since it came from a server outside Strix.
|
||||
*
|
||||
* The full result is still in the run's event data on disk either way.
|
||||
*
|
||||
* list_mcps is the other exception: its result is the engine's own inventory of
|
||||
* the run's connections (names and tool counts), short and assembled by Strix
|
||||
* rather than returned by an outside server, so it is shown inline.
|
||||
*/
|
||||
|
||||
/** Arguments one line each, as the terminal prints them. */
|
||||
@@ -25,6 +29,52 @@ function argLines(args: unknown): string[] {
|
||||
});
|
||||
}
|
||||
|
||||
/** One connection out of a list_mcps inventory. */
|
||||
interface McpListingEntry {
|
||||
name: string;
|
||||
toolCount: number | null;
|
||||
dead: boolean;
|
||||
}
|
||||
|
||||
/**
|
||||
* The connections out of a list_mcps result, which is
|
||||
* `{"connections": [{id, name, description, tool_count}, ...]}`, sometimes
|
||||
* arriving JSON-encoded as a string. Anything else yields an empty list and the
|
||||
* row shows just the header and status. Unlike other MCP results this one is
|
||||
* safe to show: the engine assembled it from the run's own registered
|
||||
* connections, so it is short and never an outside server's payload. It still
|
||||
* renders as inert text.
|
||||
*/
|
||||
function listingEntries(result: unknown): McpListingEntry[] {
|
||||
let value = result;
|
||||
if (typeof value === "string") {
|
||||
try {
|
||||
value = JSON.parse(value);
|
||||
} catch {
|
||||
return [];
|
||||
}
|
||||
}
|
||||
const connections =
|
||||
value && typeof value === "object" && !Array.isArray(value)
|
||||
? (value as Record<string, unknown>).connections
|
||||
: null;
|
||||
if (!Array.isArray(connections)) return [];
|
||||
return connections.flatMap((entry) => {
|
||||
if (!entry || typeof entry !== "object" || Array.isArray(entry)) return [];
|
||||
const record = entry as Record<string, unknown>;
|
||||
const name =
|
||||
typeof record.name === "string" && record.name.trim()
|
||||
? record.name.trim()
|
||||
: typeof record.id === "string"
|
||||
? record.id.trim()
|
||||
: "";
|
||||
if (!name) return [];
|
||||
const toolCount = typeof record.tool_count === "number" ? record.tool_count : null;
|
||||
const dead = record.dead === true;
|
||||
return [{ name, toolCount, dead }];
|
||||
});
|
||||
}
|
||||
|
||||
const MAX_ERROR_CHARS = 600;
|
||||
|
||||
function errorText(result: unknown): string | null {
|
||||
@@ -50,11 +100,17 @@ export default function McpRenderer({
|
||||
// describe_mcp inspects a connection's catalog rather than calling a tool on
|
||||
// it, so the connection is the subject and there is no underlying tool.
|
||||
const inspecting = toolName === "describe_mcp";
|
||||
// list_mcps inventories every connection rather than touching one, so it
|
||||
// carries no connection at all and is routed here by name instead.
|
||||
const listing = toolName === "list_mcps";
|
||||
const entries = listing ? listingEntries(result) : [];
|
||||
|
||||
return (
|
||||
<div>
|
||||
<div className="flex items-center gap-2 flex-wrap">
|
||||
{inspecting ? (
|
||||
{listing ? (
|
||||
<span className="text-[13px] text-[#555]">Listing connected MCP servers</span>
|
||||
) : inspecting ? (
|
||||
<>
|
||||
<span className="text-[13px] text-[#555]">Inspecting MCP server</span>
|
||||
{mcpConnection && (
|
||||
@@ -82,6 +138,26 @@ export default function McpRenderer({
|
||||
</div>
|
||||
)}
|
||||
|
||||
{entries.length > 0 && (
|
||||
<div className="mt-1 font-mono text-[13px] leading-relaxed">
|
||||
{entries.map((entry) => (
|
||||
<div key={entry.name} className={`break-all${entry.dead ? " opacity-50" : ""}`}>
|
||||
<span className="text-teal-300">{entry.name}</span>
|
||||
{entry.dead ? (
|
||||
<span className="text-red-400/80"> · offline</span>
|
||||
) : (
|
||||
entry.toolCount !== null && (
|
||||
<span className="text-[#555]">
|
||||
{" "}
|
||||
· {entry.toolCount} {entry.toolCount === 1 ? "tool" : "tools"}
|
||||
</span>
|
||||
)
|
||||
)}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div className="mt-1 text-[13px]">
|
||||
{status === "running" && <span className="text-[#666]">Running</span>}
|
||||
{status === "completed" && <span className="text-emerald-400/80">✓ Done</span>}
|
||||
|
||||
@@ -16,13 +16,6 @@ const ACTION_LABELS: Record<string, { label: string; Icon: typeof Crosshair }> =
|
||||
amend_threat_model: { label: "Threat model amended", Icon: Plus },
|
||||
};
|
||||
|
||||
/** A git sha is noise past its first bytes, and "unversioned" is not a revision. */
|
||||
function shortRevision(revision: unknown): string {
|
||||
const value = typeof revision === "string" ? revision.trim() : "";
|
||||
if (!value || value === "unversioned") return "";
|
||||
return value.slice(0, 8);
|
||||
}
|
||||
|
||||
export default function ThreatModelRenderer({ toolName, args, result }: ToolRendererProps) {
|
||||
const action = ACTION_LABELS[toolName] ?? { label: "Threat model", Icon: Crosshair };
|
||||
const ActionIcon = action.Icon;
|
||||
@@ -59,22 +52,15 @@ export default function ThreatModelRenderer({ toolName, args, result }: ToolRend
|
||||
return (
|
||||
<div>
|
||||
{header}
|
||||
<div className="mt-1.5 text-[#555] text-xs">No model cached for this target yet</div>
|
||||
<div className="mt-1.5 text-[#555] text-xs">No model derived for this target yet</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
const rawAmendments = structured?.amendments;
|
||||
const amendments: Amendment[] = Array.isArray(rawAmendments) ? (rawAmendments as Amendment[]) : [];
|
||||
const cachedRevision = shortRevision(structured?.cached_revision);
|
||||
return (
|
||||
<div>
|
||||
{header}
|
||||
{structured?.stale === true && (
|
||||
<div className="mt-1.5 flex items-center gap-1.5 text-yellow-400/80 text-xs">
|
||||
<AlertTriangle className="w-3 h-3 shrink-0" />
|
||||
<span>stale{cachedRevision ? ` — written at ${cachedRevision}` : ""}</span>
|
||||
</div>
|
||||
)}
|
||||
{amendments.length > 0 && (
|
||||
<div className="mt-2">
|
||||
<span className="text-amber-400/70 text-xs font-semibold">
|
||||
@@ -119,12 +105,10 @@ export default function ThreatModelRenderer({ toolName, args, result }: ToolRend
|
||||
}
|
||||
|
||||
const cleared = (structured?.amendments_cleared as number | undefined) ?? 0;
|
||||
const revision = shortRevision(structured?.revision);
|
||||
const content = (args.content as string) ?? "";
|
||||
return (
|
||||
<div>
|
||||
{header}
|
||||
{revision && <div className="mt-1.5 text-[#666] font-mono text-xs">at {revision}</div>}
|
||||
{/* Saving folds amendments away — the one destructive thing this tool does. */}
|
||||
{cleared > 0 && (
|
||||
<div className="mt-1.5 flex items-center gap-1.5 text-yellow-400/80 text-xs">
|
||||
|
||||
@@ -93,7 +93,9 @@ const CATEGORY_META: Record<ToolCategory, CategoryMeta> = {
|
||||
threatModel: { renderer: ThreatModelRenderer, icon: Crosshair, color: "text-blue-400", match: /threat_model/ },
|
||||
telemetry: { renderer: FallbackRenderer, icon: Wrench, color: "text-[#555]" },
|
||||
// Tools from the user's own MCP servers. Resolved from the connection on the
|
||||
// event rather than from a tool name, so this family has no names below.
|
||||
// event rather than from a tool name — except list_mcps, the engine's
|
||||
// inventory of every connection, which touches none and so carries no
|
||||
// connection to resolve from; it is the family's one name below.
|
||||
mcp: { renderer: McpRenderer, icon: Plug, color: "text-teal-400" },
|
||||
};
|
||||
|
||||
@@ -114,7 +116,7 @@ const CATEGORY_TOOLS: Record<ToolCategory, readonly string[]> = {
|
||||
filesystem: ["apply_patch", "view_image", "str_replace_editor", "list_files", "search_files"],
|
||||
// Caido proxy tools (legacy: send_request)
|
||||
proxy: ["list_requests", "view_request", "repeat_request", "list_sitemap", "view_sitemap_entry", "scope_rules", "send_request"],
|
||||
reporting: ["create_vulnerability_report", "list_reports", "get_report"],
|
||||
reporting: ["create_vulnerability_report", "update_vulnerability_report", "list_reports", "get_report"],
|
||||
thinking: ["think"],
|
||||
agents: ["create_agent", "agent_finish", "send_message_to_agent", "wait_for_agents", "view_agent_graph", "stop_agent"],
|
||||
search: ["web_search"],
|
||||
@@ -128,7 +130,7 @@ const CATEGORY_TOOLS: Record<ToolCategory, readonly string[]> = {
|
||||
// Per-target threat model, shared across the agent tree
|
||||
threatModel: ["get_threat_model", "save_threat_model", "amend_threat_model"],
|
||||
telemetry: ["sandbox_error_details", "llm_error_details"],
|
||||
mcp: [],
|
||||
mcp: ["list_mcps"],
|
||||
};
|
||||
|
||||
/** Reverse index (tool name → family), built once from CATEGORY_TOOLS. */
|
||||
|
||||
@@ -39,6 +39,39 @@ export interface Transcript {
|
||||
events: TranscriptEvent[];
|
||||
}
|
||||
|
||||
/**
|
||||
* One MCP connection's non-secret status, as persisted to run.json by the
|
||||
* engine under `mcp_connection_status` and surfaced verbatim by GET /api/run.
|
||||
* Only name / provider / tool_count / dead ride here; never config, url, or
|
||||
* token. `dead` means the connection's live session gave up reconnecting.
|
||||
*/
|
||||
export interface McpConnectionStatus {
|
||||
name: string;
|
||||
provider: string | null;
|
||||
toolCount: number;
|
||||
dead: boolean;
|
||||
}
|
||||
|
||||
/**
|
||||
* Read the MCP connection roster out of a raw run record. Tolerates the field
|
||||
* being absent (older runs, or a run with no MCP) and any malformed entry,
|
||||
* yielding an empty list rather than throwing.
|
||||
*/
|
||||
export function parseMcpConnectionStatus(raw: Record<string, unknown>): McpConnectionStatus[] {
|
||||
const list = raw?.mcp_connection_status;
|
||||
if (!Array.isArray(list)) return [];
|
||||
return list.flatMap((entry) => {
|
||||
if (!entry || typeof entry !== "object" || Array.isArray(entry)) return [];
|
||||
const record = entry as Record<string, unknown>;
|
||||
const name = typeof record.name === "string" ? record.name.trim() : "";
|
||||
if (!name) return [];
|
||||
const provider = typeof record.provider === "string" && record.provider.trim() ? record.provider.trim() : null;
|
||||
const toolCount = typeof record.tool_count === "number" ? record.tool_count : 0;
|
||||
const dead = record.dead === true;
|
||||
return [{ name, provider, toolCount, dead }];
|
||||
});
|
||||
}
|
||||
|
||||
export interface LoadedRun {
|
||||
summary: ParsedRunSummary;
|
||||
/** Whole raw run record (for llm_usage, targets_info details, etc.). */
|
||||
|
||||
@@ -20,6 +20,7 @@ from datetime import datetime
|
||||
from io import BytesIO
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from markdown_it import MarkdownIt
|
||||
from pypdf import PdfReader, PdfWriter
|
||||
from reportlab.lib import colors
|
||||
from reportlab.lib.enums import TA_CENTER
|
||||
@@ -49,6 +50,8 @@ from strix.interface.viewer.transcript import (
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
from markdown_it.token import Token
|
||||
|
||||
|
||||
# Palette lifted from the cloud report theme (styles/base.ts, docx/theme.ts).
|
||||
_INK = colors.HexColor("#000000")
|
||||
@@ -72,11 +75,21 @@ _SANS_BOLD = "Helvetica-Bold"
|
||||
_MONO = "Courier"
|
||||
|
||||
_PAGE_W, _PAGE_H = A4
|
||||
_INLINE_MD = MarkdownIt("commonmark", {"html": False, "linkify": False}).disable(
|
||||
["autolink", "image", "link"]
|
||||
)
|
||||
_UNSAFE_TEXT_RE = re.compile(r"[\x00-\x08\x0b\x0c\x0e-\x1f\x7f-\x9f\ud800-\udfff\ufffe\uffff]")
|
||||
|
||||
|
||||
def _normalize_text(value: Any) -> str:
|
||||
"""Normalize characters that ReportLab cannot safely serialize."""
|
||||
text = str(value).replace("\r\n", "\n").replace("\r", "\n")
|
||||
return _UNSAFE_TEXT_RE.sub("\ufffd", text)
|
||||
|
||||
|
||||
def _esc(value: Any) -> str:
|
||||
"""Escape a value for reportlab's Paragraph markup."""
|
||||
return html.escape(str(value)).replace("\n", "<br/>")
|
||||
return html.escape(_normalize_text(value)).replace("\n", "<br/>")
|
||||
|
||||
|
||||
class _NumberedCanvas(pdfcanvas.Canvas): # type: ignore[misc] # reportlab base is untyped
|
||||
@@ -253,7 +266,10 @@ def _duration(start: Any, end: Any) -> str:
|
||||
end_dt = _parse_time(end)
|
||||
if not start_dt or not end_dt:
|
||||
return "n/a"
|
||||
seconds = int((end_dt - start_dt).total_seconds())
|
||||
try:
|
||||
seconds = int((end_dt - start_dt).total_seconds())
|
||||
except (OverflowError, TypeError):
|
||||
return "n/a"
|
||||
if seconds < 0:
|
||||
return "n/a"
|
||||
hours, remainder = divmod(seconds, 3600)
|
||||
@@ -265,10 +281,18 @@ def _duration(start: Any, end: Any) -> str:
|
||||
return f"{secs}s"
|
||||
|
||||
|
||||
def _severity_badge(styles: dict[str, ParagraphStyle], severity: str) -> Table:
|
||||
def _normalize_severity(value: Any) -> str:
|
||||
severity = str(value or "").lower().strip()
|
||||
if severity == "informational":
|
||||
return "info"
|
||||
return severity if severity in {*_SEVERITY_COLORS, "info"} else "low"
|
||||
|
||||
|
||||
def _severity_badge(styles: dict[str, ParagraphStyle], severity: Any) -> Table:
|
||||
"""A colored pill matching .severity-badge in the cloud report."""
|
||||
severity = _normalize_severity(severity)
|
||||
color = _SEVERITY_COLORS.get(severity, _MUTED)
|
||||
cell = Paragraph(severity.upper(), styles["badge"])
|
||||
cell = Paragraph(_esc(severity.upper()), styles["badge"])
|
||||
table = Table([[cell]], colWidths=[len(severity) * 6.5 + 20])
|
||||
table.setStyle(
|
||||
TableStyle(
|
||||
@@ -406,27 +430,36 @@ def _cover(
|
||||
|
||||
|
||||
def _inline_md(text: str) -> str:
|
||||
"""Convert inline markdown (bold, italic, `code`) to reportlab markup.
|
||||
"""Render a safe subset of inline Markdown as ReportLab markup."""
|
||||
tokens = _INLINE_MD.parseInline(_normalize_text(text))[0].children or []
|
||||
return "".join(_inline_token_markup(token) for token in tokens)
|
||||
|
||||
Code spans are stashed as placeholders before bold/italic run, so bold that
|
||||
wraps a code span (``**`x`**``) works and code contents are never mangled.
|
||||
"""
|
||||
codes: list[str] = []
|
||||
|
||||
def _stash(match: re.Match[str]) -> str:
|
||||
codes.append(match.group(1))
|
||||
return f"\x00{len(codes) - 1}\x00"
|
||||
def _inline_token_markup(token: Token) -> str:
|
||||
fixed_markup = {
|
||||
"strong_open": "<b>",
|
||||
"strong_close": "</b>",
|
||||
"em_open": "<i>",
|
||||
"em_close": "</i>",
|
||||
"hardbreak": "<br/>",
|
||||
"softbreak": " ",
|
||||
}.get(token.type)
|
||||
if fixed_markup is not None:
|
||||
return fixed_markup
|
||||
if token.type == "code_inline":
|
||||
return f'<font face="{_MONO}" color="#b31d28">{html.escape(token.content)}</font>'
|
||||
# Unsupported token content remains escaped so parser extensions cannot
|
||||
# expose ReportLab tags.
|
||||
return html.escape(token.content)
|
||||
|
||||
seg = html.escape(re.sub(r"`([^`]+)`", _stash, text))
|
||||
seg = re.sub(r"\*\*(.+?)\*\*", r"<b>\1</b>", seg)
|
||||
seg = re.sub(r"__(.+?)__", r"<b>\1</b>", seg)
|
||||
seg = re.sub(r"\*(.+?)\*", r"<i>\1</i>", seg)
|
||||
|
||||
def _restore(match: re.Match[str]) -> str:
|
||||
inner = html.escape(codes[int(match.group(1))])
|
||||
return f'<font face="{_MONO}" color="#b31d28">{inner}</font>'
|
||||
|
||||
return re.sub(r"\x00(\d+)\x00", _restore, seg)
|
||||
def _markdown_paragraph(text: str, style: ParagraphStyle) -> Paragraph:
|
||||
"""Build a Markdown paragraph, falling back to escaped source text."""
|
||||
source = _normalize_text(text)
|
||||
try:
|
||||
return Paragraph(_inline_md(source), style)
|
||||
except ValueError:
|
||||
return Paragraph(_esc(source), style)
|
||||
|
||||
|
||||
def _strip_leading_heading(md: str) -> str:
|
||||
@@ -447,12 +480,12 @@ def _markdown_flowables( # noqa: PLR0915 - cohesive block parser, splitting hur
|
||||
|
||||
def flush_para() -> None:
|
||||
if para:
|
||||
flow.append(Paragraph(_inline_md(" ".join(para)), styles["body"]))
|
||||
flow.append(_markdown_paragraph(" ".join(para), styles["body"]))
|
||||
para.clear()
|
||||
|
||||
def flush_bullets() -> None:
|
||||
for marker, item in bullets:
|
||||
flow.append(Paragraph(f"{marker} {_inline_md(item)}", styles["bullet"]))
|
||||
flow.append(_markdown_paragraph(f"{marker}\u00a0{item}", styles["bullet"]))
|
||||
bullets.clear()
|
||||
|
||||
lines = md.replace("\r\n", "\n").split("\n")
|
||||
@@ -479,7 +512,7 @@ def _markdown_flowables( # noqa: PLR0915 - cohesive block parser, splitting hur
|
||||
if heading:
|
||||
flush_para()
|
||||
flush_bullets()
|
||||
flow.append(Paragraph(_inline_md(heading.group(2)), styles["md_heading"]))
|
||||
flow.append(_markdown_paragraph(heading.group(2), styles["md_heading"]))
|
||||
i += 1
|
||||
continue
|
||||
ordered = re.match(r"^(\d+)\.\s+(.*)$", stripped)
|
||||
@@ -532,7 +565,7 @@ def _finding_flowables(
|
||||
styles: dict[str, ParagraphStyle], index: int, vuln: dict[str, Any]
|
||||
) -> list[Flowable]:
|
||||
title = vuln.get("title") or "Untitled finding"
|
||||
severity = str(vuln.get("severity") or "").lower().strip() or "low"
|
||||
severity = _normalize_severity(vuln.get("severity"))
|
||||
|
||||
meta_bits = []
|
||||
if vuln.get("cvss") is not None:
|
||||
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
10
strix/interface/viewer/static/assets/index-qwPOPAGC.css
Normal file
10
strix/interface/viewer/static/assets/index-qwPOPAGC.css
Normal file
File diff suppressed because one or more lines are too long
@@ -6,8 +6,8 @@
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<meta name="color-scheme" content="dark" />
|
||||
<title>Strix Results</title>
|
||||
<script type="module" crossorigin src="./assets/index-CYf9nnT3.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="./assets/index-D0453ODW.css">
|
||||
<script type="module" crossorigin src="./assets/index-Bpn8GiSb.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="./assets/index-qwPOPAGC.css">
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
|
||||
@@ -14,6 +14,7 @@ from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import logging
|
||||
import sys
|
||||
import threading
|
||||
|
||||
|
||||
@@ -30,12 +31,38 @@ _lock = threading.Lock()
|
||||
_thread: threading.Thread | None = None
|
||||
|
||||
|
||||
def _purge_orphaned_modules(before: frozenset[str]) -> None:
|
||||
"""Remove submodules stranded by an import attempt that just failed.
|
||||
|
||||
When a package import fails partway (for example CPython's import-lock
|
||||
deadlock avoidance breaking a cross-thread cycle), the failed package is
|
||||
removed from ``sys.modules`` but submodules it already finished stay
|
||||
behind. A later import of one of those submodules then short-circuits on
|
||||
the cached entry without re-importing its parent, and re-entering the
|
||||
parent from inside a submodule crashes with "partially initialized
|
||||
module". Dropping the orphans (cached submodules whose ancestor package is
|
||||
gone) restores a clean slate, and touches nothing another thread imported
|
||||
successfully.
|
||||
"""
|
||||
added = set(sys.modules) - before
|
||||
for name in added:
|
||||
parent = name.rpartition(".")[0]
|
||||
while parent:
|
||||
if parent not in sys.modules:
|
||||
sys.modules.pop(name, None)
|
||||
logger.debug("Import warm-up purged orphaned module %r", name)
|
||||
break
|
||||
parent = parent.rpartition(".")[0]
|
||||
|
||||
|
||||
def _warm(modules: tuple[str, ...]) -> None:
|
||||
for name in modules:
|
||||
before = frozenset(sys.modules)
|
||||
try:
|
||||
importlib.import_module(name)
|
||||
except Exception: # noqa: BLE001 - a failed warm-up must never fail the run.
|
||||
logger.debug("Import warm-up for %r failed", name, exc_info=True)
|
||||
_purge_orphaned_modules(before)
|
||||
|
||||
|
||||
def start_import_warmup(modules: tuple[str, ...] = WARMUP_MODULES) -> threading.Thread:
|
||||
|
||||
@@ -1,12 +1,26 @@
|
||||
"""Report/finding helpers."""
|
||||
|
||||
from strix.report.dedupe import check_duplicate
|
||||
from importlib import import_module
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from strix.report.state import ReportState, get_global_report_state, set_global_report_state
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from strix.report.dedupe import check_duplicate
|
||||
|
||||
__all__ = [
|
||||
"ReportState",
|
||||
"check_duplicate",
|
||||
"get_global_report_state",
|
||||
"set_global_report_state",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any:
|
||||
# check_duplicate pulls in the agents SDK import graph, so it resolves
|
||||
# lazily: importing this package must stay lightweight and never enter
|
||||
# that graph (the import warm-up thread may be walking it concurrently).
|
||||
if name == "check_duplicate":
|
||||
return import_module("strix.report.dedupe").check_duplicate
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -7,7 +7,6 @@ import logging
|
||||
import re
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agents.model_settings import ModelSettings
|
||||
from agents.models.interface import ModelTracing
|
||||
from openai.types.responses import ResponseOutputMessage
|
||||
|
||||
@@ -22,6 +21,8 @@ from strix.report.state import get_global_report_state
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agents.items import ModelResponse
|
||||
from agents.model_settings import ModelSettings
|
||||
from agents.models.interface import Model
|
||||
|
||||
from strix.config.settings import DedupeSettings
|
||||
|
||||
@@ -29,30 +30,11 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _dedupe_extra_args(dedupe: DedupeSettings) -> dict[str, str]:
|
||||
"""Per-call credential + endpoint for the dedupe model.
|
||||
|
||||
Provider env vars and the global base URL are process-wide, so a
|
||||
shared-provider dedupe key or a distinct dedupe endpoint can't be installed
|
||||
globally without clobbering (or being clobbered by) the main model's
|
||||
config. Passing them per call keeps the two apart. Only applies when a
|
||||
dedicated dedupe model is configured.
|
||||
"""
|
||||
if not dedupe.model:
|
||||
return {}
|
||||
extra: dict[str, str] = {}
|
||||
if dedupe.api_key and dedupe.api_key.strip():
|
||||
extra["api_key"] = dedupe.api_key.strip()
|
||||
if dedupe.api_base and dedupe.api_base.strip():
|
||||
extra["api_base"] = dedupe.api_base.strip()
|
||||
return extra
|
||||
|
||||
|
||||
def _dedupe_model_settings(
|
||||
dedupe: DedupeSettings, model_name: str, request_timeout: float | None
|
||||
) -> ModelSettings:
|
||||
llm = load_settings().llm
|
||||
settings = make_model_settings(
|
||||
return make_model_settings(
|
||||
dedupe.reasoning_effort,
|
||||
model_name=model_name,
|
||||
force_required_tool_choice=False,
|
||||
@@ -64,10 +46,21 @@ def _dedupe_model_settings(
|
||||
extra_headers=dedupe.extra_headers if dedupe.model else llm.extra_headers,
|
||||
has_tools=False,
|
||||
)
|
||||
extra = _dedupe_extra_args(dedupe)
|
||||
if extra:
|
||||
settings = settings.resolve(ModelSettings(extra_args=extra))
|
||||
return settings
|
||||
|
||||
|
||||
def resolve_dedupe_model(dedupe: DedupeSettings, model_name: str) -> Model:
|
||||
"""Resolve the dedupe model, bound to its own endpoint when it has one.
|
||||
|
||||
Credentials can't ride on the request: every model implementation already
|
||||
passes its own ``api_key``/``base_url``, so the same keys in ``extra_args``
|
||||
collide with them and raise before anything is sent. A provider bound to the
|
||||
dedupe endpoint keeps it apart from the main model's process-wide defaults.
|
||||
"""
|
||||
api_key = (dedupe.api_key or "").strip() if dedupe.model else ""
|
||||
api_base = (dedupe.api_base or "").strip() if dedupe.model else ""
|
||||
if not (api_key or api_base):
|
||||
return StrixProvider().get_model(model_name)
|
||||
return StrixProvider(api_key=api_key or None, base_url=api_base or None).get_model(model_name)
|
||||
|
||||
|
||||
DEDUPE_SYSTEM_PROMPT = """You are an expert vulnerability report deduplication judge.
|
||||
@@ -371,7 +364,7 @@ async def check_duplicate(
|
||||
|
||||
configure_sdk_model_defaults(settings)
|
||||
resolved_model = model_name.strip()
|
||||
model = StrixProvider().get_model(resolved_model)
|
||||
model = resolve_dedupe_model(dedupe, resolved_model)
|
||||
response = await model.get_response(
|
||||
system_instructions=DEDUPE_SYSTEM_PROMPT,
|
||||
input=user_msg,
|
||||
|
||||
@@ -1,23 +1,21 @@
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import subprocess
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from datetime import UTC, datetime
|
||||
from importlib.metadata import PackageNotFoundError, version
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Optional, cast
|
||||
from uuid import uuid4
|
||||
|
||||
from agents.usage import Usage
|
||||
|
||||
from strix.config import codex
|
||||
from strix.config.loader import load_settings
|
||||
from strix.core.paths import run_dir_for, runtime_state_dir
|
||||
from strix.report.coverage import write_coverage
|
||||
from strix.report.pricing import resolve_litellm_model
|
||||
from strix.report.sarif import write_sarif
|
||||
from strix.report.usage import LLMUsageLedger
|
||||
from strix.report.writer import (
|
||||
read_run_record,
|
||||
write_executive_report,
|
||||
@@ -27,10 +25,16 @@ from strix.report.writer import (
|
||||
from strix.telemetry import posthog, scarf
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agents.usage import Usage
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_global_report_state: Optional["ReportState"] = None
|
||||
|
||||
_CONTROL_CHARS = re.compile(r"[\x00-\x1f\x7f]+")
|
||||
|
||||
|
||||
def _strix_version() -> str | None:
|
||||
"""Best-effort package version for the SARIF tool.driver.version field."""
|
||||
@@ -40,6 +44,67 @@ def _strix_version() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
# Content a revision may replace. The identity of the finding (id, timestamp,
|
||||
# finding_class) and its original author stay put. dependency_metadata is
|
||||
# replaced whole, so a caller carries the package identity over itself.
|
||||
UPDATABLE_REPORT_FIELDS = frozenset(
|
||||
{
|
||||
"title",
|
||||
"dependency_metadata",
|
||||
"severity",
|
||||
"description",
|
||||
"impact",
|
||||
"target",
|
||||
"technical_analysis",
|
||||
"poc_description",
|
||||
"poc_script_code",
|
||||
"remediation_steps",
|
||||
"evidence",
|
||||
"assumptions",
|
||||
"counterevidence",
|
||||
"confidence",
|
||||
"confidence_rationale",
|
||||
"severity_change_conditions",
|
||||
"fix_effort",
|
||||
"cvss",
|
||||
"cvss_breakdown",
|
||||
"endpoint",
|
||||
"method",
|
||||
"cve",
|
||||
"cwe",
|
||||
"code_locations",
|
||||
"fix_verification",
|
||||
"fix_pr_body",
|
||||
}
|
||||
)
|
||||
|
||||
_LOWERCASE_REPORT_FIELDS = frozenset({"severity", "confidence", "fix_effort"})
|
||||
|
||||
# Fields that only describe another field. A revision may raise the rating or
|
||||
# replace the locations without restating the reasoning behind the old one, and
|
||||
# that leftover reasoning then contradicts the finding it annotates
|
||||
# ("confidence: high" beside a rationale calling the evidence unconfirmed). When
|
||||
# the field they describe changes and the update carries no replacement, they
|
||||
# are dropped rather than kept.
|
||||
_DEPENDENT_REPORT_FIELDS: dict[str, tuple[str, ...]] = {
|
||||
"confidence": ("confidence_rationale",),
|
||||
"severity": ("severity_change_conditions",),
|
||||
"cvss": ("cvss_breakdown",),
|
||||
"code_locations": ("fix_verification",),
|
||||
}
|
||||
|
||||
|
||||
def _clean_title(title: str) -> str:
|
||||
"""Return a single-line finding title.
|
||||
|
||||
A title quotes text from the scanned target, so it can carry newlines, tabs or
|
||||
other control characters. Those break every artifact that renders the title on
|
||||
one line, such as the markdown heading, the CSV cell and the TUI list. Control
|
||||
characters become spaces and runs of whitespace collapse to one space.
|
||||
"""
|
||||
return " ".join(_CONTROL_CHARS.sub(" ", title).split())
|
||||
|
||||
|
||||
def _number(value: Any) -> int | float:
|
||||
try:
|
||||
return float(value or 0)
|
||||
@@ -131,6 +196,10 @@ class ReportState:
|
||||
|
||||
self.scan_results: dict[str, Any] | None = None
|
||||
self.scan_config: dict[str, Any] | None = None
|
||||
# Imported here so importing this module never enters the agents SDK
|
||||
# package (which the warm-up thread may be initializing concurrently).
|
||||
from strix.report.usage import LLMUsageLedger
|
||||
|
||||
self._llm_usage = LLMUsageLedger()
|
||||
self._telemetry_llm_usage_baseline: dict[str, Any] = {}
|
||||
auth_mode = codex.auth_mode(load_settings().llm.model)
|
||||
@@ -150,6 +219,7 @@ class ReportState:
|
||||
|
||||
self.caido_url: str | None = None
|
||||
self.vulnerability_found_callback: Callable[[dict[str, Any]], None] | None = None
|
||||
self.vulnerability_updated_callback: Callable[[dict[str, Any]], None] | None = None
|
||||
|
||||
self._sarif_repo_ctx: dict[str, Any] | None = None
|
||||
self._sarif_repo_ctx_ready: bool = False
|
||||
@@ -217,8 +287,21 @@ class ReportState:
|
||||
)
|
||||
self.vulnerability_reports = [r for r in data if isinstance(r, dict)]
|
||||
for r in self.vulnerability_reports:
|
||||
# A finding written before the class was persisted still carries the
|
||||
# metadata of its class, so name the class it always had.
|
||||
if not r.get("finding_class"):
|
||||
r["finding_class"] = (
|
||||
"dependency_cve" if r.get("dependency_metadata") else "dynamic"
|
||||
)
|
||||
title = r.get("title")
|
||||
stale_md = False
|
||||
if isinstance(title, str):
|
||||
r["title"] = _clean_title(title)
|
||||
stale_md = r["title"] != title
|
||||
rid = r.get("id")
|
||||
if isinstance(rid, str):
|
||||
# A finding already on disk keeps its markdown, unless cleaning
|
||||
# changed the title: the heading on disk then needs a rewrite.
|
||||
if isinstance(rid, str) and not stale_md:
|
||||
self._saved_vuln_ids.add(rid)
|
||||
logger.info(
|
||||
"report state hydrated %d vulnerability report(s)",
|
||||
@@ -261,7 +344,7 @@ class ReportState:
|
||||
|
||||
report: dict[str, Any] = {
|
||||
"id": report_id,
|
||||
"title": title.strip(),
|
||||
"title": _clean_title(title),
|
||||
"severity": severity.lower().strip(),
|
||||
"timestamp": datetime.now(UTC).strftime("%Y-%m-%d %H:%M:%S UTC"),
|
||||
}
|
||||
@@ -331,6 +414,100 @@ class ReportState:
|
||||
self.save_run_data()
|
||||
return report_id
|
||||
|
||||
def update_vulnerability_report(
|
||||
self,
|
||||
report_id: str,
|
||||
fields: dict[str, Any],
|
||||
*,
|
||||
update_reason: str | None = None,
|
||||
updated_by_agent_id: str | None = None,
|
||||
updated_by_agent_name: str | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Apply a revision to an existing report, keeping its id.
|
||||
|
||||
A field that only describes a field this update replaces is dropped when
|
||||
the update carries no replacement for it, so the revised report cannot
|
||||
state a new rating beside the superseded reasoning for the old one.
|
||||
|
||||
Returns the revised report, or ``None`` when the id is unknown or when
|
||||
nothing in ``fields`` changes it.
|
||||
"""
|
||||
report = next((r for r in self.vulnerability_reports if r.get("id") == report_id), None)
|
||||
if report is None:
|
||||
logger.warning("cannot update unknown vulnerability report %s", report_id)
|
||||
return None
|
||||
|
||||
changed: dict[str, Any] = {}
|
||||
for key, raw_value in fields.items():
|
||||
if key not in UPDATABLE_REPORT_FIELDS or raw_value is None:
|
||||
continue
|
||||
value = raw_value
|
||||
if isinstance(value, str):
|
||||
value = _clean_title(value) if key == "title" else value.strip()
|
||||
if key in _LOWERCASE_REPORT_FIELDS:
|
||||
value = value.lower()
|
||||
if not value:
|
||||
continue
|
||||
if report.get(key) == value:
|
||||
continue
|
||||
changed[key] = value
|
||||
|
||||
superseded = {
|
||||
dependent
|
||||
for primary, dependents in _DEPENDENT_REPORT_FIELDS.items()
|
||||
if primary in changed
|
||||
for dependent in dependents
|
||||
if dependent not in changed and report.get(dependent) not in (None, "", [], {})
|
||||
}
|
||||
|
||||
if not changed and not superseded:
|
||||
logger.info("update for %s carried no new content; keeping it as is", report_id)
|
||||
return None
|
||||
|
||||
entry: dict[str, Any] = {
|
||||
"timestamp": datetime.now(UTC).strftime("%Y-%m-%d %H:%M:%S UTC"),
|
||||
"fields": sorted(changed),
|
||||
}
|
||||
if superseded:
|
||||
entry["dropped_fields"] = sorted(superseded)
|
||||
if update_reason and update_reason.strip():
|
||||
entry["reason"] = update_reason.strip()[:500]
|
||||
if updated_by_agent_id:
|
||||
entry["agent_id"] = updated_by_agent_id
|
||||
if updated_by_agent_name:
|
||||
entry["agent_name"] = updated_by_agent_name
|
||||
for key in ("severity", "cvss", "confidence"):
|
||||
if key in changed and report.get(key) is not None:
|
||||
entry[f"previous_{key}"] = report[key]
|
||||
|
||||
raw_history = report.get("update_history")
|
||||
history: list[dict[str, Any]] = (
|
||||
[e for e in raw_history if isinstance(e, dict)] if isinstance(raw_history, list) else []
|
||||
)
|
||||
history.append(entry)
|
||||
|
||||
report.update(changed)
|
||||
for dependent in superseded:
|
||||
report.pop(dependent, None)
|
||||
report["update_history"] = history
|
||||
report["updated_at"] = entry["timestamp"]
|
||||
|
||||
# The markdown on disk still shows the superseded evidence, so let the
|
||||
# writer re-render it.
|
||||
self._saved_vuln_ids.discard(report_id)
|
||||
|
||||
logger.info(
|
||||
"Updated vulnerability report %s (%s)",
|
||||
report_id,
|
||||
", ".join(entry["fields"]) or "no field replaced",
|
||||
)
|
||||
|
||||
if self.vulnerability_updated_callback:
|
||||
self.vulnerability_updated_callback(report)
|
||||
|
||||
self.save_run_data()
|
||||
return report
|
||||
|
||||
def get_existing_vulnerabilities(self) -> list[dict[str, Any]]:
|
||||
return list(self.vulnerability_reports)
|
||||
|
||||
@@ -338,7 +515,7 @@ class ReportState:
|
||||
self,
|
||||
*,
|
||||
agent_id: str,
|
||||
usage: Usage | None,
|
||||
usage: "Usage | None",
|
||||
agent_name: str | None = None,
|
||||
model: str | None = None,
|
||||
) -> None:
|
||||
@@ -416,6 +593,22 @@ class ReportState:
|
||||
self.run_record["mcp_connections"] = names
|
||||
self.save_run_data()
|
||||
|
||||
def record_mcp_connection_status(self, status: list[dict[str, Any]]) -> None:
|
||||
"""Persist the run's non-secret MCP connection status roster.
|
||||
|
||||
``status`` is one entry per connection carrying only ``name``,
|
||||
``provider``, ``tool_count``, and ``dead`` (no config, url, token, or
|
||||
auth). Saved as soon as the run connects and rewritten each time a
|
||||
connection dies, so the viewer, which rebuilds its display by re-reading
|
||||
the run's files from disk, can show a live connections panel and health
|
||||
without any in-memory event sink. Kept separate from the
|
||||
``mcp_connections`` name list so neither field repurposes the other.
|
||||
"""
|
||||
if self.run_record.get("mcp_connection_status") == status:
|
||||
return
|
||||
self.run_record["mcp_connection_status"] = status
|
||||
self.save_run_data()
|
||||
|
||||
def set_scan_config(self, config: dict[str, Any]) -> None:
|
||||
self.scan_config = config
|
||||
self.run_record["status"] = "running"
|
||||
|
||||
@@ -26,10 +26,33 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
_SEVERITY_ORDER = {"critical": 0, "high": 1, "medium": 2, "low": 3, "info": 4}
|
||||
|
||||
_CSV_FORMULA_PREFIXES = ("=", "+", "-", "@", "\t", "\r")
|
||||
|
||||
_FENCE_RE = re.compile(r"^```([^\n`]*)\r?\n(.*?)\r?\n?```$", re.DOTALL)
|
||||
_BACKTICK_RUN = re.compile(r"`+")
|
||||
|
||||
|
||||
def csv_safe(value: object) -> str:
|
||||
"""Return ``value`` as a CSV cell a spreadsheet will not treat as a formula.
|
||||
|
||||
Excel, LibreOffice and Sheets evaluate a cell whose first character is one of
|
||||
``= + - @``, tab or carriage return. The :mod:`csv` module quotes CSV syntax
|
||||
but has no notion of formula triggers, so such a value reaches the cell intact
|
||||
and is executed on open (CWE-1236). Vulnerability titles quote text from the
|
||||
scanned target, which is exactly the attacker-influenced input this guards
|
||||
against.
|
||||
|
||||
Prefixing with an apostrophe is the standard mitigation (OWASP): the rest of
|
||||
the cell is kept as literal text instead of being evaluated. Excel shows the
|
||||
apostrophe when it opens a ``.csv`` directly, which is cosmetic — the point is
|
||||
that nothing runs.
|
||||
"""
|
||||
text = str(value)
|
||||
if text.startswith(_CSV_FORMULA_PREFIXES):
|
||||
return "'" + text
|
||||
return text
|
||||
|
||||
|
||||
def safe_fence(content: str) -> str:
|
||||
"""Return a backtick fence that ``content`` cannot break out of.
|
||||
|
||||
@@ -151,11 +174,11 @@ def write_vulnerabilities(
|
||||
for report in sorted_reports:
|
||||
csv_writer.writerow(
|
||||
{
|
||||
"id": report["id"],
|
||||
"title": report["title"],
|
||||
"severity": report["severity"].upper(),
|
||||
"timestamp": report["timestamp"],
|
||||
"file": f"vulnerabilities/{report['id']}.md",
|
||||
"id": csv_safe(report["id"]),
|
||||
"title": csv_safe(report["title"]),
|
||||
"severity": csv_safe(report["severity"].upper()),
|
||||
"timestamp": csv_safe(report["timestamp"]),
|
||||
"file": csv_safe(f"vulnerabilities/{report['id']}.md"),
|
||||
},
|
||||
)
|
||||
atomic_write_text(csv_path, csv_buf.getvalue())
|
||||
@@ -176,11 +199,17 @@ def write_vulnerabilities(
|
||||
|
||||
|
||||
def atomic_write_text(path: Path, payload: str) -> None:
|
||||
"""Write *payload* to *path* via a sibling temp file and an atomic rename."""
|
||||
"""Write *payload* to *path* via a sibling temp file and an atomic rename.
|
||||
|
||||
``newline=""`` disables newline translation so *payload* lands byte-for-byte:
|
||||
the CSV index carries its own ``\\r\\n`` terminators, which text mode would turn
|
||||
into ``\\r\\r\\n`` on Windows.
|
||||
"""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="w",
|
||||
encoding="utf-8",
|
||||
newline="",
|
||||
dir=str(path.parent),
|
||||
prefix=f".{path.name}.",
|
||||
suffix=".tmp",
|
||||
@@ -327,4 +356,41 @@ def render_vulnerability_md(report: dict[str, Any]) -> str: # noqa: PLR0912, PL
|
||||
lines.append(str(report["assumptions"]))
|
||||
lines.append("")
|
||||
|
||||
lines.extend(render_update_history(report.get("update_history")))
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def render_update_history(history: Any) -> list[str]:
|
||||
"""Render the audit trail of every revision a report has received."""
|
||||
if not isinstance(history, list):
|
||||
return []
|
||||
entries: list[dict[str, Any]] = [
|
||||
cast("dict[str, Any]", e) for e in history if isinstance(e, dict)
|
||||
]
|
||||
if not entries:
|
||||
return []
|
||||
|
||||
lines = ["## Update History\n"]
|
||||
for entry in entries:
|
||||
author = str(entry.get("agent_name") or entry.get("agent_id") or "an agent")
|
||||
raw_fields = entry.get("fields")
|
||||
fields: list[Any] = raw_fields if isinstance(raw_fields, list) else []
|
||||
changed = ", ".join(str(field) for field in fields)
|
||||
timestamp = str(entry.get("timestamp") or "unknown")
|
||||
lines.append(f"**{timestamp}** — {author} updated: {changed}")
|
||||
raw_dropped = entry.get("dropped_fields")
|
||||
if isinstance(raw_dropped, list) and raw_dropped:
|
||||
dropped = ", ".join(str(field) for field in raw_dropped)
|
||||
lines.append(f" Dropped as superseded: {dropped}")
|
||||
for key, label in (
|
||||
("previous_severity", "severity"),
|
||||
("previous_cvss", "CVSS"),
|
||||
("previous_confidence", "confidence"),
|
||||
):
|
||||
if entry.get(key) is not None:
|
||||
lines.append(f" Previous {label}: {entry[key]}")
|
||||
if entry.get("reason"):
|
||||
lines.append(f" Reason: {entry['reason']}")
|
||||
lines.append("")
|
||||
return lines
|
||||
|
||||
@@ -5,7 +5,9 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
@@ -13,7 +15,6 @@ from agents.sandbox.entries import BaseEntry, File, LocalDir
|
||||
from agents.sandbox.manifest import Environment, Manifest
|
||||
|
||||
from strix.config import load_settings
|
||||
from strix.core.paths import run_dir_for, runtime_state_dir
|
||||
from strix.runtime.backends import backend_supports_bind_mounts, get_backend
|
||||
from strix.runtime.caido_bootstrap import bootstrap_caido
|
||||
from strix.runtime.caido_handle import CaidoBootstrapHandle
|
||||
@@ -168,6 +169,17 @@ def build_extra_file_entries(
|
||||
return entries
|
||||
|
||||
|
||||
def extra_file_staging_dir(scan_id: str) -> Path:
|
||||
"""A fresh host staging directory for a scan's extra-file bind mounts.
|
||||
|
||||
The docker daemon resolves bind sources in its own filesystem. With a
|
||||
remote daemon (e.g. a dind sidecar) the run directory is not shared, so
|
||||
staging lives under the temp dir like every other bind-mount source.
|
||||
"""
|
||||
safe = "".join(c if c.isalnum() or c in "-_." else "-" for c in scan_id)
|
||||
return Path(tempfile.mkdtemp(prefix=f"strix-extra-files-{safe}-"))
|
||||
|
||||
|
||||
def build_extra_file_bind_mounts(
|
||||
extra_files: list[dict[str, Any]],
|
||||
staging_dir: Path,
|
||||
@@ -280,11 +292,12 @@ async def create_or_reuse(
|
||||
backend_name = load_settings().runtime.backend
|
||||
backend = get_backend(backend_name)
|
||||
|
||||
staging_dir: Path | None = None
|
||||
if backend_supports_bind_mounts(backend_name):
|
||||
bind_mounts = build_bind_mounts(local_sources)
|
||||
entries: dict[str | Path, BaseEntry] = {}
|
||||
if extra_files:
|
||||
staging_dir = runtime_state_dir(run_dir_for(scan_id)) / "extra_files"
|
||||
staging_dir = extra_file_staging_dir(scan_id)
|
||||
bind_mounts.extend(
|
||||
build_extra_file_bind_mounts(extra_files, staging_dir, local_sources)
|
||||
)
|
||||
@@ -322,44 +335,56 @@ async def create_or_reuse(
|
||||
image,
|
||||
)
|
||||
report("Starting sandbox container")
|
||||
client, session = await backend(
|
||||
image=image,
|
||||
manifest=manifest,
|
||||
exposed_ports=(_CONTAINER_CAIDO_PORT,),
|
||||
bind_mounts=bind_mounts,
|
||||
)
|
||||
|
||||
report("Setting up the proxy")
|
||||
caido_endpoint = await session.resolve_exposed_port(_CONTAINER_CAIDO_PORT)
|
||||
scheme = "https" if caido_endpoint.tls else "http"
|
||||
host_caido_url = f"{scheme}://{caido_endpoint.host}:{caido_endpoint.port}"
|
||||
logger.debug("Caido host endpoint resolved: %s", host_caido_url)
|
||||
|
||||
# The Caido login + project setup polls the guest for a couple of seconds
|
||||
# and nothing needs the client before the first proxy tool call, so it
|
||||
# runs concurrently with the rest of scan start; consumers resolve the
|
||||
# handle at first use (see CaidoBootstrapHandle).
|
||||
caido_client = CaidoBootstrapHandle(
|
||||
asyncio.create_task(
|
||||
bootstrap_caido(
|
||||
session,
|
||||
host_url=host_caido_url,
|
||||
container_url=container_caido_url,
|
||||
),
|
||||
name=f"caido-bootstrap-{scan_id}",
|
||||
try:
|
||||
client, session = await backend(
|
||||
image=image,
|
||||
manifest=manifest,
|
||||
exposed_ports=(_CONTAINER_CAIDO_PORT,),
|
||||
bind_mounts=bind_mounts,
|
||||
)
|
||||
)
|
||||
|
||||
bundle = {
|
||||
"client": client,
|
||||
"session": session,
|
||||
"caido_client": caido_client,
|
||||
}
|
||||
_SESSION_CACHE[scan_id] = bundle
|
||||
report("Setting up the proxy")
|
||||
caido_endpoint = await session.resolve_exposed_port(_CONTAINER_CAIDO_PORT)
|
||||
scheme = "https" if caido_endpoint.tls else "http"
|
||||
host_caido_url = f"{scheme}://{caido_endpoint.host}:{caido_endpoint.port}"
|
||||
logger.debug("Caido host endpoint resolved: %s", host_caido_url)
|
||||
|
||||
# The Caido login + project setup polls the guest for a couple of seconds
|
||||
# and nothing needs the client before the first proxy tool call, so it
|
||||
# runs concurrently with the rest of scan start; consumers resolve the
|
||||
# handle at first use (see CaidoBootstrapHandle).
|
||||
caido_client = CaidoBootstrapHandle(
|
||||
asyncio.create_task(
|
||||
bootstrap_caido(
|
||||
session,
|
||||
host_url=host_caido_url,
|
||||
container_url=container_caido_url,
|
||||
),
|
||||
name=f"caido-bootstrap-{scan_id}",
|
||||
)
|
||||
)
|
||||
|
||||
bundle = {
|
||||
"client": client,
|
||||
"session": session,
|
||||
"caido_client": caido_client,
|
||||
"extra_file_staging_dir": staging_dir,
|
||||
}
|
||||
_SESSION_CACHE[scan_id] = bundle
|
||||
except BaseException:
|
||||
# Until the bundle is cached, cleanup(scan_id) cannot find the
|
||||
# staging dir, so it is removed here.
|
||||
_remove_staging_dir(staging_dir)
|
||||
raise
|
||||
logger.info("Sandbox session for scan %s ready and cached", scan_id)
|
||||
return bundle
|
||||
|
||||
|
||||
def _remove_staging_dir(staging_dir: Path | None) -> None:
|
||||
if staging_dir is not None:
|
||||
shutil.rmtree(staging_dir, ignore_errors=True)
|
||||
|
||||
|
||||
async def cleanup(scan_id: str) -> None:
|
||||
"""Tear down ``scan_id``'s container and drop its cache entry.
|
||||
|
||||
@@ -373,6 +398,8 @@ async def cleanup(scan_id: str) -> None:
|
||||
logger.debug("cleanup(%s): no cached session", scan_id)
|
||||
return
|
||||
|
||||
_remove_staging_dir(bundle.get("extra_file_staging_dir"))
|
||||
|
||||
caido_client = bundle.get("caido_client")
|
||||
if caido_client is not None:
|
||||
try:
|
||||
|
||||
@@ -27,7 +27,7 @@ Before spawning agents, analyze the target from the scan config/scope and any pr
|
||||
|
||||
## Establish the Threat Model
|
||||
|
||||
Every scan needs one shared answer to "who is the attacker here, and what are they attacking" — black-box or white-box. Without it, five agents derive five different answers and their findings cannot be reconciled. Call `get_threat_model` on the target (a host, a URL, or a repository path) before you spawn hunters; if nothing is cached, derive one and persist it with `save_threat_model`. It is cached per target, so a later scan of the same host or tree reads it back instead of paying for it twice, and a model written from source is read back by an agent testing the deployment.
|
||||
Every scan needs one shared answer to "who is the attacker here, and what are they attacking" — black-box or white-box. Without it, five agents derive five different answers and their findings cannot be reconciled. Call `get_threat_model` on the target (a host, a URL, or a repository path) before you spawn hunters; if no model exists yet, derive one and share it with `save_threat_model`. It lives for this scan only — nothing carries over from an earlier run, so every scan derives its own — but within the run every agent reads the same document, and a model written from source is read back by an agent testing the deployment.
|
||||
|
||||
**When the target includes a repository**, derive it up front: the code tells you the boundaries, entrypoints, and controls before you send a single request.
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ from strix.tools.mcp.config import (
|
||||
McpAuth,
|
||||
McpConnectionConfig,
|
||||
)
|
||||
from strix.tools.mcp.failures import FailureInfo, HttpStatusRecorder, classify
|
||||
from strix.tools.mcp.loader import load_user_mcp_configs
|
||||
from strix.tools.mcp.naming import namespaced_tool_name
|
||||
from strix.tools.mcp.registry import (
|
||||
@@ -23,10 +24,12 @@ from strix.tools.mcp.registry import (
|
||||
McpCallInfo,
|
||||
McpConnectionEntry,
|
||||
McpConnectionRequest,
|
||||
McpConnectionStatus,
|
||||
McpConnectionSummary,
|
||||
McpRegistry,
|
||||
resolve_mcp_call,
|
||||
)
|
||||
from strix.tools.mcp.session import McpConnectionUnavailableError, SupervisedMcpSession
|
||||
|
||||
|
||||
__all__ = [
|
||||
@@ -36,15 +39,21 @@ __all__ = [
|
||||
"MCP_REGISTRY_CONTEXT_KEY",
|
||||
"BearerAuth",
|
||||
"ConnectedMcpServer",
|
||||
"FailureInfo",
|
||||
"HttpStatusRecorder",
|
||||
"McpAuth",
|
||||
"McpCallInfo",
|
||||
"McpConnectionConfig",
|
||||
"McpConnectionEntry",
|
||||
"McpConnectionRequest",
|
||||
"McpConnectionStatus",
|
||||
"McpConnectionSummary",
|
||||
"McpConnectionUnavailableError",
|
||||
"McpRegistry",
|
||||
"SupervisedMcpSession",
|
||||
"attach_mcp_requests",
|
||||
"call_mcp",
|
||||
"classify",
|
||||
"connect_mcp_servers",
|
||||
"describe_mcp",
|
||||
"list_mcps",
|
||||
|
||||
@@ -26,9 +26,10 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agents import RunContextWrapper, function_tool
|
||||
|
||||
from strix.tools.mcp.client import dispatch_mcp_call
|
||||
from strix.tools.mcp.client import _errored_tool_output
|
||||
from strix.tools.mcp.naming import namespaced_tool_name
|
||||
from strix.tools.mcp.registry import MCP_REGISTRY_CONTEXT_KEY, McpRegistry
|
||||
from strix.tools.mcp.session import McpConnectionUnavailableError
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -49,6 +50,13 @@ def _unknown_connection(connection: str, registry: McpRegistry) -> str:
|
||||
return f"Unknown MCP connection {connection!r}. Available connections: {available}."
|
||||
|
||||
|
||||
def _unavailable_connection(connection: str) -> str:
|
||||
return (
|
||||
f"MCP connection {connection!r} is unavailable: its live session failed and "
|
||||
"could not be reconnected, so it is unavailable for the rest of this run."
|
||||
)
|
||||
|
||||
|
||||
def _format_tool(tool: MCPTool) -> str:
|
||||
schema = json.dumps(tool.inputSchema or {"type": "object"}, indent=2, ensure_ascii=False)
|
||||
description = (tool.description or "").strip() or "(no description)"
|
||||
@@ -70,6 +78,7 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]:
|
||||
registry = _registry_from_ctx(ctx)
|
||||
if registry is None or not registry:
|
||||
return {"connections": []}
|
||||
dead_by_name = {status.name: status.dead for status in registry.statuses()}
|
||||
return {
|
||||
"connections": [
|
||||
{
|
||||
@@ -77,6 +86,7 @@ async def list_mcps(ctx: RunContextWrapper) -> dict[str, Any]:
|
||||
"name": summary.name,
|
||||
"description": summary.purpose,
|
||||
"tool_count": summary.tool_count,
|
||||
"dead": dead_by_name.get(summary.name, False),
|
||||
}
|
||||
for summary in registry.summaries()
|
||||
]
|
||||
@@ -102,7 +112,10 @@ async def describe_mcp(ctx: RunContextWrapper, connection: str) -> str:
|
||||
entry = registry.get(connection)
|
||||
if entry is None:
|
||||
return _unknown_connection(connection, registry)
|
||||
tools = await entry.server.list_tools()
|
||||
try:
|
||||
tools = await entry.session.list_tools()
|
||||
except McpConnectionUnavailableError:
|
||||
return _unavailable_connection(connection)
|
||||
if not tools:
|
||||
return f"MCP connection {connection!r} offers no tools."
|
||||
header = f"MCP connection {connection!r} offers {len(tools)} tool(s):"
|
||||
@@ -155,7 +168,10 @@ async def call_mcp(
|
||||
return invalid_arguments
|
||||
if arguments is not None and not isinstance(arguments, dict):
|
||||
return invalid_arguments
|
||||
available = await entry.server.list_tools()
|
||||
try:
|
||||
available = await entry.session.list_tools()
|
||||
except McpConnectionUnavailableError:
|
||||
return _errored_tool_output(_unavailable_connection(connection))
|
||||
valid_names = {mcp_tool.name for mcp_tool in available}
|
||||
if tool not in valid_names:
|
||||
offered = ", ".join(sorted(valid_names)) or "(none)"
|
||||
@@ -164,8 +180,7 @@ async def call_mcp(
|
||||
f"Tools this connection offers: {offered}. "
|
||||
"Call describe_mcp for their input schemas."
|
||||
)
|
||||
return await dispatch_mcp_call(
|
||||
entry.server,
|
||||
return await entry.session.dispatch(
|
||||
tool,
|
||||
arguments or {},
|
||||
label=namespaced_tool_name(connection, tool),
|
||||
|
||||
@@ -30,11 +30,17 @@ from agents.mcp import (
|
||||
create_static_tool_filter,
|
||||
)
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.shared._httpx_utils import create_mcp_http_client
|
||||
|
||||
from strix.tools.mcp.failures import HttpStatusRecorder
|
||||
from strix.tools.mcp.session import McpConnectionUnavailableError, SupervisedMcpSession
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
import httpx
|
||||
|
||||
from strix.tools.mcp.config import McpConnectionConfig
|
||||
from strix.tools.mcp.registry import McpConnectionRequest, McpRegistry
|
||||
|
||||
@@ -52,22 +58,30 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ConnectedMcpServer(NamedTuple):
|
||||
"""One successfully connected MCP server and how many tools it offers.
|
||||
"""One successfully connected MCP connection and how many tools it offers.
|
||||
|
||||
``server`` is kept so the caller can clean it up when the run ends, and so
|
||||
the caller can hand the live session to the run's
|
||||
``session`` is the :class:`~strix.tools.mcp.session.SupervisedMcpSession` that
|
||||
owns the live connection on its own task, so the caller cleans it up when the
|
||||
run ends (``await session.aclose()``) and hands it to the run's
|
||||
:class:`~strix.tools.mcp.registry.McpRegistry`; ``name`` and ``tool_count``
|
||||
let the caller show the user a startup summary and fill the prompt inventory;
|
||||
``notes`` carries the connection's optional free-text description so the
|
||||
caller can surface it as the connection's purpose in the inventory.
|
||||
"""
|
||||
|
||||
server: MCPServer
|
||||
session: SupervisedMcpSession
|
||||
name: str
|
||||
tool_count: int
|
||||
notes: str | None = None
|
||||
|
||||
|
||||
class BuiltMcpServer(NamedTuple):
|
||||
"""A constructed SDK server and its optional HTTP failure recorder."""
|
||||
|
||||
server: MCPServer
|
||||
recorder: HttpStatusRecorder | None
|
||||
|
||||
|
||||
def _auth_headers(config: McpConnectionConfig) -> dict[str, str]:
|
||||
"""Build the per-server request headers from the connection's auth."""
|
||||
auth = config.auth
|
||||
@@ -105,9 +119,12 @@ class _QuietMCPServerStdio(MCPServerStdio):
|
||||
return _quiet_stdio_streams(self.params)
|
||||
|
||||
|
||||
def _build_server(config: McpConnectionConfig) -> MCPServer:
|
||||
def _build_server(config: McpConnectionConfig) -> BuiltMcpServer:
|
||||
"""Construct (but do not connect) the SDK server for one connection.
|
||||
|
||||
The returned tuple carries the server and, for HTTP connections, a recorder
|
||||
that retains sanitized response metadata for the owning session.
|
||||
|
||||
When ``allowed_tools`` is a list the static filter means the server will not
|
||||
even list tools outside it, so it is the authoritative gate on what
|
||||
``describe_mcp`` and ``call_mcp`` can see. When it is ``None`` no filter is
|
||||
@@ -125,22 +142,43 @@ def _build_server(config: McpConnectionConfig) -> MCPServer:
|
||||
"args": config.args,
|
||||
"env": config.env,
|
||||
}
|
||||
return _QuietMCPServerStdio(
|
||||
params=stdio_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
return BuiltMcpServer(
|
||||
_QuietMCPServerStdio(
|
||||
params=stdio_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
recorder = HttpStatusRecorder()
|
||||
|
||||
def httpx_client_factory(
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: httpx.Timeout | None = None,
|
||||
auth: httpx.Auth | None = None,
|
||||
) -> httpx.AsyncClient:
|
||||
client = create_mcp_http_client(headers=headers, timeout=timeout, auth=auth)
|
||||
client.event_hooks.setdefault("response", []).append(recorder)
|
||||
return client
|
||||
|
||||
http_params: MCPServerStreamableHttpParams = {
|
||||
"url": cast("str", config.url),
|
||||
"headers": _auth_headers(config),
|
||||
"timeout": config.http_timeout_seconds,
|
||||
"sse_read_timeout": config.sse_read_timeout_seconds,
|
||||
"httpx_client_factory": httpx_client_factory,
|
||||
}
|
||||
return MCPServerStreamableHttp(
|
||||
params=http_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
return BuiltMcpServer(
|
||||
MCPServerStreamableHttp(
|
||||
params=http_params,
|
||||
name=config.name,
|
||||
tool_filter=tool_filter,
|
||||
cache_tools_list=True,
|
||||
client_session_timeout_seconds=config.session_timeout_seconds,
|
||||
),
|
||||
recorder,
|
||||
)
|
||||
|
||||
|
||||
@@ -227,66 +265,75 @@ def _errored_tool_output(tool_output: Any) -> dict[str, Any]:
|
||||
return {"success": False, "content": tool_output}
|
||||
|
||||
|
||||
async def _count_server_tools(config: McpConnectionConfig, server: MCPServer) -> int:
|
||||
"""Count a connected server's reachable tools for the startup summary.
|
||||
async def _count_session_tools(config: McpConnectionConfig, session: SupervisedMcpSession) -> int:
|
||||
"""Count a connected session's reachable tools for the startup summary.
|
||||
|
||||
``allowed_tools`` of ``None`` counts every listed tool; a list counts only
|
||||
those names. The count matches what ``describe_mcp`` will show, because the
|
||||
static tool filter built in :func:`_build_server` restricts the server's own
|
||||
``list_tools`` to the same allowlist.
|
||||
``list_tools`` to the same allowlist. The listing goes through the session's
|
||||
owning task like every other call.
|
||||
"""
|
||||
allowed = config.allowed_tools
|
||||
mcp_tools = await server.list_tools()
|
||||
mcp_tools = await session.list_tools()
|
||||
return sum(1 for mcp_tool in mcp_tools if allowed is None or mcp_tool.name in allowed)
|
||||
|
||||
|
||||
async def connect_mcp_servers(
|
||||
configs: list[McpConnectionConfig],
|
||||
) -> list[ConnectedMcpServer]:
|
||||
"""Connect to each MCP server and return its live session.
|
||||
"""Connect each MCP config on its own supervising task and return the sessions.
|
||||
|
||||
Returns one :class:`ConnectedMcpServer` per server that connected, carrying
|
||||
the SDK server (so the caller can clean it up when the run ends and hand it to
|
||||
the run's registry) plus the server name, how many tools it offers, and the
|
||||
connection's notes. Connections that fail are skipped rather than raised.
|
||||
Each connection becomes a :class:`~strix.tools.mcp.session.SupervisedMcpSession`
|
||||
that owns ``connect()``, the held-open session, and ``cleanup()`` on one
|
||||
dedicated task, so a later background failure in one session is contained to
|
||||
that task and never cancels the run. Returns one :class:`ConnectedMcpServer`
|
||||
per session that connected, carrying the session (the caller closes it with
|
||||
``await session.aclose()`` when the run ends and hands it to the run's
|
||||
registry) plus the connection name, tool count, and notes. A connection whose
|
||||
initial connect fails is skipped rather than raised (fail-open).
|
||||
|
||||
If this coroutine is itself cancelled mid-attach (the run going down), every
|
||||
session started so far is closed on its own task before the cancellation is
|
||||
re-raised, so nothing is orphaned.
|
||||
|
||||
Nothing is registered as an agent tool: the caller builds a per-run
|
||||
:class:`~strix.tools.mcp.registry.McpRegistry` from these sessions, and the
|
||||
agent reaches each tool on demand through ``describe_mcp`` / ``call_mcp``.
|
||||
"""
|
||||
connected: list[ConnectedMcpServer] = []
|
||||
for config in configs:
|
||||
server: MCPServer | None = None
|
||||
try:
|
||||
server = _build_server(config)
|
||||
await server.connect() # type: ignore[no-untyped-call]
|
||||
tool_count = await _count_server_tools(config, server)
|
||||
except Exception:
|
||||
logger.exception("Skipping MCP connection %r", config.name)
|
||||
if server is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
continue
|
||||
except BaseException:
|
||||
# A cancellation (or other non-Exception failure) mid-connect must not
|
||||
# orphan MCP subprocesses or HTTP sessions. Clean up the server being
|
||||
# connected and every server already connected, then re-raise so the
|
||||
# caller still stops. The runner only receives the list on a clean
|
||||
# return, so on an abnormal exit this function owns the cleanup.
|
||||
if server is not None:
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
for established in connected:
|
||||
with contextlib.suppress(Exception):
|
||||
await established.server.cleanup() # type: ignore[no-untyped-call]
|
||||
raise
|
||||
|
||||
logger.info("Connected MCP server %r (%d tools)", config.name, tool_count)
|
||||
connected.append(
|
||||
ConnectedMcpServer(
|
||||
server=server, name=config.name, tool_count=tool_count, notes=config.notes
|
||||
sessions: list[SupervisedMcpSession] = []
|
||||
try:
|
||||
for config in configs:
|
||||
session = SupervisedMcpSession(config)
|
||||
sessions.append(session)
|
||||
if not await session.start():
|
||||
# Initial connect failed; already logged inside the session. Drop it.
|
||||
await session.aclose()
|
||||
sessions.remove(session)
|
||||
continue
|
||||
try:
|
||||
tool_count = await _count_session_tools(config, session)
|
||||
except McpConnectionUnavailableError:
|
||||
# The session died between connecting and its first listing; skip it.
|
||||
logger.warning("MCP connection %r died before its first listing", config.name)
|
||||
await session.aclose()
|
||||
sessions.remove(session)
|
||||
continue
|
||||
logger.info("Connected MCP server %r (%d tools)", config.name, tool_count)
|
||||
connected.append(
|
||||
ConnectedMcpServer(
|
||||
session=session, name=config.name, tool_count=tool_count, notes=config.notes
|
||||
)
|
||||
)
|
||||
)
|
||||
except BaseException:
|
||||
# Cancelled or errored mid-attach: close every session started so far,
|
||||
# each on its own task, then re-raise. The runner only receives the list
|
||||
# on a clean return, so on an abnormal exit this function owns the cleanup.
|
||||
for session in sessions:
|
||||
with contextlib.suppress(BaseException):
|
||||
await session.aclose()
|
||||
raise
|
||||
|
||||
return connected
|
||||
|
||||
@@ -318,7 +365,7 @@ async def attach_mcp_requests(
|
||||
request = request_by_name[connection.name]
|
||||
registry.add(
|
||||
name=connection.name,
|
||||
server=connection.server,
|
||||
session=connection.session,
|
||||
tool_count=connection.tool_count,
|
||||
purpose=request.purpose or connection.notes,
|
||||
provider=request.provider,
|
||||
|
||||
@@ -12,6 +12,9 @@ from typing import Annotated, Literal
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
|
||||
DEFAULT_MAX_CONCURRENT_CALLS = 4
|
||||
|
||||
|
||||
class BearerAuth(BaseModel):
|
||||
"""Header-token auth, sent as ``Authorization: Bearer <token>``."""
|
||||
|
||||
@@ -65,6 +68,18 @@ class McpConnectionConfig(BaseModel):
|
||||
MCP inventory every agent renders in its prompt, so it describes the
|
||||
connection once rather than being repeated onto each of its tools."""
|
||||
|
||||
http_timeout_seconds: float = Field(default=30.0, gt=0)
|
||||
"""HTTP request timeout; the SDK's 5-second default is below tool p95s."""
|
||||
|
||||
sse_read_timeout_seconds: float = Field(default=300.0, gt=0)
|
||||
"""Stream read timeout; the SDK's 5-second default is below tool p95s."""
|
||||
|
||||
session_timeout_seconds: float = Field(default=60.0, gt=0)
|
||||
"""MCP operation timeout for SQL queries and cloud describe fan-outs."""
|
||||
|
||||
max_concurrent_calls: int = Field(default=DEFAULT_MAX_CONCURRENT_CALLS, ge=1)
|
||||
"""Maximum concurrent calls for this connection name across sessions."""
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_transport_fields(self) -> McpConnectionConfig:
|
||||
if self.transport == "http" and not self.url:
|
||||
|
||||
150
strix/tools/mcp/failures.py
Normal file
150
strix/tools/mcp/failures.py
Normal file
@@ -0,0 +1,150 @@
|
||||
"""Classify MCP connection failures without retaining sensitive request data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from email.utils import parsedate_to_datetime
|
||||
from typing import Literal, cast
|
||||
|
||||
import httpx
|
||||
from agents.exceptions import UserError
|
||||
from mcp.shared.exceptions import McpError
|
||||
|
||||
|
||||
FailureKind = Literal[
|
||||
"auth", "permission", "rate_limit", "server", "transport", "timeout", "protocol", "unknown"
|
||||
]
|
||||
|
||||
_PRIORITY: dict[FailureKind, int] = {
|
||||
"auth": 0,
|
||||
"permission": 1,
|
||||
"rate_limit": 2,
|
||||
"server": 3,
|
||||
"protocol": 4,
|
||||
"timeout": 5,
|
||||
"transport": 6,
|
||||
"unknown": 7,
|
||||
}
|
||||
_HTTP_ERROR_RE = re.compile(r"\bHTTP error\s+(\d{3})\b", re.IGNORECASE)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FailureInfo:
|
||||
"""A non-sensitive description of one connection failure."""
|
||||
|
||||
kind: FailureKind
|
||||
status: int | None = None
|
||||
reason: str | None = None
|
||||
retry_after: float | None = None
|
||||
request_method: str | None = None
|
||||
request_path: str | None = None
|
||||
|
||||
@property
|
||||
def retryable(self) -> bool:
|
||||
return self.kind not in {"auth", "permission"}
|
||||
|
||||
|
||||
def _retry_after(value: str | None) -> float | None:
|
||||
if not value:
|
||||
return None
|
||||
try:
|
||||
return max(0.0, float(value))
|
||||
except ValueError:
|
||||
pass
|
||||
try:
|
||||
date = parsedate_to_datetime(value)
|
||||
if date.tzinfo is None:
|
||||
date = date.replace(tzinfo=UTC)
|
||||
return max(0.0, (date - datetime.now(UTC)).total_seconds())
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
return None
|
||||
|
||||
|
||||
def _from_status(
|
||||
status: int,
|
||||
reason: str | None = None,
|
||||
retry_after: float | None = None,
|
||||
*,
|
||||
request_method: str | None = None,
|
||||
request_path: str | None = None,
|
||||
) -> FailureInfo:
|
||||
if status == 401:
|
||||
kind: FailureKind = "auth"
|
||||
elif status == 403:
|
||||
kind = "permission"
|
||||
elif status == 429:
|
||||
kind = "rate_limit"
|
||||
elif 500 <= status <= 599:
|
||||
kind = "server"
|
||||
elif 400 <= status <= 499:
|
||||
kind = "protocol"
|
||||
else:
|
||||
kind = "unknown"
|
||||
return FailureInfo(
|
||||
kind,
|
||||
status,
|
||||
reason,
|
||||
retry_after,
|
||||
request_method,
|
||||
request_path,
|
||||
)
|
||||
|
||||
|
||||
def _direct(exc: BaseException) -> FailureInfo | None:
|
||||
if isinstance(exc, httpx.HTTPStatusError):
|
||||
response = exc.response
|
||||
request = response.request
|
||||
return _from_status(
|
||||
response.status_code,
|
||||
response.reason_phrase,
|
||||
_retry_after(response.headers.get("Retry-After")),
|
||||
request_method=request.method,
|
||||
request_path=request.url.path,
|
||||
)
|
||||
if isinstance(exc, httpx.TimeoutException):
|
||||
return FailureInfo("timeout", reason="request timed out")
|
||||
if isinstance(exc, httpx.TransportError):
|
||||
return FailureInfo("transport", reason="transport error")
|
||||
if isinstance(exc, McpError):
|
||||
return FailureInfo("protocol", reason="MCP protocol error")
|
||||
if isinstance(exc, UserError):
|
||||
match = _HTTP_ERROR_RE.search(str(exc))
|
||||
if match:
|
||||
return _from_status(int(match.group(1)))
|
||||
return None
|
||||
|
||||
|
||||
def classify(exc: BaseException) -> FailureInfo:
|
||||
"""Return the most specific non-sensitive classification in an exception tree."""
|
||||
direct = _direct(exc)
|
||||
matches: list[FailureInfo] = [direct] if direct is not None else []
|
||||
if isinstance(exc, BaseExceptionGroup):
|
||||
group = cast("BaseExceptionGroup[BaseException]", exc)
|
||||
matches.extend(classify(child) for child in group.exceptions)
|
||||
if matches:
|
||||
return min(matches, key=lambda info: _PRIORITY[info.kind])
|
||||
return FailureInfo("unknown", reason="unknown failure")
|
||||
|
||||
|
||||
class HttpStatusRecorder:
|
||||
"""Capture the last non-success response from one HTTP connection."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._failure: FailureInfo | None = None
|
||||
|
||||
async def __call__(self, response: httpx.Response) -> None:
|
||||
if not 200 <= response.status_code < 300:
|
||||
request = response.request
|
||||
self._failure = _from_status(
|
||||
response.status_code,
|
||||
response.reason_phrase,
|
||||
_retry_after(response.headers.get("Retry-After")),
|
||||
request_method=request.method,
|
||||
request_path=request.url.path,
|
||||
)
|
||||
|
||||
def take(self) -> FailureInfo | None:
|
||||
failure, self._failure = self._failure, None
|
||||
return failure
|
||||
@@ -25,6 +25,8 @@ from __future__ import annotations
|
||||
import dataclasses
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple
|
||||
|
||||
from strix.tools.mcp.session import SupervisedMcpSession
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agents.mcp import MCPServer
|
||||
@@ -54,24 +56,45 @@ MCP_DISPATCH_TOOLS = frozenset({CALL_MCP_TOOL, DESCRIBE_MCP_TOOL})
|
||||
class McpConnectionEntry:
|
||||
"""One live MCP connection a scan may reach, keyed by ``name``.
|
||||
|
||||
``server`` is the connected SDK session the dispatch tools list tools on and
|
||||
call tools through. ``purpose`` is the human label ``list_mcps`` reports as the
|
||||
connection's description (the user's connection notes, or whatever the caller
|
||||
supplies). ``tool_count`` is how many tools the connection offers, also
|
||||
reported by ``list_mcps``. ``result_transform``, when set, runs on each call's structured result
|
||||
at the single dispatch point (strix-pro's sanitizer uses it). ``provider`` is
|
||||
an optional source label (e.g. ``"supabase"``) the caller tags the connection
|
||||
with; the command-line path leaves it ``None``, and event tagging surfaces it
|
||||
when set.
|
||||
``session`` is the :class:`~strix.tools.mcp.session.SupervisedMcpSession` that
|
||||
owns the connection on its own task; the dispatch tools list tools and call
|
||||
tools through it (``session.list_tools`` / ``session.dispatch``) so a session
|
||||
failure is contained and can reconnect. ``purpose`` is the human label
|
||||
``list_mcps`` reports as the connection's description (the user's connection
|
||||
notes, or whatever the caller supplies). ``tool_count`` is how many tools the
|
||||
connection offers, also reported by ``list_mcps``. ``result_transform``, when
|
||||
set, runs on each call's structured result at the single dispatch point
|
||||
(strix-pro's sanitizer uses it). ``provider`` is an optional source label
|
||||
(e.g. ``"supabase"``) the caller tags the connection with; the command-line
|
||||
path leaves it ``None``, and event tagging surfaces it when set.
|
||||
|
||||
The connection config the session reconnects with (and its bearer token) lives
|
||||
on ``session`` in memory only. It is reached via :attr:`config` for the
|
||||
reconnect path and is never logged, serialized into the event stream, or
|
||||
written to disk.
|
||||
"""
|
||||
|
||||
server: MCPServer
|
||||
session: SupervisedMcpSession
|
||||
name: str
|
||||
purpose: str | None = None
|
||||
tool_count: int = 0
|
||||
result_transform: ResultTransform | None = None
|
||||
provider: str | None = None
|
||||
|
||||
@property
|
||||
def server(self) -> MCPServer | None:
|
||||
"""The current live server behind the session (swapped on reconnect).
|
||||
|
||||
Kept so existing callers that read ``entry.server`` keep working; new code
|
||||
should call through ``entry.session`` so reconnect and containment apply.
|
||||
"""
|
||||
return self.session.server
|
||||
|
||||
@property
|
||||
def config(self) -> McpConnectionConfig | None:
|
||||
"""The session's reconnect config. Carries the bearer token; never log it."""
|
||||
return self.session.config
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class McpConnectionSummary:
|
||||
@@ -84,6 +107,24 @@ class McpConnectionSummary:
|
||||
provider: str | None = None
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class McpConnectionStatus:
|
||||
"""One connection's live status for the interfaces (the TUI panel, the app
|
||||
strip, and the roster signal the app consumes).
|
||||
|
||||
Non-secret by construction: only the connection ``name``, its ``provider``
|
||||
label, its ``tool_count``, and whether its live session is currently ``dead``
|
||||
(its reconnect-retry gave up). No config, token, url, or purpose rides here.
|
||||
``dead`` is read live off the connection's session at the moment this is
|
||||
built, so a fresh :meth:`McpRegistry.statuses` reflects the current health.
|
||||
"""
|
||||
|
||||
name: str
|
||||
provider: str | None
|
||||
tool_count: int
|
||||
dead: bool
|
||||
|
||||
|
||||
@dataclasses.dataclass(frozen=True)
|
||||
class McpConnectionRequest:
|
||||
"""A source-agnostic request to attach one MCP connection to a run.
|
||||
@@ -129,15 +170,28 @@ class McpRegistry:
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
server: MCPServer,
|
||||
session: SupervisedMcpSession | None = None,
|
||||
server: MCPServer | None = None,
|
||||
config: McpConnectionConfig | None = None,
|
||||
purpose: str | None = None,
|
||||
tool_count: int = 0,
|
||||
result_transform: ResultTransform | None = None,
|
||||
provider: str | None = None,
|
||||
) -> McpConnectionEntry:
|
||||
"""Register one connection under ``name`` (last write wins)."""
|
||||
"""Register one connection under ``name`` (last write wins).
|
||||
|
||||
Pass ``session`` for a session the engine already supervises (the attach
|
||||
path does this). Pass ``server`` for an already-connected server the caller
|
||||
owns (strix-pro's cloud sessions): it is adopted into a session that runs
|
||||
calls inline against it, and reconnects only when a ``config`` is also
|
||||
given. Exactly one of ``session`` or ``server`` is required.
|
||||
"""
|
||||
if session is None:
|
||||
if server is None:
|
||||
raise ValueError("McpRegistry.add requires either 'session' or 'server'")
|
||||
session = SupervisedMcpSession.adopt(server, name=name, config=config)
|
||||
entry = McpConnectionEntry(
|
||||
server=server,
|
||||
session=session,
|
||||
name=name,
|
||||
purpose=purpose,
|
||||
tool_count=tool_count,
|
||||
@@ -167,6 +221,23 @@ class McpRegistry:
|
||||
for entry in self._entries.values()
|
||||
]
|
||||
|
||||
def statuses(self) -> list[McpConnectionStatus]:
|
||||
"""One live status per connection, in insertion order.
|
||||
|
||||
Reads each connection's ``dead`` flag off its session at call time, so the
|
||||
interfaces (the TUI panel via the Python backend projection, and the
|
||||
roster signal the app consumes) get the current health each time they
|
||||
rebuild. Non-secret: name, provider, tool_count, dead only."""
|
||||
return [
|
||||
McpConnectionStatus(
|
||||
name=entry.name,
|
||||
provider=entry.provider,
|
||||
tool_count=entry.tool_count,
|
||||
dead=entry.session.is_dead,
|
||||
)
|
||||
for entry in self._entries.values()
|
||||
]
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Drop every connection (the sessions themselves are closed by the
|
||||
runner)."""
|
||||
|
||||
806
strix/tools/mcp/session.py
Normal file
806
strix/tools/mcp/session.py
Normal file
@@ -0,0 +1,806 @@
|
||||
"""Own each MCP connection's live session on its own supervising task.
|
||||
|
||||
The bug this fixes: the streamable-HTTP transport (the ``mcp`` SDK) opens an
|
||||
internal anyio task group when ``server.connect()`` runs, and that task group's
|
||||
cancel scope is entered on whatever task called ``connect()`` and stays open for
|
||||
the session's whole life. In the old code that task was the run's main task, the
|
||||
one the agent loop runs on. So when a provider returned an HTTP error on one of
|
||||
the transport's background tasks (for example a ``403`` on a background POST),
|
||||
the task group cancelled its scope, the cancellation
|
||||
propagated to the main task, and the whole scan died with a bare
|
||||
``CancelledError`` (mislabeled as a user interrupt). Teardown then raised
|
||||
"Attempted to exit cancel scope in a different task than it was entered in"
|
||||
because cleanup ran on a different task than connect.
|
||||
|
||||
The fix, mirroring how child agents run on their own ``asyncio.create_task``
|
||||
(see :func:`strix.core.execution.spawn_child_agent`): give each connection its
|
||||
own dedicated supervising task that owns ``connect()``, the session's held-open
|
||||
lifetime, and ``cleanup()``. Three consequences:
|
||||
|
||||
- **Containment.** The transport's cancel scope is now entered on the supervising
|
||||
task, so a background failure cancels only that task. The run and every other
|
||||
connection keep going.
|
||||
- **Co-located teardown.** ``connect()`` and ``cleanup()`` run on the same task,
|
||||
so the "exit cancel scope in a different task" error cannot happen.
|
||||
- **A value, not a cancellation, reaches the caller.** The agent never touches the
|
||||
live session directly. It hands a call to the supervising task over a queue and
|
||||
awaits the result as a value; if the session task dies, the caller gets a
|
||||
"connection unavailable" value instead of a cancellation propagating into the
|
||||
agent loop.
|
||||
|
||||
Failure handling follows connection-pool discipline: discard on error, rebuild on
|
||||
next use. A failure while connecting or rebuilding describes the session. A
|
||||
non-2xx response from a tool call describes that request, not the session. Permission
|
||||
and protocol failures from a call return a failed tool output while the connection
|
||||
stays usable. Other classified failures are retried on the rebuilt session and then,
|
||||
if they keep failing, temporarily quarantine the connection. Authentication failures
|
||||
and repeated transient exhaustion permanently retire a connection.
|
||||
|
||||
Security: the connection's :class:`~strix.tools.mcp.config.McpConnectionConfig`
|
||||
holds a live bearer credential and is kept here in memory only, on the same
|
||||
in-process object that already holds the live session. It is never logged,
|
||||
serialized into the run's event stream, or written to disk; :meth:`__repr__`
|
||||
omits it and the token field's own ``repr`` is already suppressed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import dataclasses
|
||||
import logging
|
||||
import secrets
|
||||
import time
|
||||
import weakref
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
from strix.tools.mcp.config import DEFAULT_MAX_CONCURRENT_CALLS
|
||||
from strix.tools.mcp.failures import FailureInfo, HttpStatusRecorder, classify
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
from agents.mcp import MCPServer
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from strix.tools.mcp.client import ResultTransform
|
||||
from strix.tools.mcp.config import McpConnectionConfig
|
||||
|
||||
# One operation to run against the live session, e.g. ``list_tools`` or a tool
|
||||
# call. Runs on the supervising task (supervised sessions) or inline (adopted
|
||||
# sessions), and its return value becomes the caller's result.
|
||||
Job = Callable[[MCPServer], Awaitable[Any]]
|
||||
|
||||
_Phase = Literal["connect", "call"]
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# How long a graceful (sentinel) shutdown waits for the serve loop to drain
|
||||
# before the supervising task is cancelled instead. Bounds teardown so a slow or
|
||||
# hung in-flight call cannot stall it forever.
|
||||
_SHUTDOWN_TIMEOUT = 10.0
|
||||
_MAX_ATTEMPTS = 3
|
||||
_SETTLE_DELAY = 0.05
|
||||
_SEMAPHORES: weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, dict[str, asyncio.Semaphore]] = (
|
||||
weakref.WeakKeyDictionary()
|
||||
)
|
||||
_JITTER = secrets.SystemRandom()
|
||||
|
||||
# Everything the SDK can surface for a failed call: ordinary errors plus the
|
||||
# transport's task-group ``BaseExceptionGroup``. Caught wholesale and handed to
|
||||
# ``classify``; ``asyncio.CancelledError`` is always handled separately first,
|
||||
# so shutdown and genuine cancellation still propagate.
|
||||
_CLASSIFIABLE: tuple[type[BaseException], ...] = (BaseExceptionGroup, Exception)
|
||||
|
||||
|
||||
def _retry_delay(attempt: int, retry_after: float | None) -> float:
|
||||
if retry_after is not None:
|
||||
return retry_after
|
||||
base = min(8.0, 0.5 * (2 ** (attempt - 1)))
|
||||
return base + _JITTER.uniform(0.0, base * 0.1) # type: ignore[no-any-return]
|
||||
|
||||
|
||||
def _call_semaphore(name: str, limit: int) -> asyncio.Semaphore:
|
||||
loop = asyncio.get_running_loop()
|
||||
semaphores = _SEMAPHORES.setdefault(loop, {})
|
||||
return semaphores.setdefault(name, asyncio.Semaphore(limit))
|
||||
|
||||
|
||||
class McpConnectionUnavailableError(RuntimeError):
|
||||
"""A dead MCP connection could not be reached and did not come back.
|
||||
|
||||
Raised by :meth:`SupervisedMcpSession.list_tools` when the connection is dead
|
||||
so the read-only dispatch tools (``describe_mcp``) can report it cleanly.
|
||||
:meth:`SupervisedMcpSession.dispatch` does not raise it: a call to a dead
|
||||
connection returns the standard failed-tool output instead.
|
||||
"""
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _Outcome:
|
||||
"""What running one job resolved to: a value, a call failure, or a dead connection."""
|
||||
|
||||
value: Any = None
|
||||
dead: bool = False
|
||||
call_failure: FailureInfo | None = None
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class _Request:
|
||||
"""One job handed to the supervising task, with the future its result lands in."""
|
||||
|
||||
job: Job
|
||||
future: asyncio.Future[_Outcome]
|
||||
phase: _Phase
|
||||
|
||||
|
||||
class SupervisedMcpSession:
|
||||
"""One MCP connection whose live session is owned by a dedicated task.
|
||||
|
||||
Built two ways:
|
||||
|
||||
- :meth:`__init__` + :meth:`start` for a *supervised* session: the engine owns
|
||||
connecting. ``start`` spawns the supervising task, which builds and connects
|
||||
the server on itself and then serves calls handed to it over a queue. This is
|
||||
the path that contains a background session failure to one task.
|
||||
- :meth:`adopt` for an *adopted* session: the caller already holds a connected
|
||||
server (strix-pro's cloud sessions, and the test fakes). There is no
|
||||
supervising task; calls run inline against the given server. Reconnect works
|
||||
only when a config was supplied.
|
||||
|
||||
Public async API used by the dispatch tools: :meth:`list_tools` and
|
||||
:meth:`dispatch`. Lifecycle: :meth:`start`, :meth:`aclose`. Read-only:
|
||||
:attr:`name`, :attr:`server`, :attr:`config`, :attr:`is_dead`.
|
||||
"""
|
||||
|
||||
def __init__(self, config: McpConnectionConfig) -> None:
|
||||
self._name = config.name
|
||||
self._config: McpConnectionConfig | None = config
|
||||
self._server: MCPServer | None = None
|
||||
self._supervised = True
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
self._queue: asyncio.Queue[_Request | None] | None = None
|
||||
self._ready: asyncio.Future[bool] | None = None
|
||||
self._pending: set[asyncio.Future[_Outcome]] = set()
|
||||
self._dead = False
|
||||
self._closing = False
|
||||
self._on_dead: Callable[[], None] | None = None
|
||||
self._recorder: HttpStatusRecorder | None = None
|
||||
self._unavailable_until: float | None = None
|
||||
self._quarantine_count = 0
|
||||
self._last_failure = FailureInfo("unknown", reason="connection unavailable")
|
||||
self._reconnect_lock = asyncio.Lock()
|
||||
self._call_semaphore: asyncio.Semaphore | None = None
|
||||
|
||||
@classmethod
|
||||
def adopt(
|
||||
cls,
|
||||
server: MCPServer,
|
||||
*,
|
||||
name: str,
|
||||
config: McpConnectionConfig | None = None,
|
||||
) -> SupervisedMcpSession:
|
||||
"""Wrap an already-connected server without a supervising task.
|
||||
|
||||
Calls run inline against ``server`` on the caller's task, matching the old
|
||||
direct-dispatch behavior. Reconnect is available only when ``config`` is
|
||||
given; otherwise a failed call can be quarantined but cannot be revived.
|
||||
"""
|
||||
self = cls.__new__(cls)
|
||||
self._name = name
|
||||
self._config = config
|
||||
self._server = server
|
||||
self._supervised = False
|
||||
self._task = None
|
||||
self._queue = None
|
||||
self._ready = None
|
||||
self._pending = set()
|
||||
self._dead = False
|
||||
self._closing = False
|
||||
self._on_dead = None
|
||||
self._recorder = None
|
||||
self._unavailable_until = None
|
||||
self._quarantine_count = 0
|
||||
self._last_failure = FailureInfo("unknown", reason="connection unavailable")
|
||||
self._reconnect_lock = asyncio.Lock()
|
||||
self._call_semaphore = None
|
||||
return self
|
||||
|
||||
# -- read-only accessors --------------------------------------------------
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return self._name
|
||||
|
||||
@property
|
||||
def server(self) -> MCPServer | None:
|
||||
"""The current live server, or ``None`` once dead. Swapped on reconnect."""
|
||||
return self._server
|
||||
|
||||
@property
|
||||
def config(self) -> McpConnectionConfig | None:
|
||||
"""The connection config kept for reconnect. Carries the bearer token, so
|
||||
never log or serialize this."""
|
||||
return self._config
|
||||
|
||||
@property
|
||||
def is_dead(self) -> bool:
|
||||
return self._dead
|
||||
|
||||
@property
|
||||
def is_unavailable(self) -> bool:
|
||||
"""Whether the connection is temporarily quarantined."""
|
||||
return (
|
||||
not self._dead
|
||||
and self._unavailable_until is not None
|
||||
and time.monotonic() < self._unavailable_until
|
||||
)
|
||||
|
||||
def set_on_dead(self, callback: Callable[[], None] | None) -> None:
|
||||
"""Register a one-shot callback fired when the connection transitions to dead.
|
||||
|
||||
The callback runs on whatever task marks the connection dead (the
|
||||
supervising task for a supervised session, the caller's task for an
|
||||
adopted one), so it must not block. It fires at most once, on the
|
||||
healthy->dead edge, and never for a connection that only ever shut down
|
||||
cleanly. The interfaces use it to push a live "offline" status without
|
||||
polling. Exceptions from the callback are swallowed (logged) so a status
|
||||
push can never take down the session task.
|
||||
"""
|
||||
self._on_dead = callback
|
||||
|
||||
def _mark_dead(self, failure: FailureInfo | None = None, *, attempt: int = 1) -> None:
|
||||
"""Flip the connection to dead and fire ``on_dead`` once on the transition."""
|
||||
if self._dead:
|
||||
return
|
||||
failure = failure or self._last_failure
|
||||
self._dead = True
|
||||
self._unavailable_until = None
|
||||
logger.error(
|
||||
"MCP connection %r permanently unavailable kind=%s status=%s reason=%s "
|
||||
"attempt=%d delay=0",
|
||||
self._name,
|
||||
failure.kind,
|
||||
failure.status,
|
||||
failure.reason,
|
||||
attempt,
|
||||
)
|
||||
callback = self._on_dead
|
||||
if callback is None:
|
||||
return
|
||||
try:
|
||||
callback()
|
||||
except Exception:
|
||||
logger.exception("MCP on_dead callback for %r failed", self._name)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
# Deliberately omits the config so the bearer token can never reach a log
|
||||
# line through an accidental repr of this object.
|
||||
return f"SupervisedMcpSession(name={self._name!r}, dead={self._dead})"
|
||||
|
||||
# -- lifecycle ------------------------------------------------------------
|
||||
|
||||
async def start(self) -> bool:
|
||||
"""Spawn the supervising task, connect on it, and wait until it is ready.
|
||||
|
||||
Returns ``True`` when the session connected, ``False`` when the initial
|
||||
connect failed (the caller then skips this connection, fail-open). Only
|
||||
valid for a supervised session.
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
self._queue = asyncio.Queue()
|
||||
self._ready = loop.create_future()
|
||||
self._task = asyncio.create_task(self._supervise(), name=f"mcp-session-{self._name}")
|
||||
return await self._ready
|
||||
|
||||
async def aclose(self) -> None:
|
||||
"""Shut the connection down and clean up its session on its owning task.
|
||||
|
||||
For a connected supervised session this signals the supervising task with a
|
||||
sentinel so ``cleanup()`` runs on the same task that ran ``connect()``,
|
||||
giving an orderly shutdown the supervisor tells apart from a session death.
|
||||
Teardown is always bounded: if the serve loop cannot drain the sentinel in
|
||||
time (a slow or hung in-flight call), or the session never finished
|
||||
connecting (including a connect cancelled mid-await), the task is cancelled
|
||||
instead. ``_closing`` is set first, so the supervisor treats that
|
||||
cancellation as shutdown and still cleans up on its own task.
|
||||
"""
|
||||
self._closing = True
|
||||
if self._supervised and self._task is not None:
|
||||
if not self._task.done():
|
||||
# A cancelled readiness future (the connect was cancelled mid-await)
|
||||
# counts as "not connected": never call ``.result()`` on it, which
|
||||
# would raise here and skip the cleanup below.
|
||||
connected = (
|
||||
self._ready is not None
|
||||
and self._ready.done()
|
||||
and not self._ready.cancelled()
|
||||
and self._ready.result()
|
||||
)
|
||||
if connected and self._queue is not None:
|
||||
# Reached the serve loop: a sentinel gives a clean, cancel-free
|
||||
# teardown, with cleanup() running on the supervising task. Bound
|
||||
# it, though: a hung in-flight call would otherwise leave the
|
||||
# sentinel queued behind it forever, so cancel the task if the
|
||||
# drain does not finish in time (wait_for cancels it on timeout).
|
||||
with contextlib.suppress(Exception):
|
||||
await self._queue.put(None)
|
||||
with contextlib.suppress(
|
||||
asyncio.TimeoutError, asyncio.CancelledError, Exception
|
||||
):
|
||||
await asyncio.wait_for(self._task, _SHUTDOWN_TIMEOUT)
|
||||
else:
|
||||
# Still stuck in connect(), never connected, or connect
|
||||
# cancelled: cancel to unstick it.
|
||||
self._task.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError, Exception):
|
||||
await self._task
|
||||
else:
|
||||
await self._safe_cleanup()
|
||||
self._fail_pending()
|
||||
|
||||
# -- caller-facing operations --------------------------------------------
|
||||
|
||||
async def list_tools(self) -> list[MCPTool]:
|
||||
"""List the connection's tools, retrying transient session failures.
|
||||
|
||||
Raises :class:`McpConnectionUnavailableError` when the connection is dead
|
||||
and never returns a call failure.
|
||||
"""
|
||||
outcome = await self._run_job(lambda server: server.list_tools(), phase="connect")
|
||||
if outcome.dead:
|
||||
raise McpConnectionUnavailableError(self._unavailable_message())
|
||||
if outcome.call_failure is not None:
|
||||
raise RuntimeError("MCP list_tools returned a call failure")
|
||||
return cast("list[MCPTool]", outcome.value)
|
||||
|
||||
async def dispatch(
|
||||
self,
|
||||
tool_name: str,
|
||||
arguments: dict[str, Any],
|
||||
*,
|
||||
label: str,
|
||||
result_transform: ResultTransform | None = None,
|
||||
) -> Any:
|
||||
"""Run one tool call with bounded retries for transient session failures.
|
||||
|
||||
Returns the tool output on success, or the standard failed-tool output
|
||||
(``success: False``) when the provider rejects the call or the connection
|
||||
is unavailable. A call rejection keeps the connection usable because the
|
||||
provider rejected the request, not the session.
|
||||
"""
|
||||
from strix.tools.mcp.client import dispatch_mcp_call
|
||||
|
||||
async def job(server: MCPServer) -> Any:
|
||||
return await dispatch_mcp_call(
|
||||
server,
|
||||
tool_name,
|
||||
arguments,
|
||||
label=label,
|
||||
result_transform=result_transform,
|
||||
)
|
||||
|
||||
outcome = await self._run_job(job, phase="call")
|
||||
if outcome.call_failure is not None:
|
||||
from strix.tools.mcp.client import _errored_tool_output
|
||||
|
||||
return _errored_tool_output(self._call_rejected_message(outcome.call_failure))
|
||||
if outcome.dead:
|
||||
from strix.tools.mcp.client import _errored_tool_output
|
||||
|
||||
return _errored_tool_output(self._unavailable_message())
|
||||
return outcome.value
|
||||
|
||||
# -- job routing ----------------------------------------------------------
|
||||
|
||||
async def _run_job(self, job: Job, *, phase: _Phase) -> _Outcome:
|
||||
"""Route one job to the owning task (supervised) or run it inline (adopted)."""
|
||||
if self._supervised:
|
||||
return await self._submit(job, phase)
|
||||
return await self._execute(job, phase)
|
||||
|
||||
async def _submit(self, job: Job, phase: _Phase) -> _Outcome:
|
||||
"""Hand a job to the supervising task and await its result as a value."""
|
||||
if self._dead or self._closing or self._task is None or self._task.done():
|
||||
return _Outcome(dead=True)
|
||||
loop = asyncio.get_running_loop()
|
||||
future: asyncio.Future[_Outcome] = loop.create_future()
|
||||
self._pending.add(future)
|
||||
if self._queue is None:
|
||||
self._pending.discard(future)
|
||||
return _Outcome(dead=True)
|
||||
await self._queue.put(_Request(job=job, future=future, phase=phase))
|
||||
# The task may have ended between the guard above and the put; ``_fail_pending``
|
||||
# would then never see this future, so resolve it here.
|
||||
if self._task.done() and not future.done():
|
||||
self._pending.discard(future)
|
||||
return _Outcome(dead=True)
|
||||
return await future
|
||||
|
||||
# -- the supervising task -------------------------------------------------
|
||||
|
||||
async def _supervise(self) -> None:
|
||||
"""Own the session for its whole life on one task: connect, serve, clean up."""
|
||||
try:
|
||||
self._server = await self._open()
|
||||
except asyncio.CancelledError:
|
||||
# The connect was cancelled (the run is going down, or the transport
|
||||
# scope cancelled mid-connect). Report not-ready so the attach path
|
||||
# treats it as a skipped connection; do not propagate.
|
||||
self._report_ready(value=False)
|
||||
await self._safe_cleanup()
|
||||
self._fail_pending()
|
||||
return
|
||||
except _CLASSIFIABLE as exc:
|
||||
failure = classify(exc)
|
||||
logger.warning(
|
||||
"Skipping MCP connection %r kind=%s status=%s attempt=1 delay=0",
|
||||
self._name,
|
||||
failure.kind,
|
||||
failure.status,
|
||||
exc_info=True,
|
||||
)
|
||||
self._report_ready(value=False)
|
||||
await self._safe_cleanup()
|
||||
self._fail_pending()
|
||||
return
|
||||
|
||||
self._report_ready(value=True)
|
||||
try:
|
||||
await self._serve_loop()
|
||||
finally:
|
||||
await self._safe_cleanup()
|
||||
self._fail_pending()
|
||||
|
||||
async def _serve_loop(self) -> None:
|
||||
assert self._queue is not None
|
||||
while True:
|
||||
try:
|
||||
request = await self._queue.get()
|
||||
except asyncio.CancelledError:
|
||||
# A cancellation while idle is the transport's task group cancelling
|
||||
# this supervising task because a background session task failed.
|
||||
# Contained here. If we are closing, this is an ordinary shutdown,
|
||||
# so let it propagate. Otherwise quarantine the failed session and
|
||||
# keep serving requests so a later call can revive it.
|
||||
if self._closing:
|
||||
raise
|
||||
failure = self._recorder.take() if self._recorder is not None else None
|
||||
failure = failure or FailureInfo("transport", reason="session cancelled")
|
||||
self._last_failure = failure
|
||||
if failure.kind in {"auth", "permission"}:
|
||||
self._mark_dead(failure, attempt=1)
|
||||
return
|
||||
await self._quarantine(failure, attempt=1)
|
||||
if self._dead:
|
||||
return
|
||||
continue
|
||||
if request is None: # shutdown sentinel
|
||||
return
|
||||
outcome = await self._execute(request.job, request.phase)
|
||||
if not request.future.done():
|
||||
request.future.set_result(outcome)
|
||||
self._pending.discard(request.future)
|
||||
if self._dead:
|
||||
return
|
||||
|
||||
# -- run one job with bounded classified retries --------------------------
|
||||
|
||||
async def _execute(self, job: Job, phase: _Phase) -> _Outcome: # noqa: PLR0912
|
||||
"""Run one job on a healthy session, disposing it the instant it errors.
|
||||
|
||||
Discard-on-error, rebuild-on-next-use is the whole discipline here, and it
|
||||
rests on one invariant: **a session object is only ever awaited while
|
||||
healthy.** The moment a call fails, the very next thing this method does,
|
||||
before any other ``await`` including the backoff sleep inside
|
||||
:meth:`_handle_failure`, is dispose that session on this task
|
||||
(:meth:`_safe_cleanup` runs the transport teardown and clears ``_server``).
|
||||
|
||||
Why the ordering is the crux, not a nicety: when a provider returns a non-2xx
|
||||
status mid-call, the streamable-HTTP transport's task group cancels its scope,
|
||||
which cancels this supervising task; the failure surfaces as a
|
||||
``CancelledError`` and the scope keeps firing (re-raising on every subsequent
|
||||
``await``) until the session is torn down. Disposing closes the transport's
|
||||
AsyncExitStack, which exits that firing scope. If instead we slept for backoff
|
||||
first, the sleep would re-raise the firing ``CancelledError``, escape this
|
||||
method, and kill the supervising task, leaving the slot wedged with
|
||||
``is_dead`` False forever. Disposing first is what turns a failure into a
|
||||
returned value and keeps the task alive to rebuild on the next attempt.
|
||||
|
||||
The rebuild itself happens lazily at the top of the loop: once a failure has
|
||||
set ``_server`` to None, the next iteration builds a fresh session (guarded by
|
||||
:meth:`_reconnect`) and retries the operation on it. A permission or protocol
|
||||
failure from a call returns immediately after disposal because it describes
|
||||
that request, not the session. A genuine shutdown (``_closing``) and a real
|
||||
external cancellation still propagate; only the transport's teardown
|
||||
cancellation is contained.
|
||||
"""
|
||||
if self._dead:
|
||||
return _Outcome(dead=True)
|
||||
if self._unavailable_until is not None:
|
||||
remaining = self._unavailable_until - time.monotonic()
|
||||
if remaining > 0:
|
||||
return _Outcome(dead=True)
|
||||
self._unavailable_until = None
|
||||
logger.info(
|
||||
"MCP connection %r revive started kind=%s status=%s attempt=1",
|
||||
self._name,
|
||||
self._last_failure.kind,
|
||||
self._last_failure.status,
|
||||
)
|
||||
|
||||
if self._call_semaphore is None:
|
||||
self._call_semaphore = _call_semaphore(
|
||||
self._name,
|
||||
(
|
||||
self._config.max_concurrent_calls
|
||||
if self._config is not None
|
||||
else DEFAULT_MAX_CONCURRENT_CALLS
|
||||
),
|
||||
)
|
||||
failure: FailureInfo | None = None
|
||||
for attempt in range(1, _MAX_ATTEMPTS + 1):
|
||||
# Lazy, atomic rebuild: a prior failure disposed the session, so build a
|
||||
# fresh one here. The rebuild lock lets concurrent callers (adopted
|
||||
# sessions dispatched from several agent tasks) share one rebuild rather
|
||||
# than each building their own.
|
||||
if self._server is None:
|
||||
reconnected, reconnect_failure = await self._reconnect()
|
||||
if not reconnected:
|
||||
failure = reconnect_failure or FailureInfo(
|
||||
"transport", reason="reconnect failed"
|
||||
)
|
||||
outcome = await self._handle_failure(failure, attempt, phase="connect")
|
||||
if outcome is not None:
|
||||
return outcome
|
||||
continue
|
||||
assert self._server is not None
|
||||
call_semaphore = self._call_semaphore
|
||||
assert call_semaphore is not None
|
||||
try:
|
||||
async with call_semaphore:
|
||||
result = await job(self._server)
|
||||
# A success clears the quarantine strikes. A connection that
|
||||
# recovered and served a call is healthy again, so transient
|
||||
# failure bursts separated by successful revivals must not
|
||||
# accumulate toward permanent retirement; only sustained failure
|
||||
# with no success in between should retire the connection.
|
||||
self._quarantine_count = 0
|
||||
return _Outcome(value=result)
|
||||
except asyncio.CancelledError:
|
||||
if not self._supervised or self._closing:
|
||||
raise
|
||||
failure = (
|
||||
self._recorder.take() if self._recorder is not None else None
|
||||
) or FailureInfo("transport", reason="session cancelled")
|
||||
# Dispose BEFORE any other await. The transport's cancel scope may be
|
||||
# firing right now; _safe_cleanup exits it so the backoff sleep below
|
||||
# cannot re-raise the cancellation and kill this task. See the
|
||||
# method docstring for why this ordering is load-bearing.
|
||||
await self._safe_cleanup()
|
||||
except _CLASSIFIABLE as exc:
|
||||
failure = classify(exc)
|
||||
if failure.kind == "unknown" and self._recorder is not None:
|
||||
failure = self._recorder.take() or failure
|
||||
# Dispose BEFORE any other await, same reason as the branch above:
|
||||
# never await on a session that has already errored.
|
||||
await self._safe_cleanup()
|
||||
|
||||
# Session is disposed and _server is None; _handle_failure may sleep for
|
||||
# backoff safely, and the next loop iteration rebuilds and retries.
|
||||
outcome = await self._handle_failure(failure, attempt, phase=phase)
|
||||
if outcome is not None:
|
||||
return outcome
|
||||
return _Outcome(dead=True)
|
||||
|
||||
async def _handle_failure(
|
||||
self, failure: FailureInfo, attempt: int, *, phase: _Phase
|
||||
) -> _Outcome | None:
|
||||
self._last_failure = failure
|
||||
if failure.kind == "auth":
|
||||
self._mark_dead(failure, attempt=attempt)
|
||||
return _Outcome(dead=True)
|
||||
if failure.kind == "permission":
|
||||
if phase == "call":
|
||||
return _Outcome(call_failure=failure)
|
||||
self._mark_dead(failure, attempt=attempt)
|
||||
return _Outcome(dead=True)
|
||||
if (
|
||||
phase == "call"
|
||||
and failure.kind == "protocol"
|
||||
and failure.status is not None
|
||||
and 400 <= failure.status <= 499
|
||||
):
|
||||
return _Outcome(call_failure=failure)
|
||||
if attempt == _MAX_ATTEMPTS:
|
||||
await self._quarantine(failure, attempt=attempt)
|
||||
return _Outcome(dead=True)
|
||||
delay = _retry_delay(attempt, failure.retry_after)
|
||||
self._log_retry(failure, attempt, delay)
|
||||
await asyncio.sleep(delay)
|
||||
return None
|
||||
|
||||
def _log_retry(self, failure: FailureInfo, attempt: int, delay: float) -> None:
|
||||
logger.warning(
|
||||
"MCP connection %r retryable failure kind=%s status=%s attempt=%d delay=%.2f",
|
||||
self._name,
|
||||
failure.kind,
|
||||
failure.status,
|
||||
attempt,
|
||||
delay,
|
||||
)
|
||||
|
||||
async def _quarantine(self, failure: FailureInfo, *, attempt: int) -> None:
|
||||
await self._safe_cleanup()
|
||||
self._quarantine_count += 1
|
||||
if self._quarantine_count >= 3:
|
||||
self._mark_dead(failure, attempt=attempt)
|
||||
return
|
||||
cooldown = 30.0 * (2 ** (self._quarantine_count - 1))
|
||||
self._unavailable_until = time.monotonic() + cooldown
|
||||
logger.warning(
|
||||
"MCP connection %r quarantined kind=%s status=%s attempt=%d delay=%.2f",
|
||||
self._name,
|
||||
failure.kind,
|
||||
failure.status,
|
||||
attempt,
|
||||
cooldown,
|
||||
)
|
||||
|
||||
async def _reconnect(self) -> tuple[bool, FailureInfo | None]:
|
||||
"""Build a fresh session under the rebuild lock, so concurrent callers share one.
|
||||
|
||||
Called only when ``_server`` is None (a prior failure already disposed the old
|
||||
session). The lock serializes rebuilds; a caller that finds the session already
|
||||
rebuilt by whoever held the lock first reuses it instead of building a second
|
||||
one. There is deliberately no cleanup of an existing ``_server`` here: this
|
||||
method never runs against a live session, because the failure path disposes
|
||||
before it ever reaches a rebuild.
|
||||
"""
|
||||
async with self._reconnect_lock:
|
||||
if self._server is not None:
|
||||
# Another caller rebuilt while we waited for the lock; share it.
|
||||
return True, None
|
||||
if self._config is None:
|
||||
return False, FailureInfo("transport", reason="no reconnect config")
|
||||
try:
|
||||
server = await self._open()
|
||||
except asyncio.CancelledError:
|
||||
if self._closing:
|
||||
raise
|
||||
self._server = None
|
||||
return False, FailureInfo("transport", reason="reconnect cancelled")
|
||||
except _CLASSIFIABLE as exc:
|
||||
self._server = None
|
||||
failure = classify(exc)
|
||||
if failure.kind == "unknown" and self._recorder is not None:
|
||||
failure = self._recorder.take() or failure
|
||||
return False, failure
|
||||
# connect() is the only readiness surface exposed by the SDK.
|
||||
self._server = server
|
||||
try:
|
||||
await asyncio.sleep(_SETTLE_DELAY)
|
||||
except asyncio.CancelledError:
|
||||
# Dispose the just-built session before returning; _safe_cleanup
|
||||
# re-raises when we are shutting down and absorbs otherwise.
|
||||
await self._safe_cleanup()
|
||||
if self._closing:
|
||||
raise
|
||||
return False, FailureInfo("transport", reason="reconnect cancelled")
|
||||
return True, None
|
||||
|
||||
async def _open(self) -> MCPServer:
|
||||
"""Build and connect the SDK server, reusing the existing setup steps.
|
||||
|
||||
If ``connect()`` fails, the just-built server is cleaned up here on this
|
||||
same task before the error propagates, so a failed connect never orphans
|
||||
an MCP subprocess or half-open HTTP session.
|
||||
"""
|
||||
from strix.tools.mcp.client import _build_server
|
||||
|
||||
if self._config is None:
|
||||
raise RuntimeError(f"MCP connection {self._name!r} has no config to connect")
|
||||
built = _build_server(self._config)
|
||||
server = built.server
|
||||
self._recorder = built.recorder
|
||||
try:
|
||||
await server.connect() # type: ignore[no-untyped-call]
|
||||
except asyncio.CancelledError:
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
raise
|
||||
except _CLASSIFIABLE:
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
raise
|
||||
return server
|
||||
|
||||
# -- helpers --------------------------------------------------------------
|
||||
|
||||
def _call_rejected_message(self, failure: FailureInfo) -> str:
|
||||
if failure.kind == "permission":
|
||||
return (
|
||||
f"MCP connection {self._name!r} rejected this call (status={failure.status}): "
|
||||
"the provider denied this specific request, not the connection. The connection "
|
||||
"is still available. Check the arguments — resource and project identifiers, "
|
||||
"and required fields — and whether the configured credential is allowed to read "
|
||||
"that resource, then retry."
|
||||
)
|
||||
if failure.kind == "protocol":
|
||||
return (
|
||||
f"MCP connection {self._name!r} rejected this call as invalid "
|
||||
f"(status={failure.status}): the request itself was malformed, not the "
|
||||
"connection. The connection is still available. Check the tool's required "
|
||||
"arguments and value formats with describe_mcp, then retry."
|
||||
)
|
||||
raise AssertionError(f"Unexpected call failure kind: {failure.kind}")
|
||||
|
||||
async def _safe_cleanup(self) -> None:
|
||||
"""Dispose the live session on this task, completing teardown even under a
|
||||
firing cancel scope.
|
||||
|
||||
Why this is delicate: the streamable-HTTP transport holds an anyio task group
|
||||
whose cancel scope was entered on this supervising task. When a background POST
|
||||
got a non-2xx status the SDK cancelled that scope, and until the scope is
|
||||
exited every ``await`` on this task re-raises ``CancelledError``.
|
||||
``server.cleanup()`` closes the AsyncExitStack that runs the task group's
|
||||
``__aexit__``, and that ``__aexit__`` is exactly what exits the scope and stops
|
||||
the firing; it also absorbs the scope's own cancellation internally, so the
|
||||
common case returns cleanly. A stray ``CancelledError`` can still surface,
|
||||
though, and ``contextlib.suppress(Exception)`` would let it through because
|
||||
``CancelledError`` is a ``BaseException``, not an ``Exception``.
|
||||
|
||||
So we catch ``CancelledError`` explicitly. During a real shutdown
|
||||
(``_closing``) that cancellation is the run going down and must propagate, so
|
||||
we re-raise it. Otherwise we absorb it and retry the close a bounded number of
|
||||
times: if a cleanup was interrupted before the exit stack finished unwinding,
|
||||
closing again continues from where it left off (the stack pops one callback at
|
||||
a time), so the scope still ends up exited and this task stays runnable for the
|
||||
next rebuild.
|
||||
"""
|
||||
server = self._server
|
||||
self._server = None
|
||||
if server is None:
|
||||
return
|
||||
for _ in range(_MAX_ATTEMPTS):
|
||||
try:
|
||||
# suppress(Exception) absorbs an ordinary cleanup error but lets a
|
||||
# CancelledError through, because it is a BaseException; the outer
|
||||
# handler below is what decides whether to propagate or retry it.
|
||||
with contextlib.suppress(Exception):
|
||||
await server.cleanup() # type: ignore[no-untyped-call]
|
||||
except asyncio.CancelledError:
|
||||
if self._closing:
|
||||
raise
|
||||
# Firing scope hit the cleanup await before the stack finished
|
||||
# unwinding; swallow this cancellation and close again to complete
|
||||
# the teardown. A fully-closed stack makes the retry a clean no-op.
|
||||
continue
|
||||
else:
|
||||
return
|
||||
|
||||
def _report_ready(self, value: bool) -> None:
|
||||
if self._ready is not None and not self._ready.done():
|
||||
self._ready.set_result(value)
|
||||
|
||||
def _fail_pending(self) -> None:
|
||||
for future in self._pending:
|
||||
if not future.done():
|
||||
future.set_result(_Outcome(dead=True))
|
||||
self._pending.clear()
|
||||
|
||||
def _unavailable_message(self) -> str:
|
||||
if self._unavailable_until is not None:
|
||||
remaining = max(0.0, self._unavailable_until - time.monotonic())
|
||||
return (
|
||||
f"MCP connection {self._name!r} is temporarily unavailable "
|
||||
f"(kind={self._last_failure.kind}, status={self._last_failure.status}); "
|
||||
f"retrying in about {remaining:.0f} seconds."
|
||||
)
|
||||
return (
|
||||
f"MCP connection {self._name!r} is unavailable "
|
||||
f"(kind={self._last_failure.kind}, status={self._last_failure.status}); "
|
||||
"it will not be retried."
|
||||
)
|
||||
@@ -12,13 +12,17 @@ import json
|
||||
import logging
|
||||
import re
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agents import RunContextWrapper, function_tool
|
||||
|
||||
from strix.tools.nullish import clean_optional
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from strix.report.state import ReportState
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -257,6 +261,329 @@ def _validate_fix_verification(
|
||||
]
|
||||
|
||||
|
||||
def _finding_class_of(report: dict[str, Any]) -> str:
|
||||
"""Resolve the class of a stored finding.
|
||||
|
||||
A finding filed before ``finding_class`` was persisted still carries the
|
||||
metadata of its class. A record with dependency metadata is a dependency
|
||||
finding even when the field is absent, so read the metadata before falling
|
||||
back to dynamic.
|
||||
"""
|
||||
declared = str(report.get("finding_class") or "").lower()
|
||||
if declared:
|
||||
return declared
|
||||
if report.get("dependency_metadata"):
|
||||
return "dependency_cve"
|
||||
return "dynamic"
|
||||
|
||||
|
||||
_UPDATE_TEXT_FIELDS = (
|
||||
"title",
|
||||
"description",
|
||||
"impact",
|
||||
"target",
|
||||
"technical_analysis",
|
||||
"poc_description",
|
||||
"poc_script_code",
|
||||
"remediation_steps",
|
||||
"evidence",
|
||||
"assumptions",
|
||||
"counterevidence",
|
||||
"confidence_rationale",
|
||||
"severity_change_conditions",
|
||||
"endpoint",
|
||||
"method",
|
||||
"fix_verification",
|
||||
"fix_pr_body",
|
||||
"contextual_cvss_reasoning",
|
||||
)
|
||||
|
||||
|
||||
def _collect_update_changes( # noqa: PLR0912
|
||||
fields: dict[str, Any],
|
||||
) -> tuple[dict[str, Any], list[str]]:
|
||||
"""Validate the fields a revision replaces and return them with any errors."""
|
||||
errors: list[str] = []
|
||||
changes: dict[str, Any] = {}
|
||||
|
||||
for name in _UPDATE_TEXT_FIELDS:
|
||||
value = clean_optional(fields.get(name))
|
||||
if value is not None:
|
||||
changes[name] = value
|
||||
|
||||
confidence = clean_optional(fields.get("confidence"))
|
||||
if confidence is not None:
|
||||
confidence = confidence.lower()
|
||||
if confidence not in _VALID_CONFIDENCE:
|
||||
errors.append(
|
||||
f"Invalid confidence: {confidence!r}. Must be one of: {sorted(_VALID_CONFIDENCE)}"
|
||||
)
|
||||
else:
|
||||
changes["confidence"] = confidence
|
||||
|
||||
fix_effort = clean_optional(fields.get("fix_effort"))
|
||||
if fix_effort is not None:
|
||||
fix_effort = fix_effort.lower()
|
||||
if fix_effort not in _VALID_FIX_EFFORT:
|
||||
errors.append(
|
||||
f"Invalid fix_effort: {fix_effort!r}. Must be one of: {sorted(_VALID_FIX_EFFORT)}"
|
||||
)
|
||||
else:
|
||||
changes["fix_effort"] = fix_effort
|
||||
|
||||
breakdown = fields.get("cvss_breakdown")
|
||||
if breakdown is not None:
|
||||
breakdown_errors = _validate_cvss_breakdown(breakdown)
|
||||
errors.extend(breakdown_errors)
|
||||
if not breakdown_errors:
|
||||
try:
|
||||
cvss_score, severity, _vector = _calculate_cvss(breakdown)
|
||||
except ValueError as exc:
|
||||
errors.append(str(exc))
|
||||
else:
|
||||
# The rating belongs to the vector, so a revised vector carries
|
||||
# its own score and severity rather than leaving the old ones.
|
||||
changes["cvss_breakdown"] = breakdown
|
||||
changes["cvss"] = cvss_score
|
||||
changes["severity"] = severity
|
||||
|
||||
raw_locations = fields.get("code_locations")
|
||||
locations = _normalize_code_locations(raw_locations)
|
||||
if locations:
|
||||
errors.extend(_validate_code_locations(locations))
|
||||
errors.extend(_validate_fix_verification(locations, changes.get("fix_verification")))
|
||||
changes["code_locations"] = locations
|
||||
elif raw_locations:
|
||||
errors.append(
|
||||
"code_locations were dropped as unusable - every location needs a relative "
|
||||
"'file' and an integer 'start_line'"
|
||||
)
|
||||
|
||||
cve, cwe, identifier_errors = _validate_identifiers(
|
||||
clean_optional(fields.get("cve")), clean_optional(fields.get("cwe"))
|
||||
)
|
||||
errors.extend(identifier_errors)
|
||||
if cve:
|
||||
changes["cve"] = cve
|
||||
if cwe:
|
||||
changes["cwe"] = cwe
|
||||
|
||||
return changes, errors
|
||||
|
||||
|
||||
# Evidence that only a dynamic finding carries. A dependency finding describes a
|
||||
# package, not a request against an endpoint.
|
||||
_DYNAMIC_ONLY_UPDATE_FIELDS = (
|
||||
"endpoint",
|
||||
"method",
|
||||
"poc_description",
|
||||
"poc_script_code",
|
||||
)
|
||||
|
||||
# A dependency finding is rated in the context of the codebase that pins it, and
|
||||
# that rating is only shown with the reasoning behind it.
|
||||
_DEPENDENCY_ONLY_UPDATE_FIELDS = ("contextual_cvss_reasoning",)
|
||||
|
||||
|
||||
def _reject_cross_class_revision(
|
||||
report_id: str,
|
||||
matched_class: str,
|
||||
offending: list[str],
|
||||
) -> dict[str, Any]:
|
||||
logger.info(
|
||||
"Revision of %s carries fields (%s) a %s finding does not hold; rejecting",
|
||||
report_id,
|
||||
", ".join(offending),
|
||||
matched_class,
|
||||
)
|
||||
return {
|
||||
"success": False,
|
||||
"error": (
|
||||
f"Report '{report_id}' is a {matched_class} finding, so it cannot carry "
|
||||
f"{', '.join(offending)}. File your proof as its own vulnerability report "
|
||||
"instead of writing it onto this one."
|
||||
),
|
||||
"report_id": report_id,
|
||||
"finding_class": matched_class,
|
||||
"rejected_fields": offending,
|
||||
}
|
||||
|
||||
|
||||
def _rate_dependency_revision(
|
||||
report_id: str,
|
||||
matched: dict[str, Any],
|
||||
changes: dict[str, Any],
|
||||
) -> dict[str, Any] | None:
|
||||
"""Turn a replacement ``cvss_breakdown`` into the contextual rating of a dependency.
|
||||
|
||||
A dependency record keeps its rating as ``cvss``/``severity`` plus the
|
||||
contextual breakdown, vector and reasoning inside ``dependency_metadata``.
|
||||
The package identity in that metadata is copied over untouched. A new
|
||||
breakdown needs its own reasoning. The reasoning alone can be corrected
|
||||
when the record already carries the breakdown it explains.
|
||||
"""
|
||||
breakdown = changes.pop("cvss_breakdown", None)
|
||||
reasoning = changes.pop("contextual_cvss_reasoning", None)
|
||||
if breakdown is None and reasoning is None:
|
||||
return None
|
||||
|
||||
metadata = dict(matched.get("dependency_metadata") or {})
|
||||
if breakdown is None and not metadata.get("contextual_cvss_breakdown"):
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Validation failed",
|
||||
"errors": [
|
||||
"cvss_breakdown is required: this dependency finding carries no "
|
||||
"contextual rating yet, so contextual_cvss_reasoning has nothing to explain"
|
||||
],
|
||||
"report_id": report_id,
|
||||
}
|
||||
if reasoning is None:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Validation failed",
|
||||
"errors": [
|
||||
"contextual_cvss_reasoning is required: a dependency finding is re-rated "
|
||||
"with the cvss_breakdown observed in this codebase together with the "
|
||||
"reasoning a reader can check"
|
||||
],
|
||||
"report_id": report_id,
|
||||
}
|
||||
|
||||
if breakdown is not None:
|
||||
score, _severity, vector = _calculate_cvss(breakdown)
|
||||
metadata["contextual_cvss_breakdown"] = breakdown
|
||||
metadata["contextual_cvss_score"] = score
|
||||
metadata["contextual_cvss_vector"] = vector
|
||||
metadata["contextual_cvss_reasoning"] = reasoning[:_MAX_CONTEXTUAL_REASONING_CHARS]
|
||||
changes["dependency_metadata"] = metadata
|
||||
return None
|
||||
|
||||
|
||||
def _fit_revision_to_class(
|
||||
report_state: ReportState,
|
||||
report_id: str,
|
||||
changes: dict[str, Any],
|
||||
) -> dict[str, Any] | None:
|
||||
"""Keep a revision inside the class of the finding it names.
|
||||
|
||||
A finding keeps its class and the metadata that belongs to it. Writing an
|
||||
exploit onto a dependency record would leave it carrying a package pin next
|
||||
to a request against an endpoint, so the proof belongs in its own dynamic
|
||||
finding instead. A dependency finding is still re-rated, through the
|
||||
contextual CVSS it was filed with.
|
||||
"""
|
||||
matched = next(
|
||||
(r for r in report_state.get_existing_vulnerabilities() if r.get("id") == report_id),
|
||||
None,
|
||||
)
|
||||
if matched is None:
|
||||
return None
|
||||
|
||||
matched_class = _finding_class_of(matched)
|
||||
foreign = (
|
||||
_DEPENDENCY_ONLY_UPDATE_FIELDS
|
||||
if matched_class == "dynamic"
|
||||
else _DYNAMIC_ONLY_UPDATE_FIELDS
|
||||
)
|
||||
offending = [name for name in foreign if name in changes]
|
||||
if offending:
|
||||
return _reject_cross_class_revision(report_id, matched_class, offending)
|
||||
if matched_class == "dynamic":
|
||||
return None
|
||||
return _rate_dependency_revision(report_id, matched, changes)
|
||||
|
||||
|
||||
def _read_revision(
|
||||
report_id: str, update_reason: str, fields: dict[str, Any]
|
||||
) -> tuple[dict[str, Any], dict[str, Any] | None]:
|
||||
"""Return the changes a revision asks for, or the reason it cannot be acted on."""
|
||||
if not report_id or not str(update_reason or "").strip():
|
||||
missing = "report_id" if not report_id else "update_reason"
|
||||
return {}, {
|
||||
"success": False,
|
||||
"error": (
|
||||
f"{missing} cannot be empty - name the report you are revising and state "
|
||||
"what you learned that it does not yet carry"
|
||||
),
|
||||
}
|
||||
|
||||
changes, errors = _collect_update_changes(fields)
|
||||
if errors:
|
||||
return {}, {"success": False, "error": "Validation failed", "errors": errors}
|
||||
if not changes:
|
||||
return {}, {
|
||||
"success": False,
|
||||
"error": "No fields to update - pass at least one field you want to replace",
|
||||
}
|
||||
return changes, None
|
||||
|
||||
|
||||
def _do_update(
|
||||
*,
|
||||
report_id: str,
|
||||
update_reason: str,
|
||||
fields: dict[str, Any],
|
||||
agent_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Apply an agent's own revision to a report it can name.
|
||||
|
||||
Editing a finding is its own operation and the only way a filed finding
|
||||
changes. Deduplication never reaches this path: it only decides whether a
|
||||
new candidate is a finding already on file.
|
||||
"""
|
||||
report_id = (report_id or "").strip()
|
||||
changes, rejection = _read_revision(report_id, update_reason, fields)
|
||||
if rejection is not None:
|
||||
return rejection
|
||||
|
||||
from strix.report.state import get_global_report_state
|
||||
|
||||
report_state = get_global_report_state()
|
||||
if report_state is None:
|
||||
return {
|
||||
"success": False,
|
||||
"error": "Report state unavailable - no reports have been filed yet",
|
||||
}
|
||||
|
||||
class_error = _fit_revision_to_class(report_state, report_id, changes)
|
||||
if class_error is not None:
|
||||
return class_error
|
||||
|
||||
updated = report_state.update_vulnerability_report(
|
||||
report_id,
|
||||
changes,
|
||||
update_reason=update_reason,
|
||||
updated_by_agent_id=agent_id,
|
||||
updated_by_agent_name=agent_name,
|
||||
)
|
||||
if updated is None:
|
||||
known = [r.get("id") for r in report_state.get_existing_vulnerabilities()]
|
||||
if report_id not in known:
|
||||
error = f"Report with id '{report_id}' not found"
|
||||
else:
|
||||
error = f"Report '{report_id}' already says this - nothing in your update changes it"
|
||||
return {"success": False, "error": error, "report_id": report_id}
|
||||
|
||||
logger.info(
|
||||
"Vulnerability report %s revised by its author: severity=%s cvss=%s fields=%s",
|
||||
report_id,
|
||||
updated.get("severity"),
|
||||
updated.get("cvss"),
|
||||
", ".join(sorted(changes)),
|
||||
)
|
||||
return {
|
||||
"success": True,
|
||||
"action": "updated",
|
||||
"message": f"Report '{report_id}' now carries your revision. Do not file it again.",
|
||||
"report_id": report_id,
|
||||
"updated_fields": sorted(changes),
|
||||
"severity": updated.get("severity"),
|
||||
"cvss_score": updated.get("cvss"),
|
||||
}
|
||||
|
||||
|
||||
async def _do_create(
|
||||
*,
|
||||
title: str,
|
||||
@@ -359,9 +686,37 @@ async def _do_create(
|
||||
"endpoint": endpoint,
|
||||
"method": method,
|
||||
}
|
||||
report_fields: dict[str, Any] = {
|
||||
"title": title,
|
||||
"description": description,
|
||||
"severity": severity,
|
||||
"impact": impact,
|
||||
"target": target,
|
||||
"technical_analysis": technical_analysis,
|
||||
"poc_description": poc_description,
|
||||
"poc_script_code": poc_script_code,
|
||||
"remediation_steps": remediation_steps,
|
||||
"evidence": evidence,
|
||||
"assumptions": assumptions,
|
||||
"counterevidence": counterevidence,
|
||||
"confidence": confidence,
|
||||
"confidence_rationale": confidence_rationale,
|
||||
"severity_change_conditions": severity_change_conditions,
|
||||
"fix_effort": fix_effort,
|
||||
"cvss": cvss_score,
|
||||
"cvss_breakdown": cvss_breakdown,
|
||||
"endpoint": endpoint,
|
||||
"method": method,
|
||||
"cve": cve,
|
||||
"cwe": cwe,
|
||||
"code_locations": parsed_locations,
|
||||
"fix_verification": fix_verification,
|
||||
"fix_pr_body": fix_pr_body,
|
||||
}
|
||||
|
||||
dedupe = await check_duplicate(candidate, existing)
|
||||
if dedupe.get("is_duplicate"):
|
||||
duplicate_id = dedupe.get("duplicate_id", "")
|
||||
duplicate_id = str(dedupe.get("duplicate_id") or "")
|
||||
duplicate_title = next(
|
||||
(r.get("title", "Unknown") for r in existing if r.get("id") == duplicate_id),
|
||||
"",
|
||||
@@ -379,31 +734,7 @@ async def _do_create(
|
||||
}
|
||||
|
||||
report_id = report_state.add_vulnerability_report(
|
||||
title=title,
|
||||
description=description,
|
||||
severity=severity,
|
||||
impact=impact,
|
||||
target=target,
|
||||
technical_analysis=technical_analysis,
|
||||
poc_description=poc_description,
|
||||
poc_script_code=poc_script_code,
|
||||
remediation_steps=remediation_steps,
|
||||
evidence=evidence,
|
||||
assumptions=assumptions,
|
||||
counterevidence=counterevidence,
|
||||
confidence=confidence,
|
||||
confidence_rationale=confidence_rationale,
|
||||
severity_change_conditions=severity_change_conditions,
|
||||
fix_effort=fix_effort,
|
||||
cvss=cvss_score,
|
||||
cvss_breakdown=cvss_breakdown,
|
||||
endpoint=endpoint,
|
||||
method=method,
|
||||
cve=cve,
|
||||
cwe=cwe,
|
||||
code_locations=parsed_locations,
|
||||
fix_verification=fix_verification,
|
||||
fix_pr_body=fix_pr_body,
|
||||
**report_fields,
|
||||
agent_id=agent_id if isinstance(agent_id, str) else None,
|
||||
agent_name=agent_name if isinstance(agent_name, str) else None,
|
||||
)
|
||||
@@ -512,7 +843,9 @@ async def create_vulnerability_report(
|
||||
Automatic LLM-based **deduplication** rejects reports that describe
|
||||
the same root cause on the same asset as an existing report. If you
|
||||
get a ``duplicate_of`` response, do NOT retry — move on to other
|
||||
areas.
|
||||
areas. When you have learned something a filed finding does not yet
|
||||
carry, revise that finding with ``update_vulnerability_report``
|
||||
instead of filing this report again.
|
||||
|
||||
**Counterevidence pass (required before filing)**: actively build the
|
||||
strongest case that this finding is NOT exploitable, or less severe
|
||||
@@ -886,6 +1219,147 @@ async def create_vulnerability_report(
|
||||
return json.dumps(result, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
@function_tool(timeout=60, strict_mode=False)
|
||||
async def update_vulnerability_report(
|
||||
ctx: RunContextWrapper,
|
||||
report_id: str,
|
||||
update_reason: str,
|
||||
title: str | None = None,
|
||||
description: str | None = None,
|
||||
impact: str | None = None,
|
||||
target: str | None = None,
|
||||
technical_analysis: str | None = None,
|
||||
poc_description: str | None = None,
|
||||
poc_script_code: str | None = None,
|
||||
remediation_steps: str | None = None,
|
||||
evidence: str | None = None,
|
||||
assumptions: str | None = None,
|
||||
counterevidence: str | None = None,
|
||||
confidence: str | None = None,
|
||||
confidence_rationale: str | None = None,
|
||||
severity_change_conditions: str | None = None,
|
||||
fix_effort: str | None = None,
|
||||
cvss_breakdown: dict[str, str] | None = None,
|
||||
endpoint: str | None = None,
|
||||
method: str | None = None,
|
||||
cve: str | None = None,
|
||||
cwe: str | None = None,
|
||||
code_locations: list[dict[str, Any]] | None = None,
|
||||
fix_verification: str | None = None,
|
||||
fix_pr_body: str | None = None,
|
||||
contextual_cvss_reasoning: str | None = None,
|
||||
) -> str:
|
||||
"""Revise a vulnerability report that is already filed, keeping its id.
|
||||
|
||||
Use this when you learn something a filed finding does not yet carry:
|
||||
|
||||
- You built the working exploit after filing the finding on static
|
||||
evidence, so the PoC and the confidence change.
|
||||
- You chained the finding with another one and the real impact is
|
||||
higher, so the impact narrative and the CVSS vector change.
|
||||
- Further testing narrowed or weakened the finding, so the severity
|
||||
must come down.
|
||||
- Counterevidence, remediation, or a code location was wrong or
|
||||
incomplete.
|
||||
|
||||
This is not deduplication. You do not need a duplicate verdict to
|
||||
revise your own finding, and you must not file a second report for a
|
||||
finding you can revise. Call ``list_reports`` or ``get_report`` first
|
||||
to find the id and read what the report already says.
|
||||
|
||||
Pass only the fields you want to replace. Every other field stays as
|
||||
it is. Reporting rules of ``create_vulnerability_report`` apply to
|
||||
every field you pass, including the markdown and tone rules.
|
||||
|
||||
Notes on specific fields:
|
||||
|
||||
- ``cvss_breakdown`` replaces the whole vector. The score and the
|
||||
severity are recalculated from it, so pass all 8 metrics. On a
|
||||
dependency finding it replaces the contextual rating and needs
|
||||
``contextual_cvss_reasoning`` with it. Pass the reasoning alone to
|
||||
correct only the explanation of the rating already on file.
|
||||
- A dependency finding never carries ``endpoint``, ``method`` or a PoC.
|
||||
File a proven exploit of the package as its own report.
|
||||
- A field that only explains another field is dropped when the field
|
||||
it explains changes and you pass no replacement. Pass
|
||||
``confidence_rationale`` with a new ``confidence``, and
|
||||
``severity_change_conditions`` with a new ``cvss_breakdown``.
|
||||
- ``code_locations`` replaces the whole list. A location carrying
|
||||
``fix_after`` needs ``fix_verification``.
|
||||
|
||||
The report keeps its id, its original author, and its filing time. The
|
||||
revision is recorded in the report as update history, so state the
|
||||
reason plainly.
|
||||
|
||||
Args:
|
||||
report_id: Id of the report to revise (format ``vuln-NNNN``).
|
||||
update_reason: What you learned that the report does not yet
|
||||
carry, in one or two sentences.
|
||||
title: Replacement title.
|
||||
description: Replacement overview.
|
||||
impact: Replacement impact narrative.
|
||||
target: Replacement affected asset.
|
||||
technical_analysis: Replacement technical details.
|
||||
poc_description: Replacement PoC steps (no code).
|
||||
poc_script_code: Replacement exploit script or payload.
|
||||
remediation_steps: Replacement remediation prose (no code).
|
||||
evidence: Replacement evidence.
|
||||
assumptions: Replacement exploitability prerequisites.
|
||||
counterevidence: Replacement case against the finding.
|
||||
confidence: ``high`` / ``medium`` / ``low``.
|
||||
confidence_rationale: The gap behind a confidence below ``high``.
|
||||
severity_change_conditions: What would move the severity now.
|
||||
fix_effort: ``trivial`` / ``low`` / ``medium`` / ``high``.
|
||||
cvss_breakdown: All 8 CVSS metrics. Replaces the score and the
|
||||
severity too.
|
||||
endpoint: Replacement endpoint.
|
||||
method: Replacement HTTP method.
|
||||
cve: Replacement CVE id.
|
||||
cwe: Replacement CWE id.
|
||||
code_locations: Replacement code locations.
|
||||
fix_verification: Verification statement for an applyable fix.
|
||||
fix_pr_body: Replacement fix PR body.
|
||||
contextual_cvss_reasoning: Dependency findings only. What you
|
||||
observed in this codebase that justifies the contextual
|
||||
``cvss_breakdown``.
|
||||
"""
|
||||
agent_id, agent_name = _caller_identity(ctx)
|
||||
result = await asyncio.to_thread(
|
||||
_do_update,
|
||||
report_id=report_id,
|
||||
update_reason=update_reason,
|
||||
fields={
|
||||
"title": title,
|
||||
"description": description,
|
||||
"impact": impact,
|
||||
"target": target,
|
||||
"technical_analysis": technical_analysis,
|
||||
"poc_description": poc_description,
|
||||
"poc_script_code": poc_script_code,
|
||||
"remediation_steps": remediation_steps,
|
||||
"evidence": evidence,
|
||||
"assumptions": assumptions,
|
||||
"counterevidence": counterevidence,
|
||||
"confidence": confidence,
|
||||
"confidence_rationale": confidence_rationale,
|
||||
"severity_change_conditions": severity_change_conditions,
|
||||
"fix_effort": fix_effort,
|
||||
"cvss_breakdown": cvss_breakdown,
|
||||
"endpoint": endpoint,
|
||||
"method": method,
|
||||
"cve": cve,
|
||||
"cwe": cwe,
|
||||
"code_locations": code_locations,
|
||||
"fix_verification": fix_verification,
|
||||
"fix_pr_body": fix_pr_body,
|
||||
"contextual_cvss_reasoning": contextual_cvss_reasoning,
|
||||
},
|
||||
agent_id=agent_id,
|
||||
agent_name=agent_name,
|
||||
)
|
||||
return json.dumps(result, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
_DEP_SEVERITY_FROM_CVSS = {
|
||||
(9.0, 10.0): "critical",
|
||||
(7.0, 9.0): "high",
|
||||
|
||||
@@ -1,28 +1,29 @@
|
||||
"""Target-scoped threat models — cached under ``~/.strix/threat-models``.
|
||||
"""Run-scoped threat models — mirrored to ``{state_dir}/threat_models.json``.
|
||||
|
||||
A threat model describes the target, not the scan: a host, an application, an
|
||||
API, a repository, or whatever else the engagement is pointed at. It stays
|
||||
valid across unrelated runs against the same target, so it is keyed by target
|
||||
identity rather than by run id — one agent derives it, every later agent in
|
||||
this run and in future runs against the same target reads it back instead of
|
||||
A threat model is the scan's shared answer to who the attacker is, where the
|
||||
trust boundaries sit, and what counts as critical for the target. One agent
|
||||
derives it and every other agent on the same run reads it back instead of
|
||||
re-deriving trust boundaries from scratch.
|
||||
|
||||
Where the target is a checkout, the model is additionally pinned to the git
|
||||
revision, so a moved ``HEAD`` marks it stale. Black-box targets have no
|
||||
revision to pin to; those age out instead.
|
||||
It does not outlive the scan. The mirror lives in the run's own state directory
|
||||
and exists only so a resumed scan keeps the baseline its earlier agents agreed
|
||||
on; a new scan against the same host or checkout starts with no model and
|
||||
derives its own. Agents do spell one target several ways within a run — the URL
|
||||
they were handed, the page they happen to be testing, a checkout path — so a
|
||||
model is keyed by a normalized target identity to keep them converging on one
|
||||
document instead of each starting a fresh one.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import subprocess
|
||||
import tempfile
|
||||
import threading
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit
|
||||
@@ -35,16 +36,19 @@ from strix.core.agents import AgentCoordinator
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_CACHE_DIR = Path.home() / ".strix" / "threat-models"
|
||||
_MAX_MODEL_BYTES = 512 * 1024
|
||||
_MIN_MODEL_CHARS = 400
|
||||
_MIN_AMENDMENT_CHARS = 80
|
||||
_MAX_AMENDMENTS = 40
|
||||
_GIT_TIMEOUT_SECONDS = 10
|
||||
_UNVERSIONED = "unversioned"
|
||||
_MAX_AGE_DAYS = 14
|
||||
_DEFAULT_PORTS = {"http": "80", "https": "443"}
|
||||
_cache_lock = threading.RLock()
|
||||
|
||||
_store_lock = threading.RLock()
|
||||
|
||||
# The whole store: target identity -> model. It holds exactly the models this
|
||||
# scan derived, and is mirrored to the run's state directory for resume.
|
||||
_MODELS: dict[str, dict[str, Any]] = {}
|
||||
_store_path: Path | None = None
|
||||
|
||||
_REQUIRED_SECTIONS = (
|
||||
"overview",
|
||||
@@ -95,7 +99,7 @@ def _remote_authority(target: str) -> str:
|
||||
|
||||
|
||||
def _normalize_remote_target(target: str) -> str:
|
||||
"""Collapse the spellings of one remote target onto a single cache key."""
|
||||
"""Collapse the spellings of one remote target onto a single key."""
|
||||
authority = _remote_authority(target)
|
||||
if not authority:
|
||||
return re.sub(r"\s+", " ", target.lower()).strip()
|
||||
@@ -132,30 +136,23 @@ def _normalize_git_remote(remote: str) -> str:
|
||||
return normalized.removesuffix(".git")
|
||||
|
||||
|
||||
def _target_identity(target: str) -> tuple[str, str]:
|
||||
"""Return the (stable identity, revision) pair a cached model is keyed on.
|
||||
def _target_identity(target: str) -> str:
|
||||
"""Return the stable identity a model is stored under.
|
||||
|
||||
A checkout is keyed on its remote (so the same repository cloned to two
|
||||
paths shares one model, and a subdirectory resolves to the whole tree) and
|
||||
pinned to ``HEAD``. Everything else — a host, a URL, an API base, a named
|
||||
scope — is keyed on its normalized form and carries no revision. Both
|
||||
routes run through the same normalization, so a checkout and the URL it
|
||||
was cloned from land on one key.
|
||||
A checkout is keyed on its remote, so the same repository checked out at
|
||||
two paths shares one model and a subdirectory resolves to the whole tree.
|
||||
Everything else — a host, a URL, an API base, a named scope — is keyed on
|
||||
its normalized form. Both routes run through the same normalization, so a
|
||||
checkout and the URL it was cloned from land on one key.
|
||||
"""
|
||||
directory = _local_directory(target)
|
||||
if directory is None:
|
||||
return _normalize_remote_target(target).removesuffix(".git"), _UNVERSIONED
|
||||
return _normalize_remote_target(target).removesuffix(".git")
|
||||
remote = _git(directory, ["config", "--get", "remote.origin.url"])
|
||||
revision = _git(directory, ["rev-parse", "HEAD"]) or _UNVERSIONED
|
||||
if remote:
|
||||
return _normalize_git_remote(remote), revision
|
||||
return _normalize_git_remote(remote)
|
||||
toplevel = _git(directory, ["rev-parse", "--show-toplevel"])
|
||||
return toplevel or str(directory), revision
|
||||
|
||||
|
||||
def _cache_path(identity: str) -> Path:
|
||||
digest = hashlib.sha256(identity.encode("utf-8")).hexdigest()[:16]
|
||||
return _CACHE_DIR / f"{digest}.json"
|
||||
return toplevel or str(directory)
|
||||
|
||||
|
||||
def _snap_to_scan_target(raw: str, scan_targets: list[str]) -> str:
|
||||
@@ -163,13 +160,13 @@ def _snap_to_scan_target(raw: str, scan_targets: list[str]) -> str:
|
||||
|
||||
Agents name the same target differently — one passes the URL it was given,
|
||||
the next the page it happens to be testing, a third the checkout path. Left
|
||||
alone those become separate cache keys, every lookup misses, and each agent
|
||||
alone those become separate keys, every lookup misses, and each agent
|
||||
quietly derives its own model, which is the exact failure the shared model
|
||||
exists to prevent. So a target that is recognisably one of the scan's own
|
||||
targets is resolved to that target instead.
|
||||
"""
|
||||
identity, _ = _target_identity(raw)
|
||||
scoped = [(target, _target_identity(target)[0]) for target in scan_targets]
|
||||
identity = _target_identity(raw)
|
||||
scoped = [(target, _target_identity(target)) for target in scan_targets]
|
||||
if any(known == identity for _, known in scoped):
|
||||
return raw
|
||||
|
||||
@@ -208,38 +205,51 @@ def _resolve_target(
|
||||
return (_snap_to_scan_target(raw, known) if known else raw), None
|
||||
|
||||
|
||||
def _is_expired(created_at: str | None) -> bool:
|
||||
if not created_at:
|
||||
return True
|
||||
try:
|
||||
created = datetime.fromisoformat(created_at)
|
||||
except ValueError:
|
||||
return True
|
||||
if created.tzinfo is None:
|
||||
created = created.replace(tzinfo=UTC)
|
||||
return datetime.now(UTC) - created > timedelta(days=_MAX_AGE_DAYS)
|
||||
|
||||
|
||||
def _missing_sections(content: str) -> list[str]:
|
||||
lowered = content.lower()
|
||||
return [section for section in _REQUIRED_SECTIONS if section not in lowered]
|
||||
|
||||
|
||||
def _read_cache(path: Path) -> dict[str, Any] | None:
|
||||
"""Load a cached model. Callers must already hold ``_cache_lock``."""
|
||||
if not path.is_file():
|
||||
return None
|
||||
try:
|
||||
cached = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
logger.exception("threat model cache at %s is unreadable", path)
|
||||
return None
|
||||
return cached if isinstance(cached, dict) else None
|
||||
|
||||
|
||||
def _write_cache(path: Path, payload: dict[str, Any]) -> str | None:
|
||||
"""Atomically persist a model. Callers must already hold ``_cache_lock``."""
|
||||
def hydrate_threat_models_from_disk(state_dir: Path) -> None:
|
||||
"""Point the store at this run's mirror and load whatever it already holds.
|
||||
|
||||
A resumed scan is the same scan, so its agents have to keep the baseline
|
||||
the earlier ones agreed on. The mirror lives under the run directory, so a
|
||||
different scan never reads it.
|
||||
"""
|
||||
global _store_path # noqa: PLW0603
|
||||
_store_path = state_dir / "threat_models.json"
|
||||
with _store_lock:
|
||||
_MODELS.clear()
|
||||
if not _store_path.is_file():
|
||||
return
|
||||
try:
|
||||
data = json.loads(_store_path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
logger.exception(
|
||||
"threat_models.json at %s is unreadable; starting with no models",
|
||||
_store_path,
|
||||
)
|
||||
return
|
||||
if not isinstance(data, dict):
|
||||
return
|
||||
_MODELS.update(
|
||||
{
|
||||
identity: model
|
||||
for identity, model in data.items()
|
||||
if isinstance(identity, str) and isinstance(model, dict)
|
||||
}
|
||||
)
|
||||
logger.info("threat models hydrated from %s (%d)", _store_path, len(_MODELS))
|
||||
|
||||
|
||||
def _persist_locked() -> None:
|
||||
"""Mirror the store to disk. Callers must already hold ``_store_lock``.
|
||||
|
||||
Serializing and renaming in one critical section keeps a writer holding an
|
||||
older serialization from winning the rename and dropping a concurrent
|
||||
agent's model or amendment.
|
||||
"""
|
||||
path = _store_path
|
||||
if path is None:
|
||||
return
|
||||
try:
|
||||
payload = json.dumps(_MODELS, ensure_ascii=False, default=str)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with tempfile.NamedTemporaryFile(
|
||||
mode="w",
|
||||
@@ -249,85 +259,61 @@ def _write_cache(path: Path, payload: dict[str, Any]) -> str | None:
|
||||
suffix=".tmp",
|
||||
delete=False,
|
||||
) as tmp:
|
||||
tmp.write(json.dumps(payload, ensure_ascii=False))
|
||||
tmp.write(payload)
|
||||
tmp_path = Path(tmp.name)
|
||||
tmp_path.replace(path)
|
||||
except OSError as exc:
|
||||
logger.exception("threat model persist to %s failed", path)
|
||||
return f"Failed to persist threat model: {exc}"
|
||||
return None
|
||||
except OSError:
|
||||
logger.exception("threat model mirror to %s failed", path)
|
||||
|
||||
|
||||
def _amendments_of(cached: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
raw = cached.get("amendments")
|
||||
def _missing_sections(content: str) -> list[str]:
|
||||
lowered = content.lower()
|
||||
return [section for section in _REQUIRED_SECTIONS if section not in lowered]
|
||||
|
||||
|
||||
def _amendments_of(model: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
raw = model.get("amendments")
|
||||
if not isinstance(raw, list):
|
||||
return []
|
||||
return [item for item in raw if isinstance(item, dict)]
|
||||
|
||||
|
||||
def _not_found(identity: str, revision: str) -> dict[str, Any]:
|
||||
def _not_found(identity: str) -> dict[str, Any]:
|
||||
return {
|
||||
"success": True,
|
||||
"found": False,
|
||||
"target": identity,
|
||||
"revision": revision,
|
||||
"message": (
|
||||
"No threat model cached for this target. Derive one — from the code if "
|
||||
"you have it, from recon output if you do not — and persist it with "
|
||||
"save_threat_model, so every agent on this scan shares one view of the "
|
||||
"trust boundaries instead of each inventing their own."
|
||||
"No threat model for this target on this scan. Nothing carries over "
|
||||
"from other scans, so derive one — from the code if you have it, from "
|
||||
"recon output if you do not — and share it with save_threat_model, so "
|
||||
"every agent on this scan works from one view of the trust boundaries "
|
||||
"instead of each inventing their own."
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _staleness(cached: dict[str, Any], revision: str) -> tuple[bool, str | None]:
|
||||
"""Decide whether a cached model can still be trusted, and why not."""
|
||||
if revision != _UNVERSIONED:
|
||||
if cached.get("revision") == revision:
|
||||
return False, None
|
||||
return True, (
|
||||
"This model was derived against a different revision. Use it as a "
|
||||
"starting point, re-check the boundaries it names against the current "
|
||||
"tree, and save the corrected version."
|
||||
)
|
||||
created_at = cached.get("created_at")
|
||||
if not _is_expired(created_at if isinstance(created_at, str) else None):
|
||||
return False, None
|
||||
return True, (
|
||||
f"This model is more than {_MAX_AGE_DAYS} days old and there is no revision "
|
||||
"to pin it to, so the target may have moved under it. Treat its surface "
|
||||
"inventory as a lead list to re-confirm during recon, not as fact, and save "
|
||||
"the corrected version."
|
||||
)
|
||||
|
||||
|
||||
def _get_impl(target: str, scan_targets: list[str] | None = None) -> dict[str, Any]:
|
||||
resolved, error = _resolve_target(target, scan_targets)
|
||||
if resolved is None:
|
||||
return {"success": False, "error": error}
|
||||
|
||||
identity, revision = _target_identity(resolved)
|
||||
path = _cache_path(identity)
|
||||
with _cache_lock:
|
||||
cached = _read_cache(path)
|
||||
if cached is None:
|
||||
return _not_found(identity, revision)
|
||||
content = cached.get("content")
|
||||
identity = _target_identity(resolved)
|
||||
with _store_lock:
|
||||
model = _MODELS.get(identity)
|
||||
if model is None:
|
||||
return _not_found(identity)
|
||||
content = model.get("content")
|
||||
amendments = list(_amendments_of(model))
|
||||
if not isinstance(content, str) or not content.strip():
|
||||
return _not_found(identity, revision)
|
||||
return _not_found(identity)
|
||||
|
||||
stale, stale_message = _staleness(cached, revision)
|
||||
result: dict[str, Any] = {
|
||||
"success": True,
|
||||
"found": True,
|
||||
"target": identity,
|
||||
"revision": revision,
|
||||
"cached_revision": cached.get("revision"),
|
||||
"created_at": cached.get("created_at"),
|
||||
"stale": stale,
|
||||
"content": content,
|
||||
}
|
||||
amendments = _amendments_of(cached)
|
||||
if amendments:
|
||||
result["amendments"] = amendments
|
||||
result["amendments_note"] = (
|
||||
@@ -335,8 +321,6 @@ def _get_impl(target: str, scan_targets: list[str] | None = None) -> dict[str, A
|
||||
"correct or extend it and have not been folded in yet - read them as "
|
||||
"part of the model, and prefer the later one where they conflict."
|
||||
)
|
||||
if stale_message:
|
||||
result["message"] = stale_message
|
||||
return result
|
||||
|
||||
|
||||
@@ -376,25 +360,21 @@ def _save_impl(
|
||||
),
|
||||
}
|
||||
|
||||
identity, revision = _target_identity(resolved)
|
||||
path = _cache_path(identity)
|
||||
payload: dict[str, Any] = {
|
||||
"target": identity,
|
||||
"revision": revision,
|
||||
"created_at": datetime.now(UTC).isoformat(),
|
||||
"created_by": agent_name,
|
||||
"content": body,
|
||||
}
|
||||
with _cache_lock:
|
||||
existing = _read_cache(path)
|
||||
identity = _target_identity(resolved)
|
||||
with _store_lock:
|
||||
existing = _MODELS.get(identity)
|
||||
folded = len(_amendments_of(existing)) if existing else 0
|
||||
error = _write_cache(path, payload)
|
||||
if error:
|
||||
return {"success": False, "error": error}
|
||||
_MODELS[identity] = {
|
||||
"target": identity,
|
||||
"written_at": datetime.now(UTC).isoformat(),
|
||||
"written_by": agent_name,
|
||||
"content": body,
|
||||
}
|
||||
_persist_locked()
|
||||
|
||||
message = (
|
||||
"Threat model saved. Subagents should call get_threat_model before they "
|
||||
"start, and treat its trust boundaries as the shared baseline."
|
||||
"Threat model shared with this scan. Subagents should call get_threat_model "
|
||||
"before they start, and treat its trust boundaries as the shared baseline."
|
||||
)
|
||||
if folded:
|
||||
message += (
|
||||
@@ -404,34 +384,35 @@ def _save_impl(
|
||||
return {
|
||||
"success": True,
|
||||
"target": identity,
|
||||
"revision": revision,
|
||||
"amendments_cleared": folded,
|
||||
"message": message,
|
||||
}
|
||||
|
||||
|
||||
def _append_amendment(
|
||||
path: Path, amendment: dict[str, Any]
|
||||
identity: str, amendment: dict[str, Any]
|
||||
) -> tuple[list[dict[str, Any]] | None, str | None]:
|
||||
"""Add an amendment to the cached model. Returns (amendments, error)."""
|
||||
with _cache_lock:
|
||||
cached = _read_cache(path)
|
||||
if cached is None or not str(cached.get("content", "")).strip():
|
||||
"""Add an amendment to the stored model. Returns (amendments, error)."""
|
||||
with _store_lock:
|
||||
model = _MODELS.get(identity)
|
||||
if model is None or not str(model.get("content", "")).strip():
|
||||
return None, (
|
||||
"No threat model exists for this target yet, so there is nothing to "
|
||||
"amend. Derive the base model and call save_threat_model instead."
|
||||
)
|
||||
amendments = _amendments_of(cached)
|
||||
amendments = _amendments_of(model)
|
||||
if len(amendments) >= _MAX_AMENDMENTS:
|
||||
return None, (
|
||||
f"This model already carries {len(amendments)} amendments. Fold them "
|
||||
"into the base model with save_threat_model before adding more."
|
||||
)
|
||||
amendments.append(amendment)
|
||||
cached["amendments"] = amendments
|
||||
if len(json.dumps(cached, ensure_ascii=False).encode("utf-8")) > _MAX_MODEL_BYTES:
|
||||
candidate = [*amendments, amendment]
|
||||
sized = {**model, "amendments": candidate}
|
||||
if len(json.dumps(sized, ensure_ascii=False).encode("utf-8")) > _MAX_MODEL_BYTES:
|
||||
return None, "Threat model with this amendment exceeds 512KB; tighten it."
|
||||
return amendments, _write_cache(path, cached)
|
||||
model["amendments"] = candidate
|
||||
_persist_locked()
|
||||
return candidate, None
|
||||
|
||||
|
||||
def _amend_impl(
|
||||
@@ -455,13 +436,12 @@ def _amend_impl(
|
||||
),
|
||||
}
|
||||
|
||||
identity, revision = _target_identity(resolved)
|
||||
identity = _target_identity(resolved)
|
||||
amendments, amend_error = _append_amendment(
|
||||
_cache_path(identity),
|
||||
identity,
|
||||
{
|
||||
"at": datetime.now(UTC).isoformat(),
|
||||
"by": agent_name,
|
||||
"revision": revision,
|
||||
"content": body,
|
||||
},
|
||||
)
|
||||
@@ -471,7 +451,6 @@ def _amend_impl(
|
||||
return {
|
||||
"success": True,
|
||||
"target": identity,
|
||||
"revision": revision,
|
||||
"amendment_count": len(amendments),
|
||||
"message": (
|
||||
"Amendment recorded. Agents calling get_threat_model will now see it "
|
||||
@@ -500,25 +479,25 @@ def _scan_targets(ctx: RunContextWrapper) -> list[str]:
|
||||
|
||||
@function_tool(timeout=30)
|
||||
async def get_threat_model(ctx: RunContextWrapper, target: str) -> str:
|
||||
"""Read the cached threat model for a target, if one exists.
|
||||
"""Read this scan's threat model for a target, if an agent has derived one.
|
||||
|
||||
A threat model belongs to the target, not to this scan — the same
|
||||
trust boundaries hold across unrelated runs against the same host
|
||||
or application. Call this before you start hunting so you inherit
|
||||
the shared view instead of re-deriving it, and so every agent on
|
||||
this run agrees on what "attacker-controlled" means here.
|
||||
The threat model is this run's shared answer to who the attacker
|
||||
is, where the trust boundaries sit, and what counts as critical
|
||||
here. Call it before you start hunting so you inherit the shared
|
||||
view instead of re-deriving it, and so every agent on this run
|
||||
agrees on what "attacker-controlled" means.
|
||||
|
||||
It is scoped to this scan and nothing is carried over from an
|
||||
earlier run, so an empty result means no agent has derived one yet.
|
||||
|
||||
Works black-box or white-box. The target can be a host, a URL, an
|
||||
API base, or a repository path; equivalent spellings of the same
|
||||
host resolve to the same model, and a checkout resolves to its
|
||||
remote, so a model derived white-box is read back by a black-box
|
||||
agent testing the deployment.
|
||||
remote, so a model derived white-box by one agent is read back by
|
||||
another testing the deployment.
|
||||
|
||||
Returns ``found: false`` when nothing is cached — derive one and
|
||||
persist it with ``save_threat_model``. ``stale: true`` means the
|
||||
checkout moved to a different revision, or that a model with no
|
||||
revision to pin to has aged out: use it as a starting point,
|
||||
re-confirm what it claims, and save the corrected version.
|
||||
Returns ``found: false`` when nothing has been derived yet — derive
|
||||
one and share it with ``save_threat_model``.
|
||||
|
||||
Any ``amendments`` in the response are corrections other agents
|
||||
recorded after the base model was written. They are part of the
|
||||
@@ -540,10 +519,10 @@ async def get_threat_model(ctx: RunContextWrapper, target: str) -> str:
|
||||
|
||||
@function_tool(timeout=30)
|
||||
async def save_threat_model(ctx: RunContextWrapper, target: str, content: str) -> str:
|
||||
"""Persist a target-scoped threat model for reuse by other agents.
|
||||
"""Share a target-scoped threat model with the other agents on this scan.
|
||||
|
||||
Keyed by target identity, so a later scan of the same host or tree
|
||||
reads it back instead of paying to derive it again.
|
||||
The model lives for this run only — it is not written to disk and a
|
||||
later scan of the same host or tree starts without it.
|
||||
|
||||
**This replaces the whole document, and clears any amendments** —
|
||||
it is for the agent establishing the baseline (normally root,
|
||||
@@ -561,9 +540,9 @@ async def save_threat_model(ctx: RunContextWrapper, target: str, content: str) -
|
||||
necessarily provisional — say which parts are inferred rather than
|
||||
observed, and let later agents amend it as the picture fills in.
|
||||
|
||||
**Scope it to the target, not to this scan.** Do not centre it on
|
||||
the diff you were handed, the subsystem you were assigned, or the
|
||||
one host that happened to answer first. With source, distinguish
|
||||
**Scope it to the target, not to your slice of it.** Do not centre
|
||||
it on the diff you were handed, the subsystem you were assigned, or
|
||||
the one host that happened to answer first. With source, distinguish
|
||||
real product and runtime surfaces from test, docs, example, and
|
||||
developer-tooling paths — in a monorepo, do not let ``tests/`` or
|
||||
one-off scripts become the centre of gravity unless the code shows
|
||||
|
||||
@@ -23,9 +23,26 @@ def write_secret_text(path: Path, text: str) -> None:
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as handle:
|
||||
handle.write(text)
|
||||
except BaseException:
|
||||
with contextlib.suppress(OSError):
|
||||
tmp.unlink()
|
||||
except BaseException as exc:
|
||||
_cleanup_tmp(tmp, exc)
|
||||
raise
|
||||
|
||||
tmp.replace(path)
|
||||
try:
|
||||
tmp.replace(path)
|
||||
except BaseException as exc:
|
||||
_cleanup_tmp(tmp, exc)
|
||||
raise
|
||||
|
||||
|
||||
def _cleanup_tmp(tmp: Path, cause: BaseException) -> None:
|
||||
"""Delete the temporary secret file. A failed delete must not stay silent."""
|
||||
try:
|
||||
tmp.unlink()
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError:
|
||||
message = (
|
||||
f"could not store the secret, and the temporary file {tmp} "
|
||||
f"still holds it. Delete the file manually."
|
||||
)
|
||||
raise OSError(message) from cause
|
||||
|
||||
@@ -228,7 +228,10 @@ def test_resume_still_requires_targets_or_a_workspace(
|
||||
|
||||
assert "has no targets_info" in capsys.readouterr().err
|
||||
|
||||
def test_resume_non_object_run_json_exits(tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str]) -> None:
|
||||
|
||||
def test_resume_non_object_run_json_exits(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str]
|
||||
) -> None:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
run_dir = tmp_path / "strix_runs" / "pentest_abcd"
|
||||
run_dir.mkdir(parents=True)
|
||||
|
||||
3569
tests/test_cloud_cli.py
Normal file
3569
tests/test_cloud_cli.py
Normal file
File diff suppressed because it is too large
Load Diff
1119
tests/test_cloud_cli_runtime.py
Normal file
1119
tests/test_cloud_cli_runtime.py
Normal file
File diff suppressed because it is too large
Load Diff
203
tests/test_cloud_idempotency.py
Normal file
203
tests/test_cloud_idempotency.py
Normal file
@@ -0,0 +1,203 @@
|
||||
"""Durable retry behavior for managed scan-launch commands."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.interface import cloud
|
||||
from strix.interface.cloud import http, runner
|
||||
from strix.interface.completions import completion_candidates
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, payload: Any, *, status_code: int = 200) -> None:
|
||||
self._payload = payload
|
||||
self.status_code = status_code
|
||||
self.headers = {"content-type": "application/json"}
|
||||
self.text = json.dumps(payload)
|
||||
self.closed = False
|
||||
|
||||
def json(self) -> Any:
|
||||
return self._payload
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _token(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("STRIX_API_TOKEN", "idempotency-test-token")
|
||||
monkeypatch.setattr(runner.time, "sleep", lambda _seconds: None)
|
||||
|
||||
|
||||
def test_scan_start_generates_and_sends_one_stable_key(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: Any,
|
||||
) -> None:
|
||||
seen: list[dict[str, Any]] = []
|
||||
monkeypatch.setattr(runner, "uuid4", lambda: "generated-key")
|
||||
|
||||
def request(_method: str, _path: str, **kwargs: Any) -> FakeResponse:
|
||||
seen.append(kwargs)
|
||||
return FakeResponse({"scan_id": "scan-1", "status": "running"})
|
||||
|
||||
monkeypatch.setattr(http, "request", request)
|
||||
assert cloud.run_cloud(["scans", "start", "--domain-ids", "domain-1", "--json"]) == 0
|
||||
assert json.loads(capsys.readouterr().out)["scan_id"] == "scan-1"
|
||||
assert len(seen) == 1
|
||||
assert seen[0]["idempotency_key"] == "generated-key"
|
||||
assert seen[0]["body"]["engagement_type"] == "live_test"
|
||||
|
||||
|
||||
def test_exact_transport_retry_reuses_key_and_body(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
seen: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
def request(_method: str, _path: str, **kwargs: Any) -> FakeResponse:
|
||||
seen.append((kwargs["idempotency_key"], kwargs["body"]))
|
||||
if len(seen) == 1:
|
||||
raise http.CloudTransportError("response lost")
|
||||
return FakeResponse({"scan_id": "scan-1", "status": "running"})
|
||||
|
||||
monkeypatch.setattr(http, "request", request)
|
||||
command = [
|
||||
"scans",
|
||||
"start",
|
||||
"--domain-ids",
|
||||
"domain-1",
|
||||
"--idempotency-key",
|
||||
"retry-key",
|
||||
"--json",
|
||||
]
|
||||
assert cloud.run_cloud(command) == 0
|
||||
assert len(seen) == 2
|
||||
assert seen[0] == seen[1]
|
||||
assert seen[0][0] == "retry-key"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"payload,status",
|
||||
[
|
||||
({"code": "idempotency_request_in_progress", "retry_safe": True}, 409),
|
||||
({"code": "idempotency_outcome_unknown", "retry_safe": True}, 503),
|
||||
({"detail": "gateway unavailable"}, 502),
|
||||
({"detail": "rate limited"}, 429),
|
||||
],
|
||||
)
|
||||
def test_retryable_responses_are_closed_and_replayed(
|
||||
payload: dict[str, Any],
|
||||
status: int,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
first = FakeResponse(payload, status_code=status)
|
||||
responses = iter((first, FakeResponse({"scan_id": "scan-1", "status": "running"})))
|
||||
keys: list[str] = []
|
||||
|
||||
def request(_method: str, _path: str, **kwargs: Any) -> FakeResponse:
|
||||
keys.append(kwargs["idempotency_key"])
|
||||
return next(responses)
|
||||
|
||||
monkeypatch.setattr(http, "request", request)
|
||||
assert (
|
||||
cloud.run_cloud(
|
||||
[
|
||||
"scans",
|
||||
"rerun",
|
||||
"scan-old",
|
||||
"--idempotency-key",
|
||||
"same-key",
|
||||
"--json",
|
||||
]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
assert keys == ["same-key", "same-key"]
|
||||
assert first.closed is True
|
||||
|
||||
|
||||
def test_terminal_key_conflict_is_not_retried(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: Any,
|
||||
) -> None:
|
||||
calls = 0
|
||||
|
||||
def request(_method: str, _path: str, **_kwargs: Any) -> FakeResponse:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
return FakeResponse(
|
||||
{
|
||||
"detail": "key belongs to another request",
|
||||
"code": "idempotency_key_conflict",
|
||||
"terminal": True,
|
||||
},
|
||||
status_code=409,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(http, "request", request)
|
||||
assert (
|
||||
cloud.run_cloud(["scans", "rerun", "scan-old", "--idempotency-key", "conflict", "--json"])
|
||||
== http.EXIT_ERROR
|
||||
)
|
||||
assert calls == 1
|
||||
assert json.loads(capsys.readouterr().out)["code"] == "idempotency_key_conflict"
|
||||
|
||||
|
||||
def test_exhausted_ambiguous_launch_reports_safe_recovery_key(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: Any,
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
http,
|
||||
"request",
|
||||
lambda *_args, **_kwargs: (_ for _ in ()).throw(http.CloudTransportError("response lost")),
|
||||
)
|
||||
assert (
|
||||
cloud.run_cloud(["scans", "rerun", "scan-old", "--idempotency-key", "recover-me", "--json"])
|
||||
== http.EXIT_ERROR
|
||||
)
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["idempotency_key"] == "recover-me"
|
||||
assert payload["retry_safe"] is True
|
||||
assert payload["retry_same_request"] is True
|
||||
assert "--idempotency-key recover-me" in payload["error"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("key", ["", " white", "bad key", "x\nheader", "x" * 201])
|
||||
def test_invalid_idempotency_key_is_usage_error_before_request(
|
||||
key: str,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(http, "request", lambda *_a, **_k: pytest.fail("must not request"))
|
||||
assert (
|
||||
cloud.run_cloud(["scans", "rerun", "scan-old", "--idempotency-key", key, "--json"])
|
||||
== http.EXIT_USAGE
|
||||
)
|
||||
|
||||
|
||||
def test_idempotency_flag_is_completed_only_for_keyed_commands() -> None:
|
||||
assert "--idempotency-key" in completion_candidates(["cloud", "scans", "start", "--idemp"])
|
||||
assert "--idempotency-key" in completion_candidates(
|
||||
["cloud", "scans", "rerun", "scan-1", "--idemp"]
|
||||
)
|
||||
assert "--idempotency-key" not in completion_candidates(["cloud", "scans", "list", "--idemp"])
|
||||
assert "--idempotency-key" in completion_candidates(["cloud", "schedules", "create", "--idemp"])
|
||||
assert "--idempotency-key" in completion_candidates(
|
||||
["cloud", "schedules", "trigger", "schedule-1", "--idemp"]
|
||||
)
|
||||
|
||||
|
||||
def test_http_client_places_key_in_the_header(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
seen: dict[str, Any] = {}
|
||||
|
||||
def request(_method: str, _url: str, **kwargs: Any) -> FakeResponse:
|
||||
seen.update(kwargs)
|
||||
return FakeResponse({"ok": True})
|
||||
|
||||
monkeypatch.setattr(http.requests, "request", request)
|
||||
http.request("POST", "/scans", body={}, idempotency_key="header-key")
|
||||
assert seen["headers"]["Idempotency-Key"] == "header-key"
|
||||
assert seen["headers"]["Authorization"] == "Bearer idempotency-test-token"
|
||||
140
tests/test_cloud_payment_proxy.py
Normal file
140
tests/test_cloud_payment_proxy.py
Normal file
@@ -0,0 +1,140 @@
|
||||
"""Security tests for the wallet payment loopback bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.interface.cloud import payment_proxy
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
|
||||
|
||||
class _StreamingResponse:
|
||||
status_code = 200
|
||||
|
||||
def __init__(self, chunks: list[bytes]) -> None:
|
||||
self.chunks = chunks
|
||||
self.closed = False
|
||||
self.headers = {"Content-Type": "application/json"}
|
||||
|
||||
def iter_content(self, *, chunk_size: int) -> Iterator[bytes]:
|
||||
assert chunk_size > 0
|
||||
yield from self.chunks
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
def _post(url: str, body: bytes, headers: dict[str, str] | None = None) -> bytes:
|
||||
request = urllib.request.Request( # noqa: S310
|
||||
url,
|
||||
data=body,
|
||||
headers={"Content-Type": "application/json", **(headers or {})},
|
||||
method="POST",
|
||||
)
|
||||
with urllib.request.urlopen(request, timeout=2) as response: # noqa: S310
|
||||
return response.read()
|
||||
|
||||
|
||||
def test_bridge_bounds_decompressed_upstream_response(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
response = _StreamingResponse([b"1234", b"5"])
|
||||
|
||||
def fake_request(*_args: Any, **kwargs: Any) -> _StreamingResponse:
|
||||
assert kwargs["stream"] is True
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(payment_proxy, "_MAX_UPSTREAM_RESPONSE_BYTES", 4)
|
||||
monkeypatch.setattr(payment_proxy.requests, "request", fake_request)
|
||||
|
||||
with payment_proxy.wallet_payment_bridge(
|
||||
upstream_url="https://app.example.test/api/v1/billing/topup",
|
||||
api_token="strix-secret", # noqa: S106
|
||||
expected_body=b"{}",
|
||||
) as wallet_url:
|
||||
request = urllib.request.Request( # noqa: S310
|
||||
wallet_url,
|
||||
data=b"{}",
|
||||
headers={"Content-Type": "application/json"},
|
||||
method="POST",
|
||||
)
|
||||
with pytest.raises(urllib.error.HTTPError) as exc_info:
|
||||
urllib.request.urlopen(request, timeout=2) # noqa: S310
|
||||
|
||||
assert exc_info.value.code == 502
|
||||
assert response.closed is True
|
||||
|
||||
|
||||
def test_bridge_forwards_only_the_approved_request_and_protected_headers(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: list[dict[str, Any]] = []
|
||||
observed: list[payment_proxy.WalletUpstreamResponse] = []
|
||||
|
||||
def fake_request(*_args: Any, **kwargs: Any) -> _StreamingResponse:
|
||||
captured.append(kwargs)
|
||||
return _StreamingResponse([b'{"ok":true}'])
|
||||
|
||||
monkeypatch.setattr(payment_proxy.requests, "request", fake_request)
|
||||
with payment_proxy.wallet_payment_bridge(
|
||||
upstream_url="https://app.example.test/api/v1/billing/topup",
|
||||
api_token="strix-secret", # noqa: S106
|
||||
workspace_id="org_trusted",
|
||||
expected_body=b'{"credits":5}',
|
||||
response_observer=observed.append,
|
||||
) as wallet_url:
|
||||
result = _post(
|
||||
wallet_url,
|
||||
b'{"credits":5}',
|
||||
{
|
||||
"Authorization": "Payment wallet-proof",
|
||||
"Proxy-Authorization": "Basic drop-me",
|
||||
"X-Strix-Authorization": "Bearer attacker",
|
||||
"X-Strix-Workspace": "org_attacker",
|
||||
},
|
||||
)
|
||||
|
||||
assert result == b'{"ok":true}'
|
||||
headers = captured[0]["headers"]
|
||||
assert headers["Authorization"] == "Payment wallet-proof"
|
||||
assert headers["X-Strix-Authorization"] == "Bearer strix-secret"
|
||||
assert headers["X-Strix-Workspace"] == "org_trusted"
|
||||
assert "Proxy-Authorization" not in headers
|
||||
assert not any(name.lower() in {"host", "content-length"} for name in headers)
|
||||
assert observed == [
|
||||
payment_proxy.WalletUpstreamResponse(status_code=200, body=b'{"ok":true}')
|
||||
]
|
||||
|
||||
with pytest.raises(urllib.error.HTTPError) as wrong_body:
|
||||
_post(wallet_url, b'{"credits":500}')
|
||||
assert wrong_body.value.code == 403
|
||||
assert len(captured) == 1
|
||||
|
||||
|
||||
def test_bridge_limits_valid_wallet_attempts(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
calls = 0
|
||||
|
||||
def fake_request(*_args: Any, **_kwargs: Any) -> _StreamingResponse:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
return _StreamingResponse([b"{}"])
|
||||
|
||||
monkeypatch.setattr(payment_proxy.requests, "request", fake_request)
|
||||
with payment_proxy.wallet_payment_bridge(
|
||||
upstream_url="https://app.example.test/api/v1/billing/topup",
|
||||
api_token="strix-secret", # noqa: S106
|
||||
expected_body=b"{}",
|
||||
) as wallet_url:
|
||||
assert _post(wallet_url, b"{}") == b"{}"
|
||||
assert _post(wallet_url, b"{}") == b"{}"
|
||||
assert _post(wallet_url, b"{}") == b"{}"
|
||||
with pytest.raises(urllib.error.HTTPError) as extra_request:
|
||||
_post(wallet_url, b"{}")
|
||||
|
||||
assert extra_request.value.code == 429
|
||||
assert calls == payment_proxy._MAX_WALLET_REQUESTS
|
||||
165
tests/test_cloud_session.py
Normal file
165
tests/test_cloud_session.py
Normal file
@@ -0,0 +1,165 @@
|
||||
"""CLI-session lifecycle, scope, and workspace-race behavior."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
from rich.console import Console
|
||||
|
||||
from strix.interface import cloud, platform_cli, platform_identity
|
||||
from strix.interface.cloud import http
|
||||
from strix.interface.cloud import session as cloud_session
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class Response:
|
||||
def __init__(self, payload: Any = None, status_code: int = 200) -> None:
|
||||
self._payload = payload
|
||||
self.status_code = status_code
|
||||
self.ok = 200 <= status_code < 400
|
||||
self.text = json.dumps(payload) if payload is not None else ""
|
||||
self.headers = {"content-type": "application/json"}
|
||||
|
||||
def json(self) -> Any:
|
||||
return self._payload
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def auth_path(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path:
|
||||
path = tmp_path / "platform-auth.json"
|
||||
monkeypatch.setattr(platform_cli, "AUTH_PATH", path)
|
||||
monkeypatch.delenv("STRIX_API_TOKEN", raising=False)
|
||||
monkeypatch.delenv("STRIX_WORKSPACE_ID", raising=False)
|
||||
return path
|
||||
|
||||
|
||||
def test_http_workspace_pin_is_captured_once(
|
||||
auth_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
platform_cli.save_record(
|
||||
{
|
||||
"api_token": "secret",
|
||||
"organization_id": "org_start",
|
||||
"app_url": "https://app.example.test",
|
||||
}
|
||||
)
|
||||
assert auth_path.exists()
|
||||
sent: list[dict[str, str]] = []
|
||||
|
||||
def fake_request(*_args: Any, **kwargs: Any) -> Response:
|
||||
sent.append(dict(kwargs["headers"]))
|
||||
return Response({})
|
||||
|
||||
monkeypatch.setattr(http.requests, "request", fake_request)
|
||||
http.configure()
|
||||
platform_cli.save_record(
|
||||
{
|
||||
"api_token": "secret",
|
||||
"organization_id": "org_changed_elsewhere",
|
||||
"app_url": "https://app.example.test",
|
||||
}
|
||||
)
|
||||
http.request("GET", "/scans")
|
||||
assert sent[0]["X-Strix-Workspace"] == "org_start"
|
||||
|
||||
|
||||
def test_session_scope_update_persists_only_for_stored_session(
|
||||
auth_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
platform_cli.save_record(
|
||||
{
|
||||
"api_token": "secret",
|
||||
"organization_id": "org_1",
|
||||
"app_url": "https://app.example.test",
|
||||
"scopes": ["scans:read"],
|
||||
}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
http,
|
||||
"request",
|
||||
lambda *_args, **_kwargs: Response(
|
||||
{
|
||||
"scopes": ["scans:read", "scans:write", "billing:read"],
|
||||
"requested_scopes": ["scans:read", "scans:write", "billing:read"],
|
||||
"scope_ceiling": ["scans:read", "scans:write", "billing:read"],
|
||||
"scope_profile": "minimal",
|
||||
}
|
||||
),
|
||||
)
|
||||
assert cloud.run_cloud(["session", "scopes", "set", "minimal", "--json"]) == 0
|
||||
stored = platform_cli.read_record()
|
||||
assert stored is not None
|
||||
assert stored["scope_profile"] == "minimal"
|
||||
assert json.loads(capsys.readouterr().out)["scope_profile"] == "minimal"
|
||||
|
||||
before = auth_path.read_text(encoding="utf-8")
|
||||
assert (
|
||||
cloud.run_cloud(["session", "scopes", "set", "minimal", "--token", "override", "--json"])
|
||||
== 0
|
||||
)
|
||||
assert auth_path.read_text(encoding="utf-8") == before
|
||||
|
||||
|
||||
def test_logout_keeps_local_token_when_remote_outcome_is_not_definitive(
|
||||
auth_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
platform_cli.save_record(
|
||||
{
|
||||
"api_token": "secret",
|
||||
"organization_id": "org_1",
|
||||
"app_url": "https://app.example.test",
|
||||
}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
platform_cli.requests,
|
||||
"delete",
|
||||
lambda *_args, **_kwargs: Response({"detail": "unavailable"}, 503),
|
||||
)
|
||||
assert cloud.run_cloud(["logout", "--json"]) == 1
|
||||
assert auth_path.exists()
|
||||
assert json.loads(capsys.readouterr().out)["removed"] is False
|
||||
|
||||
|
||||
def test_local_only_logout_is_explicit_and_recoverable(auth_path: Path, capsys: Any) -> None:
|
||||
platform_cli.save_record({"api_token": "secret"})
|
||||
assert cloud.run_cloud(["logout", "--local-only", "--json"]) == 0
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["local_only"] is True
|
||||
assert payload["remotely_revoked"] is False
|
||||
assert not auth_path.exists()
|
||||
|
||||
|
||||
def test_session_json_errors_preserve_machine_readable_server_details(capsys: Any) -> None:
|
||||
error = http.CloudError(
|
||||
"workspace changed",
|
||||
payload={
|
||||
"detail": "workspace changed",
|
||||
"code": "workspace_session_changed",
|
||||
"current_organization_id": "org_current",
|
||||
},
|
||||
)
|
||||
|
||||
assert cloud_session._error(Console(), error, as_json=True) == http.EXIT_ERROR
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload == {
|
||||
"code": "workspace_session_changed",
|
||||
"current_organization_id": "org_current",
|
||||
"error": "workspace changed",
|
||||
}
|
||||
|
||||
|
||||
def test_cli_device_identity_is_stable_and_privacy_safe(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
path = tmp_path / "cli-identity.json"
|
||||
monkeypatch.setattr(platform_identity, "IDENTITY_PATH", path)
|
||||
first = platform_identity.read_or_create_identity()
|
||||
second = platform_identity.read_or_create_identity(device_name=" Build laptop ")
|
||||
assert second["client_instance_id"] == first["client_instance_id"]
|
||||
assert second["device_name"] == "Build laptop"
|
||||
assert path.stat().st_mode & 0o777 == 0o600
|
||||
657
tests/test_cloud_source_upload.py
Normal file
657
tests/test_cloud_source_upload.py
Normal file
@@ -0,0 +1,657 @@
|
||||
"""Local-source packaging and scan upload tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import zipfile
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.interface import cloud
|
||||
from strix.interface.cloud import http, source_upload
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, payload: Any, status_code: int = 200) -> None:
|
||||
self.status_code = status_code
|
||||
self._payload = payload
|
||||
self.text = json.dumps(payload)
|
||||
self.content = b""
|
||||
self.ok = 200 <= status_code < 400
|
||||
self.headers = {"content-type": "application/json"}
|
||||
|
||||
def json(self) -> Any:
|
||||
return self._payload
|
||||
|
||||
|
||||
class MalformedJsonResponse(FakeResponse):
|
||||
def json(self) -> Any:
|
||||
raise ValueError("malformed JSON")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _token_env(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("STRIX_API_TOKEN", "test-token")
|
||||
|
||||
|
||||
def _git_source(tmp_path: Path) -> Path:
|
||||
git = shutil.which("git")
|
||||
assert git is not None
|
||||
subprocess.run([git, "init", "-q", str(tmp_path)], check=True) # noqa: S603
|
||||
(tmp_path / "app.py").write_text("print('hello')\n", encoding="utf-8")
|
||||
(tmp_path / "README.md").write_text("hello\n", encoding="utf-8")
|
||||
(tmp_path / ".gitignore").write_text("ignored.log\n", encoding="utf-8")
|
||||
(tmp_path / "ignored.log").write_text("ignored\n", encoding="utf-8")
|
||||
subprocess.run( # noqa: S603
|
||||
[git, "-C", str(tmp_path), "add", "app.py", ".gitignore"], check=True
|
||||
)
|
||||
return tmp_path
|
||||
|
||||
|
||||
def test_source_defaults_are_private_and_git_aware(tmp_path: Path) -> None:
|
||||
source = _git_source(tmp_path)
|
||||
(source / ".hidden.py").write_text("hidden\n", encoding="utf-8")
|
||||
(source / ".env").write_text("TOKEN=secret\n", encoding="utf-8")
|
||||
(source / "private.pem").write_text("secret\n", encoding="utf-8")
|
||||
(source / "fixture.zip").write_bytes(b"not really a zip")
|
||||
(source / "node_modules").mkdir()
|
||||
(source / "node_modules" / "dep.js").write_text("dep\n", encoding="utf-8")
|
||||
(source / "linked.py").symlink_to(source / "app.py")
|
||||
|
||||
bundle = source_upload.prepare_source(
|
||||
str(source),
|
||||
include_hidden=False,
|
||||
include_sensitive=False,
|
||||
include_archives=False,
|
||||
exclude=[],
|
||||
)
|
||||
try:
|
||||
names = [item.archive_name for item in bundle.manifest.files]
|
||||
assert names == ["README.md", "app.py"]
|
||||
assert bundle.manifest.total_bytes > 0
|
||||
assert bundle.archive_bytes <= source_upload.MAX_ARCHIVE_BYTES
|
||||
assert bundle.manifest.excluded["hidden"] == 3
|
||||
assert bundle.manifest.excluded["sensitive_filename"] == 1
|
||||
assert bundle.manifest.excluded["nested_archive"] == 1
|
||||
assert bundle.manifest.excluded["dependency_or_build_output"] == 1
|
||||
assert bundle.manifest.excluded["symlink_or_non_file"] == 1
|
||||
with zipfile.ZipFile(bundle.archive_path) as archive:
|
||||
assert archive.namelist() == names
|
||||
finally:
|
||||
source_upload.remove_bundle(bundle)
|
||||
|
||||
|
||||
def test_hidden_and_sensitive_files_need_separate_opt_ins(tmp_path: Path) -> None:
|
||||
source = _git_source(tmp_path)
|
||||
(source / ".env").write_text("TOKEN=secret\n", encoding="utf-8")
|
||||
(source / ".github").mkdir()
|
||||
(source / ".github" / "workflow.yml").write_text("name: test\n", encoding="utf-8")
|
||||
|
||||
hidden = source_upload.select_source(source, include_hidden=True)
|
||||
hidden_names = {item.archive_name for item in hidden.files}
|
||||
assert ".github/workflow.yml" in hidden_names
|
||||
assert ".env" not in hidden_names
|
||||
|
||||
sensitive = source_upload.select_source(source, include_hidden=True, include_sensitive=True)
|
||||
assert ".env" in {item.archive_name for item in sensitive.files}
|
||||
assert all(
|
||||
not name.startswith(".git/") for name in (item.archive_name for item in sensitive.files)
|
||||
)
|
||||
|
||||
|
||||
def test_hidden_opt_in_still_excludes_common_credential_paths(tmp_path: Path) -> None:
|
||||
paths = [
|
||||
".aws/credentials",
|
||||
".git-credentials",
|
||||
".docker/config.json",
|
||||
".config/gcloud/application_default_credentials.json",
|
||||
".config/gcloud/credentials.db",
|
||||
".azure/accessTokens.json",
|
||||
".kube/config",
|
||||
]
|
||||
for relative in paths:
|
||||
path = tmp_path / relative
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text("credential material\n", encoding="utf-8")
|
||||
|
||||
hidden_only = source_upload.select_source(tmp_path, include_hidden=True)
|
||||
assert not ({item.archive_name for item in hidden_only.files} & set(paths))
|
||||
assert hidden_only.excluded["sensitive_filename"] == len(paths)
|
||||
|
||||
explicitly_sensitive = source_upload.select_source(
|
||||
tmp_path, include_hidden=True, include_sensitive=True
|
||||
)
|
||||
assert set(paths) <= {item.archive_name for item in explicitly_sensitive.files}
|
||||
|
||||
|
||||
def test_hidden_opt_in_cannot_reenable_dependency_cache_or_build_dirs(tmp_path: Path) -> None:
|
||||
excluded_dirs = [
|
||||
".venv",
|
||||
"env",
|
||||
".tox",
|
||||
".pytest_cache",
|
||||
".mypy_cache",
|
||||
".ruff_cache",
|
||||
".next",
|
||||
".nuxt",
|
||||
".gradle",
|
||||
]
|
||||
for directory in excluded_dirs:
|
||||
path = tmp_path / directory / "artifact.txt"
|
||||
path.parent.mkdir(parents=True)
|
||||
path.write_text("generated\n", encoding="utf-8")
|
||||
(tmp_path / ".github" / "workflow.yml").parent.mkdir()
|
||||
(tmp_path / ".github" / "workflow.yml").write_text("name: test\n", encoding="utf-8")
|
||||
|
||||
manifest = source_upload.select_source(tmp_path, include_hidden=True)
|
||||
|
||||
names = {item.archive_name for item in manifest.files}
|
||||
assert ".github/workflow.yml" in names
|
||||
assert not any(name.split("/", 1)[0] in excluded_dirs for name in names)
|
||||
assert manifest.excluded["dependency_or_build_output"] == len(excluded_dirs)
|
||||
|
||||
|
||||
def test_strixignore_and_cli_excludes_are_applied(tmp_path: Path) -> None:
|
||||
(tmp_path / "keep.py").write_text("keep\n", encoding="utf-8")
|
||||
(tmp_path / "generated.py").write_text("generated\n", encoding="utf-8")
|
||||
(tmp_path / "test_app.py").write_text("test\n", encoding="utf-8")
|
||||
(tmp_path / ".strixignore").write_text("generated.py\n", encoding="utf-8")
|
||||
|
||||
manifest = source_upload.select_source(tmp_path, exclude=["test_*.py"])
|
||||
assert [item.archive_name for item in manifest.files] == ["keep.py"]
|
||||
assert manifest.excluded["user_pattern"] == 2
|
||||
|
||||
|
||||
def test_strixignore_trailing_slash_excludes_the_whole_directory(tmp_path: Path) -> None:
|
||||
(tmp_path / "keep.py").write_text("keep\n", encoding="utf-8")
|
||||
private = tmp_path / "private" / "nested"
|
||||
private.mkdir(parents=True)
|
||||
(private / "secret.txt").write_text("do not upload\n", encoding="utf-8")
|
||||
cache = tmp_path / "packages" / "cache"
|
||||
cache.mkdir(parents=True)
|
||||
(cache / "artifact.txt").write_text("do not upload\n", encoding="utf-8")
|
||||
(tmp_path / ".strixignore").write_text("private/\n", encoding="utf-8")
|
||||
|
||||
manifest = source_upload.select_source(tmp_path, exclude=["cache/"])
|
||||
|
||||
assert [item.archive_name for item in manifest.files] == ["keep.py"]
|
||||
assert manifest.excluded["user_pattern"] >= 2
|
||||
|
||||
|
||||
def test_source_limits_expanded_bytes_before_compression(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(source_upload, "MAX_TOTAL_BYTES", 5)
|
||||
(tmp_path / "large.py").write_bytes(b"a" * 6)
|
||||
with pytest.raises(http.CloudError, match="expanded-size limit"):
|
||||
source_upload.select_source(tmp_path)
|
||||
|
||||
|
||||
def test_source_rejects_archives_by_suffix_and_actual_bytes(tmp_path: Path) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
(tmp_path / "dependency.jar").write_bytes(b"not-even-a-valid-archive")
|
||||
(tmp_path / "renamed-source.txt").write_bytes(b"PK\x03\x04" + b"x" * 32)
|
||||
tar_header = bytearray(512)
|
||||
tar_header[257:262] = b"ustar"
|
||||
(tmp_path / "renamed-tar.bin").write_bytes(tar_header)
|
||||
|
||||
manifest = source_upload.select_source(tmp_path)
|
||||
|
||||
assert [item.archive_name for item in manifest.files] == ["app.py"]
|
||||
assert manifest.excluded["nested_archive"] == 3
|
||||
|
||||
|
||||
def test_source_enumeration_is_bounded_before_filtering(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setattr(source_upload, "MAX_CANDIDATE_PATHS", 2)
|
||||
for index in range(3):
|
||||
(tmp_path / f"file-{index}.py").write_text("safe\n", encoding="utf-8")
|
||||
|
||||
with pytest.raises(http.CloudError, match="enumeration exceeded 2 paths"):
|
||||
source_upload.select_source(tmp_path)
|
||||
|
||||
|
||||
def test_strixignore_size_and_pattern_counts_are_bounded(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
ignore = tmp_path / ".strixignore"
|
||||
monkeypatch.setattr(source_upload, "MAX_IGNORE_BYTES", 4)
|
||||
ignore.write_text("12345", encoding="utf-8")
|
||||
with pytest.raises(http.CloudError, match="larger than the 4-byte limit"):
|
||||
source_upload.select_source(tmp_path)
|
||||
|
||||
monkeypatch.setattr(source_upload, "MAX_IGNORE_BYTES", 1_000)
|
||||
monkeypatch.setattr(source_upload, "MAX_IGNORE_PATTERNS", 1)
|
||||
ignore.write_text("one\ntwo\n", encoding="utf-8")
|
||||
with pytest.raises(http.CloudError, match="more than 1 exclusion patterns"):
|
||||
source_upload.select_source(tmp_path)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not hasattr(os, "mkfifo"), reason="named pipes are not supported")
|
||||
def test_strixignore_must_be_a_nonblocking_regular_file(tmp_path: Path) -> None:
|
||||
(tmp_path / "app.py").write_text("print('ok')\n", encoding="utf-8")
|
||||
os.mkfifo(tmp_path / ".strixignore")
|
||||
|
||||
with pytest.raises(http.CloudError, match="must be a regular file"):
|
||||
source_upload.select_source(tmp_path)
|
||||
|
||||
|
||||
def test_source_archive_rejects_a_path_swapped_after_manifest_review(tmp_path: Path) -> None:
|
||||
source_path = tmp_path / "app.py"
|
||||
source_path.write_bytes(b"safe")
|
||||
manifest = source_upload.select_source(tmp_path)
|
||||
|
||||
replacement = tmp_path / "replacement"
|
||||
replacement.write_bytes(b"oops")
|
||||
replacement.replace(source_path)
|
||||
|
||||
with pytest.raises(http.CloudError, match="changed while the source archive was being built"):
|
||||
source_upload._write_archive(tmp_path / "source.zip", manifest.files)
|
||||
|
||||
|
||||
def test_source_archive_rejects_same_inode_same_size_change_after_review(tmp_path: Path) -> None:
|
||||
source_path = tmp_path / "app.py"
|
||||
source_path.write_bytes(b"safe")
|
||||
manifest = source_upload.select_source(tmp_path)
|
||||
|
||||
source_path.write_bytes(b"evil")
|
||||
selected = manifest.files[0]
|
||||
os.utime(
|
||||
source_path,
|
||||
ns=(selected.mtime_ns + 1_000_000, selected.mtime_ns + 1_000_000),
|
||||
)
|
||||
|
||||
with pytest.raises(http.CloudError, match="changed while the source archive was being built"):
|
||||
source_upload._write_archive(tmp_path / "source.zip", manifest.files)
|
||||
|
||||
|
||||
def test_source_dry_run_never_calls_the_api(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
|
||||
def fail_request(*_args: Any, **_kwargs: Any) -> Any:
|
||||
raise AssertionError("dry-run must not make an API request")
|
||||
|
||||
monkeypatch.setattr(http, "request", fail_request)
|
||||
assert (
|
||||
cloud.run_cloud(
|
||||
["scans", "start", "--source", str(tmp_path), "--dry-run", "--show-files", "--json"]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["source"]["files"] == ["app.py"]
|
||||
assert payload["source"]["archive_sha256"]
|
||||
|
||||
|
||||
def test_noninteractive_source_upload_requires_yes(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
monkeypatch.setattr(
|
||||
http,
|
||||
"request",
|
||||
lambda *_args, **_kwargs: pytest.fail("approval must happen before any API request"),
|
||||
)
|
||||
assert cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--json"]) == 1
|
||||
output = capsys.readouterr().out
|
||||
assert "requires explicit approval" in output
|
||||
assert "--approve-sha256 <reviewed hash>" in output
|
||||
assert "--yes" in output
|
||||
assert "one-shot approval" in output
|
||||
|
||||
|
||||
def test_source_digest_approval_rejects_a_changed_snapshot(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
source = tmp_path / "app.py"
|
||||
source.write_text("print('reviewed')\n", encoding="utf-8")
|
||||
assert (
|
||||
cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--dry-run", "--json"]) == 0
|
||||
)
|
||||
approved = json.loads(capsys.readouterr().out)["source"]["archive_sha256"]
|
||||
source.write_text("print('changed')\n", encoding="utf-8")
|
||||
monkeypatch.setattr(
|
||||
http,
|
||||
"request",
|
||||
lambda *_args, **_kwargs: pytest.fail("a changed snapshot must not reach the API"),
|
||||
)
|
||||
|
||||
assert (
|
||||
cloud.run_cloud(
|
||||
[
|
||||
"scans",
|
||||
"start",
|
||||
"--source",
|
||||
str(tmp_path),
|
||||
"--approve-sha256",
|
||||
approved,
|
||||
"--json",
|
||||
]
|
||||
)
|
||||
== http.EXIT_ERROR
|
||||
)
|
||||
assert "does not match" in json.loads(capsys.readouterr().out)["error"]
|
||||
|
||||
|
||||
def test_source_upload_is_completed_and_attached_to_scan(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
calls: list[tuple[str, str, dict[str, Any]]] = []
|
||||
uploaded_path: Path | None = None
|
||||
|
||||
def fake_request(method: str, path: str, **kwargs: Any) -> FakeResponse:
|
||||
calls.append((method, path, kwargs))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-1",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/complete":
|
||||
return FakeResponse({"id": "upload-1"})
|
||||
if path == "/scans":
|
||||
return FakeResponse({"scan_id": "scan-1", "status": "pending"})
|
||||
raise AssertionError(path)
|
||||
|
||||
def fake_upload(_url: str, _token: str, path: Path) -> None:
|
||||
nonlocal uploaded_path
|
||||
uploaded_path = path
|
||||
assert path.exists()
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(http, "upload_file", fake_upload)
|
||||
|
||||
assert (
|
||||
cloud.run_cloud(
|
||||
["scans", "start", "--source", str(tmp_path), "--yes", "--show-files", "--json"]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["upload_id"] == "upload-1"
|
||||
assert payload["scan"]["scan_id"] == "scan-1"
|
||||
assert payload["source"]["files"] == ["app.py"]
|
||||
scan_call = next(call for call in calls if call[1] == "/scans")
|
||||
assert scan_call[2]["body"] == {
|
||||
"engagement_type": "code_review",
|
||||
"upload_ids": ["upload-1"],
|
||||
}
|
||||
assert uploaded_path is not None and not uploaded_path.exists()
|
||||
|
||||
|
||||
def test_source_upload_with_domain_is_a_live_test(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
calls: list[tuple[str, str, dict[str, Any]]] = []
|
||||
|
||||
def fake_request(method: str, path: str, **kwargs: Any) -> FakeResponse:
|
||||
calls.append((method, path, kwargs))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-1",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/complete":
|
||||
return FakeResponse({"id": "upload-1"})
|
||||
if path == "/scans":
|
||||
return FakeResponse({"scan_id": "scan-1", "status": "pending"})
|
||||
raise AssertionError(path)
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(http, "upload_file", lambda *_args, **_kwargs: None)
|
||||
|
||||
assert (
|
||||
cloud.run_cloud(
|
||||
[
|
||||
"scans",
|
||||
"start",
|
||||
"--source",
|
||||
str(tmp_path),
|
||||
"--domain-ids",
|
||||
"domain-1",
|
||||
"--yes",
|
||||
"--json",
|
||||
]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
scan_call = next(call for call in calls if call[1] == "/scans")
|
||||
assert scan_call[2]["body"] == {
|
||||
"engagement_type": "live_test",
|
||||
"domain_ids": ["domain-1"],
|
||||
"upload_ids": ["upload-1"],
|
||||
}
|
||||
assert json.loads(capsys.readouterr().out)["scan"]["scan_id"] == "scan-1"
|
||||
|
||||
|
||||
def test_failed_scan_deletes_completed_source_upload(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
paths: list[tuple[str, str]] = []
|
||||
|
||||
def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse:
|
||||
paths.append((method, path))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-1",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/complete":
|
||||
return FakeResponse({"id": "upload-1"})
|
||||
if path == "/scans":
|
||||
return FakeResponse({"detail": "not enough credits"}, status_code=402)
|
||||
if path == "/uploads/upload-1":
|
||||
return FakeResponse({"ok": True})
|
||||
raise AssertionError(path)
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(http, "upload_file", lambda *_args, **_kwargs: None)
|
||||
assert cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--yes", "--json"]) == 5
|
||||
assert ("DELETE", "/uploads/upload-1") in paths
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"failure",
|
||||
["network", "server", "malformed_success", "malformed_json_success", "wrong_shape_success"],
|
||||
)
|
||||
def test_ambiguous_scan_launch_retains_completed_source_upload(
|
||||
failure: str,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: Any,
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
paths: list[tuple[str, str]] = []
|
||||
|
||||
def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse:
|
||||
paths.append((method, path))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-ambiguous",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/complete":
|
||||
return FakeResponse({"id": "upload-ambiguous"})
|
||||
if path == "/scans":
|
||||
if failure == "network":
|
||||
raise http.CloudError("connection closed before a response")
|
||||
if failure == "server":
|
||||
return FakeResponse({"detail": "temporary failure"}, status_code=500)
|
||||
if failure == "malformed_json_success":
|
||||
return MalformedJsonResponse("accepted")
|
||||
if failure == "wrong_shape_success":
|
||||
return FakeResponse({})
|
||||
response = FakeResponse("accepted")
|
||||
response.headers = {"content-type": "text/html"}
|
||||
return response
|
||||
if path == "/uploads/upload-ambiguous":
|
||||
pytest.fail("an upload with an ambiguous launch must not be deleted")
|
||||
raise AssertionError(path)
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(http, "upload_file", lambda *_args, **_kwargs: None)
|
||||
|
||||
assert (
|
||||
cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--yes", "--json"])
|
||||
== http.EXIT_ERROR
|
||||
)
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["upload_id"] == "upload-ambiguous"
|
||||
assert payload["upload_retained"] is True
|
||||
assert payload["launch_outcome_unknown"] is True
|
||||
assert "outcome is unknown" in payload["error"]
|
||||
assert ("DELETE", "/uploads/upload-ambiguous") not in paths
|
||||
|
||||
|
||||
def test_mismatched_upload_completion_response_is_cleaned_before_launch(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
paths: list[tuple[str, str]] = []
|
||||
|
||||
def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse:
|
||||
paths.append((method, path))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-expected",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/complete":
|
||||
return FakeResponse({"id": "upload-different"})
|
||||
if path == "/uploads/upload-expected":
|
||||
return FakeResponse({"ok": True})
|
||||
if path == "/scans":
|
||||
pytest.fail("a scan must not launch before upload completion is confirmed")
|
||||
raise AssertionError(path)
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(http, "upload_file", lambda *_args, **_kwargs: None)
|
||||
|
||||
assert cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--yes"]) == 1
|
||||
assert ("DELETE", "/uploads/upload-expected") in paths
|
||||
|
||||
|
||||
def test_interrupted_scan_launch_retains_completed_source_upload(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
paths: list[tuple[str, str]] = []
|
||||
|
||||
def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse:
|
||||
paths.append((method, path))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-interrupted",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/complete":
|
||||
return FakeResponse({"id": "upload-interrupted"})
|
||||
if path == "/scans":
|
||||
raise KeyboardInterrupt
|
||||
if path == "/uploads/upload-interrupted":
|
||||
pytest.fail("an upload with an interrupted launch must not be deleted")
|
||||
raise AssertionError(path)
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(http, "upload_file", lambda *_args, **_kwargs: None)
|
||||
|
||||
assert cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--yes", "--json"]) == 130
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["interrupted"] is True
|
||||
assert payload["upload_id"] == "upload-interrupted"
|
||||
assert payload["upload_retained"] is True
|
||||
assert payload["launch_outcome_unknown"] is True
|
||||
assert "scans list" in payload["error"]
|
||||
assert ("DELETE", "/uploads/upload-interrupted") not in paths
|
||||
|
||||
|
||||
@pytest.mark.parametrize("cleanup_failure", ["timeout", "server"])
|
||||
def test_failed_automatic_upload_cleanup_reports_retained_id(
|
||||
cleanup_failure: str,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
capsys: Any,
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
|
||||
def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse:
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse(
|
||||
{
|
||||
"upload_id": "upload-orphaned",
|
||||
"signed_url": "https://storage.test/object",
|
||||
"token": "signed",
|
||||
}
|
||||
)
|
||||
if path == "/uploads/upload-orphaned" and method == "DELETE":
|
||||
if cleanup_failure == "timeout":
|
||||
raise http.CloudError("cleanup timed out")
|
||||
return FakeResponse({"detail": "cleanup unavailable"}, status_code=500)
|
||||
raise AssertionError((method, path))
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
monkeypatch.setattr(
|
||||
http,
|
||||
"upload_file",
|
||||
lambda *_args, **_kwargs: (_ for _ in ()).throw(http.CloudError("upload failed")),
|
||||
)
|
||||
|
||||
assert (
|
||||
cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--yes", "--json"])
|
||||
== http.EXIT_ERROR
|
||||
)
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["upload_id"] == "upload-orphaned"
|
||||
assert payload["upload_retained"] is True
|
||||
assert payload["cleanup_unknown"] is True
|
||||
assert "uploads delete upload-orphaned" in payload["error"]
|
||||
|
||||
|
||||
def test_incomplete_upload_credentials_delete_the_reserved_upload(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
(tmp_path / "app.py").write_text("print('safe')\n", encoding="utf-8")
|
||||
paths: list[tuple[str, str]] = []
|
||||
|
||||
def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse:
|
||||
paths.append((method, path))
|
||||
if path == "/uploads/request":
|
||||
return FakeResponse({"upload_id": "upload-incomplete"})
|
||||
if path == "/uploads/upload-incomplete":
|
||||
return FakeResponse({"ok": True})
|
||||
raise AssertionError(path)
|
||||
|
||||
monkeypatch.setattr(http, "request", fake_request)
|
||||
assert cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--yes", "--json"]) == 1
|
||||
assert ("DELETE", "/uploads/upload-incomplete") in paths
|
||||
118
tests/test_cloud_wallet.py
Normal file
118
tests/test_cloud_wallet.py
Normal file
@@ -0,0 +1,118 @@
|
||||
"""Tests for the Stripe Link wallet setup path of `strix cloud billing topup`."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from rich.console import Console
|
||||
|
||||
from strix.interface.cloud import billing
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import pytest
|
||||
|
||||
|
||||
_MIN_LINK_CONTEXT_CHARS = 100
|
||||
|
||||
|
||||
def _completed(stdout: str) -> subprocess.CompletedProcess[str]:
|
||||
return subprocess.CompletedProcess(args=["link-cli"], returncode=0, stdout=stdout, stderr="")
|
||||
|
||||
|
||||
def test_payment_context_is_long_enough_for_link_approval() -> None:
|
||||
context = billing._payment_context({"credits": 5})
|
||||
assert len(context) >= _MIN_LINK_CONTEXT_CHARS
|
||||
assert "5" in context
|
||||
|
||||
|
||||
def test_mppx_wallet_configured_follows_environment(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("MPPX_ACCOUNT", raising=False)
|
||||
monkeypatch.delenv("MPPX_STRIPE_SECRET_KEY", raising=False)
|
||||
assert billing._mppx_wallet_configured() is False
|
||||
monkeypatch.setenv("MPPX_ACCOUNT", "agent")
|
||||
assert billing._mppx_wallet_configured() is True
|
||||
|
||||
|
||||
def test_link_wallet_authenticated_reads_status_list(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(
|
||||
billing,
|
||||
"_run_link_cli",
|
||||
lambda *_args, **_kwargs: _completed('[{"authenticated": true}]'),
|
||||
)
|
||||
assert billing._link_wallet_authenticated("npx") is True
|
||||
|
||||
|
||||
def test_link_wallet_authenticated_handles_unusable_output(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(billing, "_run_link_cli", lambda *_args, **_kwargs: _completed("not json"))
|
||||
assert billing._link_wallet_authenticated("npx") is False
|
||||
|
||||
|
||||
def test_link_wallet_authenticated_handles_launch_failure(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
def explode(*_args: Any, **_kwargs: Any) -> subprocess.CompletedProcess[str]:
|
||||
raise OSError
|
||||
|
||||
monkeypatch.setattr(billing, "_run_link_cli", explode)
|
||||
assert billing._link_wallet_authenticated("npx") is False
|
||||
|
||||
|
||||
def test_pending_spend_request_reads_the_created_record() -> None:
|
||||
stdout = (
|
||||
'[{"id": "lsrq_123", "status": "pending_approval", '
|
||||
'"approval_url": "https://app.link.com/activity/approve/lsrq_123"}]'
|
||||
)
|
||||
assert billing._pending_spend_request(stdout) == (
|
||||
"lsrq_123",
|
||||
"https://app.link.com/activity/approve/lsrq_123",
|
||||
)
|
||||
assert billing._pending_spend_request('[{"id": "lsrq_1", "status": "approved"}]') is None
|
||||
assert billing._pending_spend_request("not json") is None
|
||||
|
||||
|
||||
def test_pending_spend_request_tolerates_banner_text_around_pretty_json() -> None:
|
||||
stdout = (
|
||||
"Update available for @stripe/link-cli: 0.13.1 -> 0.16.0\n"
|
||||
"[\n {\n"
|
||||
' "id": "lsrq_9",\n'
|
||||
' "status": "pending_approval",\n'
|
||||
' "approval_url": "https://app.link.com/activity/approve/lsrq_9"\n'
|
||||
" }\n]"
|
||||
)
|
||||
assert billing._pending_spend_request(stdout) == (
|
||||
"lsrq_9",
|
||||
"https://app.link.com/activity/approve/lsrq_9",
|
||||
)
|
||||
|
||||
|
||||
def test_final_spend_request_status_reads_the_last_poll_line() -> None:
|
||||
stdout = '{"status": "pending_approval"}\n{"status": "approved"}\n'
|
||||
assert billing._final_spend_request_status(stdout) == "approved"
|
||||
assert billing._final_spend_request_status("") is None
|
||||
|
||||
|
||||
def test_final_spend_request_status_unwraps_chunk_envelopes() -> None:
|
||||
stdout = (
|
||||
'{"type":"chunk","data":{"id":"lsrq_9","status":"pending_approval"}}\n'
|
||||
'{"type":"chunk","data":{"id":"lsrq_9","status":"approved"}}\n'
|
||||
'{"type":"done","ok":true,"meta":{"command":"spend-request retrieve"}}\n'
|
||||
)
|
||||
assert billing._final_spend_request_status(stdout) == "approved"
|
||||
|
||||
|
||||
def test_prepare_link_wallet_skips_login_when_connected(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(billing, "_link_wallet_authenticated", lambda _npx: True)
|
||||
assert billing._prepare_link_wallet(Console(), "npx", as_json=True) is None
|
||||
|
||||
|
||||
def test_prepare_link_wallet_explains_setup_without_a_terminal(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(billing, "_link_wallet_authenticated", lambda _npx: False)
|
||||
message = billing._prepare_link_wallet(Console(), "npx", as_json=True)
|
||||
assert message is not None
|
||||
assert "https://link.com/agents" in message
|
||||
164
tests/test_completions.py
Normal file
164
tests/test_completions.py
Normal file
@@ -0,0 +1,164 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from strix.interface.completions import completion_candidates, run_completions
|
||||
|
||||
|
||||
def test_root_completion_candidates() -> None:
|
||||
assert completion_candidates(["cl"]) == ["cloud"]
|
||||
assert "completions" in completion_candidates([""])
|
||||
|
||||
|
||||
def test_cloud_group_and_alias_candidates() -> None:
|
||||
candidates = completion_candidates(["cloud", "work"])
|
||||
assert candidates == ["workspace", "workspaces"]
|
||||
|
||||
|
||||
def test_cloud_verb_candidates_include_multiword_prefixes() -> None:
|
||||
assert "test-users" in completion_candidates(["cloud", "domains", ""])
|
||||
assert completion_candidates(["cloud", "domains", "test-users", "in"]) == [
|
||||
"inbox",
|
||||
"inbox-message",
|
||||
]
|
||||
|
||||
|
||||
def test_cloud_leaf_flag_candidates_come_from_command_spec() -> None:
|
||||
candidates = completion_candidates(["cloud", "scans", "start", "--"])
|
||||
assert "--domain-ids" in candidates
|
||||
assert "--json" in candidates
|
||||
assert "--wait" in candidates
|
||||
assert "--source" in candidates
|
||||
assert "--approve-sha256" in candidates
|
||||
assert "--dry-run" in candidates
|
||||
assert "--include-hidden" in candidates
|
||||
|
||||
|
||||
def test_boolean_completion_includes_positive_and_negative_flags() -> None:
|
||||
candidates = completion_candidates(["cloud", "billing", "auto-topup", "update", "--"])
|
||||
assert "--enabled" in candidates
|
||||
assert "--no-enabled" in candidates
|
||||
assert "--no-monthly-cap" in candidates
|
||||
|
||||
|
||||
def test_session_and_workspace_use_completions_include_their_real_flags() -> None:
|
||||
credit_flags = completion_candidates(["cloud", "credits", "--"])
|
||||
assert {"--json", "--token", "--app-url", "--timeout", "--help"} <= set(credit_flags)
|
||||
assert "--json" in completion_candidates(["cloud", "logout", "--"])
|
||||
|
||||
workspace_use = completion_candidates(["cloud", "workspace", "use", "--"])
|
||||
assert {"--scopes", "--json", "--token", "--app-url", "--timeout"} <= set(workspace_use)
|
||||
|
||||
|
||||
def test_leaf_flags_remain_available_after_options_and_positionals() -> None:
|
||||
after_option = completion_candidates(["cloud", "scans", "list", "--status", "running", "--"])
|
||||
assert {"--page", "--limit", "--json"} <= set(after_option)
|
||||
|
||||
after_positional = completion_candidates(["cloud", "scans", "get", "scan-1", "--"])
|
||||
assert {"--json", "--token", "--app-url", "--timeout"} <= set(after_positional)
|
||||
|
||||
|
||||
def test_default_verbs_complete_flags_without_an_explicit_verb() -> None:
|
||||
audit = completion_candidates(["cloud", "audit", "--"])
|
||||
assert {"--page", "--limit", "--json"} <= set(audit)
|
||||
|
||||
after_option = completion_candidates(["cloud", "audit", "--page", "2", "--"])
|
||||
assert {"--limit", "--format", "--json"} <= set(after_option)
|
||||
|
||||
workspaces = completion_candidates(["cloud", "workspace", "--"])
|
||||
assert {"--json", "--token", "--app-url", "--timeout"} <= set(workspaces)
|
||||
|
||||
|
||||
def test_completion_does_not_offer_flags_while_an_option_value_is_empty() -> None:
|
||||
assert completion_candidates(["cloud", "scans", "list", "--page", ""]) == []
|
||||
assert completion_candidates(["cloud", "scans", "start", "--approve-sha256", ""]) == []
|
||||
|
||||
|
||||
def test_exact_verbs_that_are_also_prefixes_keep_their_subverbs() -> None:
|
||||
candidates = completion_candidates(["cloud", "billing", "auto-topup", ""])
|
||||
assert "update" in candidates
|
||||
assert "--json" in candidates
|
||||
|
||||
|
||||
def test_contract_fix_flags_are_completed() -> None:
|
||||
integration_connect = completion_candidates(
|
||||
["cloud", "integrations", "connect", "gitlab", "--"]
|
||||
)
|
||||
assert {
|
||||
"--provider-token",
|
||||
"--instance-url",
|
||||
"--account-email",
|
||||
"--installation-id",
|
||||
} <= set(integration_connect)
|
||||
|
||||
disconnect = completion_candidates(["cloud", "integrations", "disconnect", "--"])
|
||||
assert "--installation-id" in disconnect
|
||||
|
||||
connector = completion_candidates(["cloud", "connectors", "get", "connector-1", "--"])
|
||||
assert "--include-command" in connector
|
||||
assert "--no-include-command" in connector
|
||||
|
||||
scan_wait = completion_candidates(["cloud", "scans", "start", "--"])
|
||||
assert "--wait-timeout" in scan_wait
|
||||
|
||||
audit_export = completion_candidates(["cloud", "audit", "--"])
|
||||
assert {"--output", "--force"} <= set(audit_export)
|
||||
|
||||
token_create = completion_candidates(["cloud", "tokens", "create", "--"])
|
||||
assert {"--expires-at", "--rbac-scopes"} <= set(token_create)
|
||||
|
||||
|
||||
def test_filesystem_completion_for_source_output_and_data(tmp_path: Any, monkeypatch: Any) -> None:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
(tmp_path / "source tree").mkdir()
|
||||
(tmp_path / "source.txt").write_text("source", encoding="utf-8")
|
||||
(tmp_path / "request.json").write_text("{}", encoding="utf-8")
|
||||
|
||||
source = completion_candidates(["cloud", "scans", "start", "--source", "sou"])
|
||||
assert source == ["source tree/"]
|
||||
|
||||
output = completion_candidates(["cloud", "scans", "report", "scan-1", "--output", "req"])
|
||||
assert output == ["request.json"]
|
||||
|
||||
audit_output = completion_candidates(["cloud", "audit", "--output", "req"])
|
||||
assert audit_output == ["request.json"]
|
||||
|
||||
data = completion_candidates(["cloud", "scans", "start", "--data", "@req"])
|
||||
assert data == ["@request.json"]
|
||||
|
||||
|
||||
def test_filesystem_completion_omits_terminal_control_names(
|
||||
tmp_path: Any, monkeypatch: Any, capsys: Any
|
||||
) -> None:
|
||||
monkeypatch.chdir(tmp_path)
|
||||
(tmp_path / "safe.json").write_text("{}", encoding="utf-8")
|
||||
(tmp_path / "unsafe\nname.json").write_text("{}", encoding="utf-8")
|
||||
(tmp_path / "unsafe\x1b]52;c;payload\x07.json").write_text("{}", encoding="utf-8")
|
||||
|
||||
words = ["cloud", "scans", "start", "--data", "@"]
|
||||
assert completion_candidates(words) == ["@safe.json"]
|
||||
assert run_completions(["--candidates", *words]) == 0
|
||||
assert capsys.readouterr().out == "@safe.json\n"
|
||||
|
||||
|
||||
def test_completion_scripts_cover_supported_shells(capsys: Any) -> None:
|
||||
for shell in ("zsh", "bash", "fish"):
|
||||
assert run_completions([shell]) == 0
|
||||
output = capsys.readouterr().out
|
||||
assert "completions --candidates" in output
|
||||
|
||||
|
||||
def test_bash_completion_preserves_candidates_with_spaces(capsys: Any) -> None:
|
||||
assert run_completions(["bash"]) == 0
|
||||
output = capsys.readouterr().out
|
||||
assert 'COMPREPLY=("${candidates[@]}")' in output
|
||||
assert "while IFS= read -r candidate" in output
|
||||
assert "mapfile" not in output
|
||||
|
||||
|
||||
def test_completion_rejects_unknown_shell(capsys: Any) -> None:
|
||||
assert run_completions(["powershell\x1b]52;c;payload\x07"]) == 2
|
||||
error = capsys.readouterr().err
|
||||
assert "Choose zsh, bash, or fish" in error
|
||||
assert "\x1b" not in error
|
||||
assert "\\x1b" in error
|
||||
@@ -7,7 +7,7 @@ from typing import TYPE_CHECKING
|
||||
|
||||
from strix.config import loader
|
||||
from strix.config.settings import DedupeSettings
|
||||
from strix.report.dedupe import _dedupe_model_settings
|
||||
from strix.report.dedupe import _dedupe_model_settings, resolve_dedupe_model
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -16,32 +16,49 @@ if TYPE_CHECKING:
|
||||
import pytest
|
||||
|
||||
|
||||
def test_dedupe_key_sent_per_call_not_via_global_env() -> None:
|
||||
def _unwrap(model: object) -> object:
|
||||
while hasattr(model, "_inner"):
|
||||
model = model._inner
|
||||
return model
|
||||
|
||||
|
||||
def test_dedupe_key_bound_to_model_client_not_global_env() -> None:
|
||||
dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap", DEDUPE_LLM_API_KEY="dedupe-key")
|
||||
settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300)
|
||||
# The key rides on the request, so a shared-provider main key can't clobber
|
||||
# it (and vice versa) through the global provider env var.
|
||||
assert (settings.extra_args or {})["api_key"] == "dedupe-key"
|
||||
model = _unwrap(resolve_dedupe_model(dedupe, "deepseek/cheap"))
|
||||
# The key is bound to the dedupe model's own client, so a shared-provider
|
||||
# main key can't clobber it (and vice versa) through the process globals —
|
||||
# and it never rides on the request, where every model implementation's own
|
||||
# api_key kwarg would collide with it.
|
||||
assert model.api_key == "dedupe-key" # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_dedupe_settings_omit_api_key_when_unset() -> None:
|
||||
dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap")
|
||||
def test_dedupe_settings_carry_no_request_credentials() -> None:
|
||||
dedupe = DedupeSettings(
|
||||
STRIX_DEDUPE_MODEL="deepseek/cheap",
|
||||
DEDUPE_LLM_API_KEY="dedupe-key",
|
||||
DEDUPE_LLM_API_BASE="https://dedupe.example/v1",
|
||||
)
|
||||
settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300)
|
||||
assert "api_key" not in (settings.extra_args or {})
|
||||
assert "api_base" not in (settings.extra_args or {})
|
||||
|
||||
|
||||
def test_dedupe_endpoint_sent_per_call() -> None:
|
||||
def test_dedupe_endpoint_bound_to_model_client() -> None:
|
||||
dedupe = DedupeSettings(
|
||||
STRIX_DEDUPE_MODEL="openai/cheap",
|
||||
DEDUPE_LLM_API_KEY="dedupe-key",
|
||||
DEDUPE_LLM_API_BASE="https://dedupe.example/v1",
|
||||
)
|
||||
settings = _dedupe_model_settings(dedupe, "openai/cheap", 300)
|
||||
# A distinct dedupe endpoint rides on the request instead of the
|
||||
# process-wide base URL, so it can't clobber the main model's endpoint.
|
||||
assert (settings.extra_args or {})["api_base"] == "https://dedupe.example/v1"
|
||||
assert (settings.extra_args or {})["api_key"] == "dedupe-key"
|
||||
model = _unwrap(resolve_dedupe_model(dedupe, "openai/cheap"))
|
||||
client = model._client # type: ignore[attr-defined]
|
||||
assert client.api_key == "dedupe-key"
|
||||
assert str(client.base_url).startswith("https://dedupe.example/v1")
|
||||
|
||||
|
||||
def test_dedupe_without_credentials_uses_default_provider() -> None:
|
||||
dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap")
|
||||
model = _unwrap(resolve_dedupe_model(dedupe, "deepseek/cheap"))
|
||||
assert model.api_key is None # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def test_dedicated_dedupe_model_uses_own_headers_not_main() -> None:
|
||||
|
||||
@@ -458,6 +458,49 @@ async def test_non_user_send_does_not_resume_budget_pause(tmp_path: Any) -> None
|
||||
session.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_send_starts_fresh_resume_attempt_after_failure() -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
await coordinator.register("root", "strix", parent_id=None)
|
||||
await coordinator.register("child", "recon", parent_id="root")
|
||||
await coordinator.park_waiting("child", wait_kind="stalled")
|
||||
await coordinator.record_recovery("child")
|
||||
await coordinator.record_idle_resume("child")
|
||||
await coordinator.set_status("child", "failed", error="provider rejected request")
|
||||
assert await coordinator.claim_parent_notice("child") is True
|
||||
|
||||
delivered = await coordinator.send("child", {"from": "user", "content": "try again"})
|
||||
|
||||
assert delivered is True
|
||||
assert coordinator.statuses["child"] == "waiting"
|
||||
assert coordinator.pending_counts["child"] == 1
|
||||
assert "child" not in coordinator.errors
|
||||
assert "child" not in coordinator.wait_kinds
|
||||
assert "child" not in coordinator.recovery_counts
|
||||
assert "child" not in coordinator.idle_resume_counts
|
||||
assert await coordinator.claim_parent_notice("child") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_user_send_preserves_failed_resume_state() -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
await coordinator.register("root", "strix", parent_id=None)
|
||||
await coordinator.register("child", "recon", parent_id="root")
|
||||
await coordinator.park_waiting("child", wait_kind="stalled")
|
||||
await coordinator.record_recovery("child")
|
||||
await coordinator.record_idle_resume("child")
|
||||
await coordinator.set_status("child", "failed", error="provider rejected request")
|
||||
|
||||
delivered = await coordinator.send("child", {"from": "root", "content": "status"})
|
||||
|
||||
assert delivered is True
|
||||
assert coordinator.statuses["child"] == "failed"
|
||||
assert coordinator.errors["child"] == "provider rejected request"
|
||||
assert coordinator.wait_kinds["child"] == "stalled"
|
||||
assert coordinator.recovery_counts["child"] == 1
|
||||
assert coordinator.idle_resume_counts["child"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reset_budget_stops_clears_pause_and_normalizes_statuses() -> None:
|
||||
coordinator = AgentCoordinator()
|
||||
|
||||
@@ -857,6 +857,37 @@ async def test_agent_state_sync_does_not_mask_root_failure_with_completed_report
|
||||
assert runtime.controller.error == "finalization failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_state_sync_clears_root_failure_after_user_resume() -> None:
|
||||
runtime = GoTuiRuntime(args())
|
||||
await runtime.coordinator.register("root", "Strix", parent_id=None)
|
||||
await runtime.coordinator.set_status("root", "failed", error="provider rejected request")
|
||||
|
||||
await runtime._sync_agent_state()
|
||||
assert runtime.controller.scan_state == "failed"
|
||||
assert runtime.live_view.agents["root"]["error_message"] == "provider rejected request"
|
||||
|
||||
await runtime.coordinator.send("root", {"from": "user", "content": "try again"})
|
||||
await runtime._sync_agent_state()
|
||||
|
||||
assert runtime.controller.scan_state == "running"
|
||||
assert runtime.controller.error is None
|
||||
root = runtime.live_view.agents["root"]
|
||||
assert root["status"] == "waiting"
|
||||
assert "error_message" not in root
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_state_sync_does_not_reopen_stopped_scan_with_active_root() -> None:
|
||||
runtime = GoTuiRuntime(args())
|
||||
runtime.controller.scan_state = "stopped"
|
||||
await runtime.coordinator.register("root", "Strix", parent_id=None)
|
||||
|
||||
await runtime._sync_agent_state()
|
||||
|
||||
assert runtime.controller.scan_state == "stopped"
|
||||
|
||||
|
||||
def _direct_launch_args() -> argparse.Namespace:
|
||||
launch_args = args()
|
||||
launch_args.needs_setup = False
|
||||
|
||||
99
tests/test_import_warmup.py
Normal file
99
tests/test_import_warmup.py
Normal file
@@ -0,0 +1,99 @@
|
||||
"""The import warm-up thread must never leave the import system poisoned.
|
||||
|
||||
Field failure: the warm-up thread's ``strix.core.runner`` import and the main
|
||||
thread's ``strix.report`` import both walked the agents SDK graph, and the two
|
||||
held each other's import locks (report -> dedupe -> agents while runner ->
|
||||
hooks -> report.state). CPython's deadlock avoidance breaks such a cycle by
|
||||
failing one import, which strands finished submodules in ``sys.modules`` with
|
||||
their parent package gone — and the next import of one of those submodules
|
||||
crashes with "partially initialized module".
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
|
||||
from strix.llm import warmup
|
||||
|
||||
|
||||
def _run(code: str) -> subprocess.CompletedProcess[str]:
|
||||
return subprocess.run( # noqa: S603
|
||||
[sys.executable, "-c", textwrap.dedent(code)],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=False,
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
|
||||
def test_strix_report_does_not_import_the_agents_graph() -> None:
|
||||
result = _run(
|
||||
"""
|
||||
import sys
|
||||
|
||||
import strix.report
|
||||
|
||||
agents_modules = [m for m in sys.modules if m == "agents" or m.startswith("agents.")]
|
||||
assert not agents_modules, agents_modules
|
||||
assert "strix.report.dedupe" not in sys.modules
|
||||
"""
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
||||
|
||||
def test_check_duplicate_resolves_lazily() -> None:
|
||||
result = _run(
|
||||
"""
|
||||
import strix.report
|
||||
from strix.report import check_duplicate
|
||||
from strix.report.dedupe import check_duplicate as direct
|
||||
|
||||
assert strix.report.check_duplicate is direct is check_duplicate
|
||||
"""
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
||||
|
||||
def test_failed_warm_import_purges_orphaned_submodules() -> None:
|
||||
result = _run(
|
||||
"""
|
||||
import sys
|
||||
|
||||
from strix.llm.warmup import _warm
|
||||
|
||||
# A package whose import fails after a submodule already completed:
|
||||
# CPython removes the package but leaves the submodule stranded.
|
||||
import pathlib
|
||||
import tempfile
|
||||
|
||||
root = pathlib.Path(tempfile.mkdtemp())
|
||||
pkg = root / "stranded_pkg"
|
||||
pkg.mkdir()
|
||||
(pkg / "ok.py").write_text("VALUE = 1")
|
||||
(pkg / "__init__.py").write_text("from . import ok\\nraise RuntimeError('boom')")
|
||||
sys.path.insert(0, str(root))
|
||||
|
||||
_warm(("stranded_pkg",))
|
||||
|
||||
assert "stranded_pkg" not in sys.modules
|
||||
assert "stranded_pkg.ok" not in sys.modules, "orphan survived the purge"
|
||||
|
||||
# And the subtree imports cleanly afterwards up to the real error.
|
||||
try:
|
||||
import stranded_pkg # noqa: F401
|
||||
except RuntimeError:
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("expected the package's own error")
|
||||
"""
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
||||
|
||||
def test_purge_does_not_touch_preexisting_or_healthy_modules() -> None:
|
||||
before = frozenset(sys.modules) - {"strix.llm.warmup"}
|
||||
warmup._purge_orphaned_modules(before)
|
||||
assert "strix.llm.warmup" in sys.modules # parent chain intact -> kept
|
||||
assert "strix" in sys.modules
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user