diff --git a/AGENTS.md b/AGENTS.md index fdaae975..451d0e91 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -38,17 +38,22 @@ 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 - strix cloud login --scopes scans:read scans:write billing:read # device sign-in, no prompts + 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 --wait - strix cloud scans start --source . --dry-run --show-files --json # review local upload - strix cloud scans start --source . --yes --engagement-type code_review --wait + strix cloud scans start --source . --dry-run --show-files --json # review + capture source.archive_sha256 + SOURCE_SHA256="" + strix cloud scans start --source . --approve-sha256 "$SOURCE_SHA256" --wait strix cloud vulns list --severity critical - strix cloud billing topup --credits 20 # buy credits when a scan returns exit code 5 + strix cloud billing topup --credits 20 --yes # explicit approval after exit code 5 ``` - - Account setup runs from the CLI too: `strix cloud workspaces list|create|use`, `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 `. 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, the installation, or the DNS change for them. - - Every REST operation has a `strix cloud ` 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. Exit codes: `0` success, `1` error, `2` usage, `4` auth or plan limit, `5` payment required. `--token` or `STRIX_API_TOKEN` overrides the stored sign-in. `--data` adds extra request fields as JSON, and accepts `@file` or `-` for standard input. - - Local source uploads require `uploads:write`. Review with `scans start --source . --dry-run --show-files --json`, then approve with `--yes`. 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. + - 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 org members invite`, `strix cloud billing subscribe --plan strix_cloud`, `strix cloud billing portal`, `strix cloud integrations install github`, and `strix cloud domains verify `. 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 ` 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` overrides the stored sign-in. `--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`, 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). Managed API docs for LLMs: https://docs.app.strix.ai/llms.txt. diff --git a/README.md b/README.md index 29cbbba2..f5ba6541 100644 --- a/README.md +++ b/README.md @@ -325,13 +325,16 @@ strix auth logout # forget the sign-in 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 # opens the browser, then pick a workspace and scopes -strix cloud login --scopes scans:read scans:write # request specific token scopes -strix cloud login --workspace "My Team" # select a workspace by name or ID -strix cloud whoami # show the active sign-in -strix cloud logout # forget the sign-in +strix cloud login # browser approval, then workspace + scope picker +strix cloud login --workspace "My Team" # select a workspace by name or ID +strix cloud whoami # show the active sign-in and scopes +strix cloud logout # forget the sign-in ``` +The default **Recommended** scope preset supports normal scan work, local source uploads, +and workspace switching. For strict least privilege, pass an explicit list such as +`--scopes scans:read scans:write uploads:write billing:read`. + Every operation of the [REST API](https://docs.app.strix.ai) has a matching command in the form `strix cloud `: ```bash @@ -339,17 +342,20 @@ strix cloud # list all resources strix cloud scans # 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 --wait +strix cloud scans start --source . --dry-run --show-files --json # review + capture source.archive_sha256 +SOURCE_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 # buy credits (agent payment, HTTP 402) +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 # your workspaces +strix cloud workspaces list # numbered list; `workspace` is also accepted strix cloud workspaces create --name "My Team" -strix cloud workspaces use "My Team" # store a token for another workspace +strix cloud workspaces use 2 # switch by list number, exact name, or ID 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 @@ -358,7 +364,7 @@ strix cloud domains verify # 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 are agent friendly. Output is JSON when stdout is not a terminal or when you pass `--json`. There are no prompts when stdin is not a terminal. Exit codes: `0` success, `1` error, `2` invalid usage, `4` authentication required, `5` payment required. Set the token with `--token` or `STRIX_API_TOKEN` to skip the stored sign-in. +The commands work for humans and agents: terminal output favors names, branches, states, and numbered selectors, while redirected output (or `--json`) preserves complete machine-readable records and IDs. 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. Set the token with `--token` or `STRIX_API_TOKEN` to skip the stored sign-in. Write commands take request fields as flags, and every write command also accepts one JSON object with `--data`: @@ -368,7 +374,39 @@ strix cloud scans start --data @request.json # read a fi cat request.json | strix cloud scans start --data - # read standard input ``` -The platform enforces plan and role limits. Report downloads need the Enterprise plan, schedules need the Pro plan, and billing writes need an admin token. Those commands return the platform message and exit code `4`. +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, 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 diff --git a/docs/cloud/overview.mdx b/docs/cloud/overview.mdx index 2429c7e2..1dc5c69b 100644 --- a/docs/cloud/overview.mdx +++ b/docs/cloud/overview.mdx @@ -40,16 +40,19 @@ Skip the setup. Run Strix in the cloud at [app.strix.ai](https://app.strix.ai). Send a local working tree to the managed white-box scanner without connecting a source-control provider: ```bash -# Review the exact file manifest first. Nothing is uploaded. +# 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="" -# Approve the reviewed selection, upload it, and wait for the scan. -strix cloud scans start --source . --yes --engagement-type code_review --wait +# 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. Non-interactive environments must pass `--yes`, so an agent cannot upload a workspace without an explicit approval flag. +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`. Run your first pentest in minutes. diff --git a/skills/managed-pentesting-with-strix/SKILL.md b/skills/managed-pentesting-with-strix/SKILL.md index a8fb70a7..cc1f5411 100644 --- a/skills/managed-pentesting-with-strix/SKILL.md +++ b/skills/managed-pentesting-with-strix/SKILL.md @@ -16,7 +16,15 @@ There are two equivalent interfaces. Prefer the CLI: - **`strix cloud` CLI** — every REST operation has a command in the form `strix cloud `. Install with `curl -sSL https://strix.ai/install | bash`. Run `strix cloud` to list all resources and `strix cloud ` to list its verbs. - **REST API** — base URL `https://app.strix.ai/api/v1`, `Authorization: Bearer ` 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`. -The CLI is agent friendly. Output is JSON when stdout is not a terminal, or when you pass `--json`. There are no interactive prompts when stdin is not a terminal. Exit codes: `0` success, `1` error, `2` invalid usage, `4` authentication required, `5` payment required. +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, states, and numbered selectors. 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 has a useful default list/read 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: @@ -33,10 +41,12 @@ The platform enforces plan and role limits, and the CLI passes the platform mess 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 ``` -The user approves the sign-in in the browser. With `--scopes` (and optionally `--workspace `) there are no 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). +The user approves the sign-in in the browser. With `--scopes` (and optionally `--workspace `) 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, and workspace switching; use explicit scopes for a narrower automation token. - `strix cloud whoami` shows the active sign-in; add `--json` when another agent will parse it. `strix cloud logout` removes it. - Every other `strix cloud` command uses the stored token automatically. To use a different token (for example one created in the dashboard at **Settings → API Access**), pass `--token ` or set `STRIX_API_TOKEN`. @@ -51,14 +61,32 @@ The user approves the sign-in in the browser. With `--scopes` (and optionally `- | `schedules:read` / `:write` | read schedules · create/trigger recurring scans | | `pr_reviews:write` | trigger PR security reviews | | `webhooks:read` / `:write` | manage webhook subscriptions | + | `uploads:write` | upload local source or documents for a scan | + | `organizations:read` | list and switch workspaces | | `tokens:write` | create/revoke API tokens | + | `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 +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"}]' +``` + +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 -Scans consume org credits. Check the balance before a scan (`billing:read`): +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 @@ -66,13 +94,17 @@ 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 `mppx` client when Node.js is available — the user approves the spend in their agent wallet, for example the [Link Agent Wallet](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 # --yes skips the confirmation prompt +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 agent wallet. Tell the user to set it up one time at [link.com/agents](https://link.com/agents). After setup, the user approves each payment in the Link app, and no keys or variables are necessary. +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: @@ -88,13 +120,14 @@ An omitted `--monthly-cap-credits` keeps the stored cap. Pass `--no-monthly-cap` Manage workspaces with a personal token from `strix cloud login`: ```bash -strix cloud workspaces list # id, name, role, and the active one +strix cloud workspaces list # numbered name/role/current list strix cloud workspaces create --name "My Team" -strix cloud workspaces use "My Team" # store a token for that workspace +strix cloud workspaces use 2 # displayed number, exact name, or ID +strix cloud workspace use "My Team" # singular `workspace` alias also works strix cloud org members invite --email dev@example.com --role analyst ``` -`workspaces use` mints a new token for a workspace the user already belongs to, and the role in that workspace limits the scopes. Add `--scopes` to request a smaller set. +`workspaces use` rotates the current personal token to a workspace the user already belongs to and stores the updated workspace metadata; the bearer secret stays unchanged. The role in the target workspace limits the scopes. Add `--scopes` to request a smaller set. ### Handoffs a person must finish @@ -161,8 +194,8 @@ Useful flags (each maps to a `CreateScanRequest` field): | `--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 HTTP headers (API keys, for example) for the target (JSON map) | -| `--focus` / `--concerns` / `--context` | steer the agents | +| `--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 | @@ -170,20 +203,50 @@ The response is `{ scan_id, title, status }` with `status` = `pending`. ### Scan a local workspace in the cloud -Review the exact local upload before sending it. An agent or CI process must never skip this review merely because it can pass `--yes`: +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 -strix cloud scans start --source . --dry-run --show-files --json -strix cloud scans start --source . --yes --engagement-type code_review --wait +strix cloud scans start --source . --exclude 'private/' --dry-run --show-files --json +# After reviewing the output, capture its source.archive_sha256 value: +SOURCE_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 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, and enforces file-count, per-file, expanded-size, and compressed-size limits. Add project exclusions to `.strixignore` (one exclude glob per line) or repeat `--exclude GLOB`. +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. -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`. Non-interactive uploads require `--yes`; interactive terminals show one summary and confirmation prompt. The CLI deletes its temporary archive after the request and best-effort deletes the uploaded object if scan creation fails. +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 ` (`scans:read`). Status flow: `pending → running → completed` (or `failed` / `cancelled`). Scans take minutes to hours — poll on an interval, do not block. +Pass `--wait` to `scans start` to poll until the scan reaches a final state, or poll yourself with `strix cloud scans get ` (`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 @@ -214,6 +277,13 @@ strix cloud scans sarif --output findings.sarif strix cloud scans report --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`). 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: @@ -235,6 +305,8 @@ List/inspect with `strix cloud pr-reviews list` and `strix cloud pr-reviews get 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. diff --git a/strix/interface/cloud/__init__.py b/strix/interface/cloud/__init__.py index 7ce2ce81..4db4f13b 100644 --- a/strix/interface/cloud/__init__.py +++ b/strix/interface/cloud/__init__.py @@ -8,12 +8,19 @@ required. from __future__ import annotations -from rich.console import Console +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.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 [arguments] @@ -29,16 +36,46 @@ _USAGE_HEADER = """[bold]Usage:[/] strix cloud [arguments] _USAGE_FOOTER = """ Run [bold]strix cloud help[/] to list its verbs. Common read-only commands may also run their default verb when no verb is given. -Every command accepts [bold]--json[/] and [bold]--token[/]. Write commands -accept [bold]--data[/] with a JSON object of extra request fields. +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() - if not argv or argv[0] in ("-h", "--help", "help"): - _print_usage(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:] @@ -49,31 +86,40 @@ def run_cloud(argv: list[str]) -> int: if group == "credits": group, rest = "billing", ["credits", *rest] if group == "workspaces" and rest and rest[0] == "use": - return run_workspace_use(rest[1:]) + 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: - console.print(f"[red]Unknown command:[/] {group}") + 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 = bool(rest and rest[0] in ("-h", "--help", "help")) + group_help = _is_help_request(rest) resolved = None if group_help else resolve(group, rest) if resolved is None: - _print_verbs(console, group) - return 0 if not rest or rest[0] in ("-h", "--help", "help") else 2 + 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] in ("-h", "--help", "help"): - return run_login(["--help"]) - if group == "logout" and rest: - console.print("[red]Unknown argument for logout:[/] " + " ".join(rest)) - return 2 +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"], + "logout": ["logout", *rest], "whoami": ["status", *rest], } return run_login(session_argv[group]) @@ -86,9 +132,34 @@ def _print_usage(console: Console) -> None: console.print(_USAGE_FOOTER) -def _print_verbs(console: Console, group: str) -> None: +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", "credits"], + "resource_commands": [{"name": group, "help": GROUP_HELP.get(group, "")} for group in SPEC], + } + sys.stdout.write(json.dumps(payload, indent=2) + "\n") diff --git a/strix/interface/cloud/arguments.py b/strix/interface/cloud/arguments.py new file mode 100644 index 00000000..d1924572 --- /dev/null +++ b/strix/interface/cloud/arguments.py @@ -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, + ) diff --git a/strix/interface/cloud/billing.py b/strix/interface/cloud/billing.py new file mode 100644 index 00000000..a3392d5a --- /dev/null +++ b/strix/interface/cloud/billing.py @@ -0,0 +1,374 @@ +"""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 +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" +_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", + "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 + 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" + ) + if not payment_method and ( + not as_json + and not os.environ.get("MPPX_ACCOUNT") + and not os.environ.get("MPPX_STRIPE_SECRET_KEY") + ): + console.print( + "[dim]Tip: payments need a wallet. Set up a Stripe agent wallet at " + "https://link.com/agents, and the user approves each payment in the Link app. " + "If the user does not want a wallet, run " + "`strix cloud billing subscribe --plan strix_top_up` for a hosted checkout link.[/]" + ) + try: + wallet_result = _run_wallet_client( + npx, + args, + body, + token=token, + payment_method=payment_method, + 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: + if as_json: + emit(console, confirmed_receipt, as_json=True) + return http.EXIT_OK + + 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.[/]" + ) + return http.EXIT_PAYMENT + + stdout = str(getattr(result, "stdout", "") or "").strip() + stderr = str(getattr(result, "stderr", "") or "").strip() + 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( + npx: str, + args: argparse.Namespace, + body: dict[str, Any], + *, + token: str | None, + payment_method: str | None, + capture_output: bool, +) -> _WalletClientResult: + """Run mppx through the loopback bridge without exposing the API token to it.""" + 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) + with wallet_payment_bridge( + upstream_url=upstream_url, + api_token=http.api_token(token), + expected_body=body_json.encode(), + timeout=getattr(args, "timeout", None), + response_observer=upstream_responses.append, + ) as wallet_url: + command = [ + npx, + "--yes", + f"--registry={_NPM_REGISTRY}", + "--ignore-scripts", + f"--userconfig={user_config}", + f"--globalconfig={global_config}", + f"--cache={wallet_root / 'npm-cache'}", + _MPPX_PACKAGE, + wallet_url, + "--fail", + "-J", + body_json, + ] + if payment_method: + command += ["-M", f"paymentMethod={payment_method}"] + process = subprocess.run( # noqa: S603 + command, + check=False, + capture_output=capture_output, + text=True, + env=wallet_env, + cwd=wallet_root, + ) + return _WalletClientResult(process=process, upstream_responses=tuple(upstream_responses)) + + +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("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 diff --git a/strix/interface/cloud/http.py b/strix/interface/cloud/http.py index 14560350..56f3a4d3 100644 --- a/strix/interface/cloud/http.py +++ b/strix/interface/cloud/http.py @@ -2,8 +2,12 @@ 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 @@ -16,11 +20,15 @@ if TYPE_CHECKING: _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 _timeout_s: float = _DEFAULT_TIMEOUT_S EXIT_OK = 0 EXIT_ERROR = 1 +EXIT_USAGE = 2 EXIT_AUTH = 4 EXIT_PAYMENT = 5 @@ -34,19 +42,49 @@ class CloudError(Exception): self.payload = payload -def configure(*, base_url: str | None = None, timeout: float | None = None) -> None: +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, +) -> None: """Set the platform URL and the request timeout for this process.""" - global _app_url_override, _timeout_s # noqa: PLW0603 - if base_url: - _app_url_override = base_url.rstrip("/") - if timeout: + global _app_url_override, _timeout_s, _token_override_active # noqa: PLW0603 + _app_url_override = base_url.rstrip("/") if base_url else None + _token_override_active = token_override + 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 - return load_settings().viewer.app_url.rstrip("/") + 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() + ) + 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: @@ -56,6 +94,7 @@ def api_token(override: str | None = None) -> str: 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( @@ -65,6 +104,31 @@ def api_token(override: str | None = None) -> str: 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, @@ -72,25 +136,38 @@ def request( 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)}"} + if idempotency_key is not None: + headers["Idempotency-Key"] = idempotency_key try: response = requests.request( method, url, headers=headers, - params={k: v for k, v in (query or {}).items() if v is not None} or None, + 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 CloudError(f"could not reach {app_url()}: {exc}") from exc + raise CloudTransportError(f"could not reach {app_url()}: {exc}") from exc return response 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( @@ -101,19 +178,95 @@ def upload_file(signed_url: str, upload_token: str, path: Path) -> None: "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 - 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})") + 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.""" + target = _parse_origin_url(signed_url, label="source upload URL") + 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) -> 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 + 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: @@ -128,7 +281,7 @@ def parsed(response: requests.Response) -> Any: def check(response: requests.Response) -> Any: data = parsed(response) - if response.ok: + if 200 <= response.status_code < 300: content_type = response.headers.get("content-type", "").lower() if "application/json" not in content_type: raise CloudError( @@ -137,10 +290,19 @@ def check(response: requests.Response) -> Any: ) return data 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: diff --git a/strix/interface/cloud/payment_proxy.py b/strix/interface/cloud/payment_proxy.py new file mode 100644 index 00000000..36e31598 --- /dev/null +++ b/strix/interface/cloud/payment_proxy.py @@ -0,0 +1,280 @@ +"""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 two +requests (challenge and 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 = 2 +_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 + 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 request and its 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-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 + 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, + expected_body: bytes, + timeout: float | None = None, + response_observer: Callable[[WalletUpstreamResponse], None] | None = None, +) -> Generator[str, None, None]: + """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}", + 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) diff --git a/strix/interface/cloud/render.py b/strix/interface/cloud/render.py index 484ae4fe..f96aee42 100644 --- a/strix/interface/cloud/render.py +++ b/strix/interface/cloud/render.py @@ -3,20 +3,28 @@ from __future__ import annotations import json +import re import sys -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, TypeGuard from rich.markup import escape from rich.table import Table +from strix.interface.terminal_text import sanitize_terminal_text + if TYPE_CHECKING: + from collections.abc import Iterable + from rich.console import Console _MAX_TABLE_COLUMNS = 8 _MAX_CELL_LENGTH = 60 +_MAX_DETAIL_FIELDS = 36 +_MAX_NESTED_PREVIEW = 5 _NARROW_TABLE_WIDTH = 120 +_CAMEL_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])") _INTERNAL_COLUMNS = frozenset( { "organization_id", @@ -24,9 +32,22 @@ _INTERNAL_COLUMNS = frozenset( "userId", "installation_id", "created_by", + "connected_by", "avatarUrl", } ) +_LOSSLESS_DETAIL_KEYS = frozenset( + { + "api_token", + "command", + "docker_command", + "enrollment_command", + "secret", + "signing_secret", + "token", + "webhook_secret", + } +) _PREFERRED_KEYS = ( "name", @@ -57,6 +78,11 @@ _PREFERRED_KEYS = ( "url", "provider", "secret_prefix", + "events", + "action", + "resource_type", + "response_status", + "attempts", "scan_type", "engagement_type", "estimated_credits", @@ -71,13 +97,126 @@ _PREFERRED_KEYS = ( "id", ) +_LIST_ENVELOPE_KEYS = frozenset( + { + "items", + "data", + "scans", + "vulnerabilities", + "findings", + "domains", + "repositories", + "repos", + "schedules", + "reviews", + "pr_reviews", + "workspaces", + "members", + "invitations", + "integrations", + "connectors", + "webhooks", + "deliveries", + "entries", + "documents", + "policies", + "tokens", + "uploads", + "events", + "audit_logs", + "logs", + } +) +_ENVELOPE_METADATA_KEYS = frozenset( + { + "total", + "total_count", + "totalCount", + "count", + "page", + "limit", + "page_size", + "pageSize", + "has_more", + "hasMore", + "next_cursor", + "nextCursor", + "meta", + "pagination", + "summary", + "stats", + "scansThisMonth", + } +) + +_VIEW_COLUMNS: dict[str, tuple[str, ...]] = { + "GET /scans": ( + "title", + "target", + "engagement_type", + "scan_type", + "status", + "findings_count", + "created_at", + "id", + ), + "GET /vulnerabilities": ( + "display_number", + "title", + "severity", + "status", + "target", + "cvss", + "finding_type", + "id", + ), + "GET /pr-reviews": ( + "repository", + "pull_request", + "branches", + "status", + "verdict", + "findings", + "updated_at", + "id", + ), + "GET /integrations": ( + "provider", + "account_login", + "installation_id", + "instance_url", + "status", + "repository_selection", + "default_collection_name", + "connected_at", + ), + "GET /webhooks": ("url", "events", "is_active", "last_delivery_at", "created_at", "id"), + "GET /webhooks/{webhookId}/deliveries": ( + "event", + "status", + "response_status", + "attempts", + "created_at", + "delivered_at", + "id", + ), +} + + +def _is_record(value: object) -> TypeGuard[dict[str, Any]]: + return isinstance(value, dict) + + +def _is_list(value: object) -> TypeGuard[list[Any]]: + return isinstance(value, list) + def json_mode(*, flag: bool) -> bool: """JSON output is on when the flag is set or when stdout is not a terminal.""" return flag or not sys.stdout.isatty() -def emit( +def emit( # noqa: PLR0911, PLR0912 console: Console, data: Any, *, @@ -85,10 +224,71 @@ def emit( row_numbers: bool = False, omit_columns: frozenset[str] = frozenset(), hint: str | None = None, + view: str | None = None, + warning: str | None = None, ) -> None: if as_json: sys.stdout.write(json.dumps(data, indent=2, default=str) + "\n") return + if warning: + console.print(f"[bold yellow]Save this now:[/] {escape(sanitize_terminal_text(warning))}") + if view == "source_manifest" and _is_record(data): + _print_source_manifest(console, data) + return + if view == "GET /analytics/scan-frequency": + _print_scan_frequency(console, data) + return + if view in {"GET /analytics/overview", "GET /analytics/stats"} and _is_record(data): + _print_analytics(console, data) + return + if view == "GET /integrations": + integration_rows = _integration_rows(data) + if integration_rows is not None: + _print_table( + console, + integration_rows, + row_numbers=row_numbers, + omit_columns=omit_columns, + hint=hint, + view=view, + ) + return + if view == "GET /scans": + scan_rows = _scan_rows(data) + if scan_rows is not None: + _print_table( + console, + scan_rows, + row_numbers=False, + omit_columns=omit_columns, + hint="Inspect one scan with `strix cloud scans get ID`.", + view=view, + ) + return + if view == "GET /vulnerabilities": + vulnerability_rows = _list_of_dicts(data) + if vulnerability_rows is not None: + _print_table( + console, + vulnerability_rows, + row_numbers=False, + omit_columns=omit_columns | frozenset({"scan_id"}), + hint="Inspect one finding with `strix cloud vulns get ID`.", + view=view, + ) + return + if view == "GET /pr-reviews": + review_rows = _pr_review_rows(data) + if review_rows is not None: + _print_table( + console, + review_rows, + row_numbers=False, + omit_columns=omit_columns, + hint="Use `strix cloud pr-reviews get ID` for one review.", + view=view, + ) + return rows = _list_of_dicts(data) if rows is not None: _print_table( @@ -97,28 +297,143 @@ def emit( row_numbers=row_numbers, omit_columns=omit_columns, hint=hint, + view=view, ) return if isinstance(data, str): - console.print(data) + console.print(sanitize_terminal_text(data), markup=False) return - if isinstance(data, dict): + if _is_record(data): _print_detail(console, data) return console.print_json(json.dumps(data, default=str)) def _list_of_dicts(data: Any) -> list[dict[str, Any]] | None: - if isinstance(data, dict) and len(data) >= 1: - lists = [v for v in data.values() if isinstance(v, list)] - scalars = [v for v in data.values() if not isinstance(v, list | dict)] - if len(lists) == 1 and not scalars: - data = lists[0] - if not isinstance(data, list) or not data: + """Extract a record list from a raw list or a common paginated envelope.""" + if _is_record(data): + # Some endpoints wrap the actual envelope in a top-level ``data`` or + # ``result`` object. Only recurse through an object wrapper; a list in + # ``data`` is handled with the other named envelope keys below. + for wrapper in ("data", "result"): + nested = data.get(wrapper) + if _is_record(nested): + nested_rows = _list_of_dicts(nested) + if nested_rows is not None: + return nested_rows + candidates = [ + (key, value) + for key, value in data.items() + if key in _LIST_ENVELOPE_KEYS + and _is_list(value) + and all(_is_record(item) for item in value) + ] + if len(candidates) == 1: + list_key, records = candidates[0] + other_keys = set(data) - {list_key} + if list_key in {"items", "data"} or other_keys <= _ENVELOPE_METADATA_KEYS: + data = records + if not _is_list(data): return None - if not all(isinstance(item, dict) for item in data): + if not data: + return [] + records = [item for item in data if _is_record(item)] + if len(records) != len(data): return None - return data + return records + + +def _integration_rows(data: Any) -> list[dict[str, Any]] | None: + """Flatten the two integration collections into one compact human view.""" + if not _is_record(data): + return _list_of_dicts(data) + rows: list[dict[str, Any]] = [] + found_collection = False + for key in ("integrations", "merge_accounts"): + collection = data.get(key) + if not _is_list(collection): + continue + found_collection = True + for item in collection: + if not _is_record(item): + continue + row = dict(item) + if not row.get("account_login") and row.get("account_email"): + row["account_login"] = row["account_email"] + rows.append(row) + return rows if found_collection else None + + +def _scan_rows(data: Any) -> list[dict[str, Any]] | None: + """Flatten the nested target and finding summaries returned by scan lists.""" + records = _list_of_dicts(data) + if records is None: + return None + rows: list[dict[str, Any]] = [] + for record in records: + row = dict(record) + if not row.get("title") and row.get("name"): + row["title"] = row["name"] + if not row.get("id") and isinstance(row.get("scan_id"), str): + row["id"] = row["scan_id"] + + targets: list[str] = [] + urls = record.get("urls") + if _is_list(urls): + targets.extend(url.strip() for url in urls if isinstance(url, str) and url.strip()) + repositories = record.get("repositories") + if _is_list(repositories): + for repository in repositories: + if not _is_record(repository): + continue + identifier = str( + repository.get("full_name") + or repository.get("name") + or repository.get("url") + or "" + ).strip() + branch = str(repository.get("branch") or "").strip() + if identifier: + targets.append(f"{identifier} @ {branch}" if branch else identifier) + if targets: + visible_targets = targets[:2] + summary = " | ".join(visible_targets) + if len(targets) > len(visible_targets): + summary += f" (+{len(targets) - len(visible_targets)} more)" + row["target"] = summary + + findings = record.get("findings") + if _is_record(findings) and findings.get("total") is not None: + row["findings_count"] = findings["total"] + rows.append(row) + return rows + + +def _pr_review_rows(data: Any) -> list[dict[str, Any]] | None: + """Collapse related PR fields into an eight-column, action-oriented human view.""" + records = _list_of_dicts(data) + if records is None: + return None + rows: list[dict[str, Any]] = [] + for record in records: + row = dict(record) + number = record.get("pr_number") + title = str(record.get("pr_title") or "").strip() + row["repository"] = record.get("repository_full_name") or record.get("repository") + row["pull_request"] = " ".join( + part for part in (f"#{number}" if number is not None else "", title) if part + ) + head = str(record.get("head_branch") or "").strip() + base = str(record.get("base_branch") or "").strip() + row["branches"] = f"{head} → {base}" if head and base else head or base + total = record.get("findings_count") + opened = record.get("open_findings_count") + if isinstance(total, int) and isinstance(opened, int): + row["findings"] = f"{opened} open / {total} total" + elif isinstance(total, int): + row["findings"] = total + rows.append(row) + return rows def _print_table( @@ -128,12 +443,20 @@ def _print_table( row_numbers: bool = False, omit_columns: frozenset[str] = frozenset(), hint: str | None = None, + view: str | None = None, ) -> None: - omit_columns = omit_columns | _INTERNAL_COLUMNS + if not rows: + console.print("[dim]No items.[/]") + if hint: + console.print(f"[dim]{escape(sanitize_terminal_text(hint))}[/]") + return + integration_view = view == "GET /integrations" + visible_internal: set[str] = {"installation_id"} if integration_view else set() + view_omissions: set[str] = {"id"} if integration_view else set() + omit_columns = omit_columns | (_INTERNAL_COLUMNS - visible_internal) | view_omissions + preferred = _VIEW_COLUMNS.get(view or "", _PREFERRED_KEYS) columns: list[str] = [ - key - for key in _PREFERRED_KEYS - if key not in omit_columns and any(key in row for row in rows) + key for key in preferred if key not in omit_columns and any(key in row for row in rows) ] for row in rows: for key in row: @@ -149,22 +472,22 @@ def _print_table( _print_cards(console, rows, columns, row_numbers=row_numbers) console.print(f"[dim]{len(rows)} item(s). Use --json for the full records.[/]") if hint: - console.print(f"[dim]{hint}[/]") + console.print(f"[dim]{escape(sanitize_terminal_text(hint))}[/]") return table = Table(show_lines=False) if row_numbers: table.add_column("#", justify="right", style="cyan", no_wrap=True) for column in columns: - table.add_column(column) + table.add_column(escape(_human_label(column))) for index, row in enumerate(rows, start=1): - cells = [_cell(row.get(column)) for column in columns] + cells = [escape(_cell(row.get(column))) for column in columns] if row_numbers: cells.insert(0, str(index)) table.add_row(*cells) console.print(table) console.print(f"[dim]{len(rows)} item(s). Use --json for the full records.[/]") if hint: - console.print(f"[dim]{hint}[/]") + console.print(f"[dim]{escape(sanitize_terminal_text(hint))}[/]") def _print_cards( @@ -182,7 +505,13 @@ def _print_cards( if row.get(column) is not None ] prefix = f"[cyan]{index}.[/] " if row_numbers else "[cyan]•[/] " - console.print(prefix + " [dim]·[/] ".join(parts), soft_wrap=False) + if not parts: + console.print(prefix.rstrip()) + continue + console.print(prefix + parts[0], soft_wrap=True) + continuation = " " if row_numbers else " " + for part in parts[1:]: + console.print(continuation + part, soft_wrap=True) def _print_detail(console: Console, data: dict[str, Any]) -> None: @@ -192,20 +521,164 @@ def _print_detail(console: Console, data: dict[str, Any]) -> None: table = Table(show_header=False, show_edge=False, box=None, padding=(0, 2)) table.add_column("field", style="bold cyan", no_wrap=True) table.add_column("value", overflow="fold") - for key in keys: + populated_keys = [key for key in keys if data.get(key) is not None] + visible_keys = populated_keys[:_MAX_DETAIL_FIELDS] + lossless_fields: list[tuple[str, Any]] = [] + for key in visible_keys: value = data.get(key) - if value is None: + if _is_lossless_detail(key, value): + lossless_fields.append((key, value)) continue - if isinstance(value, (dict, list)): - rendered = json.dumps(value, indent=2, default=str) - else: - rendered = _cell(value) - table.add_row(_human_label(key), rendered) - console.print(table) + rendered = _nested_summary(value) if _is_record(value) or _is_list(value) else _cell(value) + table.add_row(escape(_human_label(key)), escape(rendered)) + if table.row_count: + console.print(table) + for key, value in lossless_fields: + console.print(f"{_human_label(key)}:", style="bold cyan", markup=False) + console.print(_lossless_detail_value(value), markup=False, soft_wrap=True) + if len(populated_keys) > len(visible_keys): + console.print( + f"[dim]{len(populated_keys) - len(visible_keys)} additional field(s) omitted from " + "this view.[/]" + ) console.print("[dim]Use --json for the lossless machine-readable record.[/]") +def _is_lossless_detail(key: str, value: Any) -> bool: + """Keep one-time credentials and enrollment commands complete and copyable.""" + sensitive_key = key in _LOSSLESS_DETAIL_KEYS or key.endswith(("_token", "_secret")) + return sensitive_key and not _is_record(value) and not _is_list(value) + + +def _lossless_detail_value(value: Any) -> str: + """Preserve structural newlines while making every other control byte visible.""" + return "\n".join(sanitize_terminal_text(line) for line in str(value).split("\n")) + + +def _nested_summary(value: dict[str, Any] | list[Any]) -> str: + """Bound nested records so one detail response cannot flood a terminal.""" + if _is_record(value): + scalar_items = [ + (nested_key, nested_value) + for nested_key, nested_value in value.items() + if not isinstance(nested_value, dict | list) and nested_value is not None + ] + lines = [ + f"{_human_label(str(nested_key))}: {_cell(nested_value)}" + for nested_key, nested_value in scalar_items[:_MAX_NESTED_PREVIEW] + ] + omitted = len(value) - len(lines) + if omitted > 0: + lines.append(f"… {omitted} more field(s)") + return "\n".join(lines) if lines else f"{len(value)} nested field(s)" + if not _is_list(value): + return "none" + if not value: + return "none" + if all(not isinstance(item, dict | list) for item in value): + preview = ", ".join(_cell(item) for item in value[:12]) + if len(value) > 12: + preview += f", … {len(value) - 12} more" + return preview + records = [item for item in value if _is_record(item)] + lines = [f"{len(value)} item(s)"] + for record in records[:_MAX_NESTED_PREVIEW]: + label = record.get("title") or record.get("name") or record.get("message") + severity = record.get("severity") + status = record.get("status") or record.get("state") + prefix = " / ".join(_cell(part) for part in (severity, status) if part) + summary = str(label or record.get("id") or "record") + lines.append(f"- {prefix + ': ' if prefix else ''}{_cell(summary)}") + if len(value) > len(records[:_MAX_NESTED_PREVIEW]): + lines.append(f"… {len(value) - len(records[:_MAX_NESTED_PREVIEW])} more; use --json") + return "\n".join(lines) + + +def _print_source_manifest(console: Console, data: dict[str, Any]) -> None: + source = data.get("source") + manifest = source if _is_record(source) else data + files = manifest.get("files") + summary = {key: value for key, value in manifest.items() if key != "files"} + _print_detail(console, summary) + if _is_list(files): + console.print(f"\n[bold]Selected files ({len(files):,})[/]") + for path in files: + console.print(f" {escape(sanitize_terminal_text(path))}", soft_wrap=True) + + +def _print_analytics(console: Console, data: dict[str, Any]) -> None: + rows = list(_flatten_summary(data)) + table = Table(show_header=False, show_edge=False, box=None, padding=(0, 2)) + table.add_column("metric", style="bold cyan") + table.add_column("value", overflow="fold") + for label, value in rows[:_MAX_DETAIL_FIELDS]: + table.add_row(escape(label), escape(value)) + console.print(table) + if len(rows) > _MAX_DETAIL_FIELDS: + console.print( + f"[dim]Showing {_MAX_DETAIL_FIELDS} of {len(rows)} summary metrics. " + "Use --json for all data.[/]" + ) + else: + console.print("[dim]Use --json for the complete analytics record.[/]") + + +def _flatten_summary(value: Any, prefix: str = "", depth: int = 0) -> Iterable[tuple[str, str]]: + if _is_record(value) and depth < 4: + for key, nested in value.items(): + label = f"{prefix} / {_human_label(key)}" if prefix else _human_label(key) + yield from _flatten_summary(nested, label, depth + 1) + return + if _is_list(value): + if all(not isinstance(item, dict | list) for item in value): + yield prefix, _nested_summary(value) + else: + yield prefix, f"{len(value)} data point(s)" + return + yield prefix or "value", _cell(value) + + +def _print_scan_frequency(console: Console, data: Any) -> None: + rows = _find_record_series(data) + if rows is None: + if _is_record(data): + _print_analytics(console, data) + else: + console.print_json(json.dumps(data, default=str)) + return + nonzero = [row for row in rows if _row_has_activity(row)] + selected = (nonzero[-30:] if nonzero else rows[-14:]) if rows else [] + _print_table(console, selected, view="GET /analytics/scan-frequency") + if rows: + qualifier = "non-zero" if nonzero else "most recent" + console.print( + f"[dim]Showing {len(selected)} {qualifier} point(s) from {len(rows)} total. " + "Use --json for the full series.[/]" + ) + + +def _find_record_series(data: Any) -> list[dict[str, Any]] | None: + direct = _list_of_dicts(data) + if direct is not None: + return direct + if _is_record(data): + candidates = [ + series for value in data.values() if (series := _find_record_series(value)) is not None + ] + if candidates: + return max(candidates, key=len) + return None + + +def _row_has_activity(row: dict[str, Any]) -> bool: + count_keys = ("count", "scans", "scan_count", "total", "value") + return any(isinstance(row.get(key), int | float) and row[key] > 0 for key in count_keys) + + def _human_label(column: str) -> str: + column = sanitize_terminal_text(column) + if column == "secret_prefix": + return "prefix" labels = { "repository_full_name": "repo", "pr_number": "PR", @@ -219,9 +692,8 @@ def _human_label(column: str) -> str: "updated_at": "updated", "expires_at": "expires", "last_used_at": "last used", - "secret_prefix": "prefix", } - return labels.get(column, column.replace("_", " ")) + return labels.get(column, _CAMEL_BOUNDARY.sub(" ", column).replace("_", " ").lower()) def _cell(value: Any) -> str: @@ -229,7 +701,13 @@ def _cell(value: Any) -> str: return "" if isinstance(value, bool): return "yes" if value else "no" - text = str(value) + if _is_list(value) and all(not _is_record(item) and not _is_list(item) for item in value): + text = ", ".join(str(item) for item in value) + elif _is_record(value) or _is_list(value): + text = f"{len(value)} item(s)" + else: + text = str(value) + text = sanitize_terminal_text(text) if len(text) > _MAX_CELL_LENGTH: return text[: _MAX_CELL_LENGTH - 1] + "…" return text diff --git a/strix/interface/cloud/runner.py b/strix/interface/cloud/runner.py index 4b32629c..717ac0fa 100644 --- a/strix/interface/cloud/runner.py +++ b/strix/interface/cloud/runner.py @@ -8,29 +8,41 @@ from __future__ import annotations import argparse import json +import math import os import re -import shutil -import subprocess import sys +import tempfile import time import webbrowser from contextlib import suppress from pathlib import Path -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast from urllib.parse import quote +from uuid import uuid4 +import requests from rich.console import Console +from rich.markup import escape -from strix.interface.cloud import http +import strix.interface.cloud.http as http # noqa: PLR0402 +from strix.interface.cloud.arguments import CloudArgumentParser +from strix.interface.cloud.billing import run_topup from strix.interface.cloud.render import emit, json_mode -from strix.interface.cloud.source_upload import SourceBundle, prepare_source, remove_bundle +from strix.interface.cloud.source_scan import LocalSourceScan from strix.interface.cloud.spec import DEFAULT_VERBS, SPEC, Cmd, P +from strix.interface.terminal_text import sanitize_terminal_text +from strix.interface.url_safety import is_safe_web_url + + +if TYPE_CHECKING: + from collections.abc import Iterator _PLACEHOLDER = re.compile(r"\{([^{}]+)\}") _CAMEL_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])") _WAIT_POLL_S = 15 +_DEFAULT_WAIT_TIMEOUT_S = 4 * 60 * 60 _TERMINAL_STATUSES = frozenset( { "completed", @@ -43,6 +55,12 @@ _TERMINAL_STATUSES = frozenset( "succeeded", } ) +_DEFINITIVE_SCAN_REJECTION_STATUSES = frozenset({400, 401, 402, 403, 404, 409, 413, 415, 422}) +_IDEMPOTENCY_KEY = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/=-]{0,199}$") +_IDEMPOTENCY_RETRY_DELAYS_S = (0.25, 1.0) +_RETRYABLE_IDEMPOTENCY_CODES = frozenset( + {"idempotency_request_in_progress", "idempotency_outcome_unknown"} +) def _dest(name: str) -> str: @@ -53,6 +71,82 @@ def _metavar(name: str) -> str: return _CAMEL_BOUNDARY.sub("_", name).upper() +def _positive_seconds(value: str) -> float: + try: + parsed = float(value) + except ValueError as exc: + raise argparse.ArgumentTypeError("must be a number greater than 0") from exc + if not math.isfinite(parsed) or parsed <= 0: + raise argparse.ArgumentTypeError("must be a finite number greater than 0") + return parsed + + +def _resolve_idempotency_key(cmd: Cmd, args: argparse.Namespace) -> str | None: + if not cmd.idempotent: + return None + supplied = getattr(args, "idempotency_key", None) + key = supplied if isinstance(supplied, str) else str(uuid4()) + if not _IDEMPOTENCY_KEY.fullmatch(key): + raise http.CloudError( + "--idempotency-key must be 1-200 characters, start with a letter or digit, and " + "contain only letters, digits, '.', '_', ':', '/', '=', or '-'.", + exit_code=http.EXIT_USAGE, + ) + return key + + +def _audit_export_format(cmd: Cmd, query: dict[str, Any]) -> str | None: + if cmd.method != "GET" or cmd.path != "/audit": + return None + value = query.get("format") + if not isinstance(value, str): + return None + normalized = value.strip().lower() + return normalized if normalized in {"csv", "ndjson", "jsonl", "snowflake", "splunk"} else None + + +def _contains_response_key(value: Any, keys: frozenset[str], *, depth: int = 0) -> bool: + if depth > 2: + return False + if isinstance(value, dict): + fields = cast("dict[str, Any]", value) + if any(key in fields and fields[key] not in (None, "") for key in keys): + return True + return any(_contains_response_key(item, keys, depth=depth + 1) for item in fields.values()) + return False + + +def _one_time_secret_warning(cmd: Cmd, args: argparse.Namespace, result: Any) -> str | None: + if ( + cmd.method == "POST" + and cmd.path == "/tokens" + and _contains_response_key(result, frozenset({"token", "api_token", "secret"})) + ): + return "This API token is shown only once. Store it securely before leaving this output." + if cmd.path.startswith("/webhooks") and _contains_response_key( + result, frozenset({"secret", "signing_secret", "webhook_secret"}) + ): + return ( + "This webhook signing secret is shown only once. Store it securely before leaving " + "this output." + ) + connector_command_requested = cmd.method == "POST" or bool( + getattr(args, "include_command", False) + ) + if ( + cmd.path.startswith("/connectors") + and connector_command_requested + and _contains_response_key( + result, frozenset({"command", "enrollment_command", "docker_command", "token"}) + ) + ): + return ( + "This connector enrollment command contains one-time credentials. Store it securely " + "and do not share it." + ) + return None + + def resolve(group: str, tokens: list[str]) -> tuple[Cmd, list[str]] | None: """Find the command for a verb. Two-word verbs match before one-word verbs.""" commands = SPEC.get(group) @@ -73,8 +167,21 @@ def resolve(group: str, tokens: list[str]) -> tuple[Cmd, list[str]] | None: def run(group: str, verb_label: str, cmd: Cmd, argv: list[str]) -> int: console = Console() parser = _build_parser(group, verb_label, cmd) + as_json = json_mode(flag="--json" in argv) + raw_binary_stdout = _argv_uses_raw_binary_stdout(cmd, argv) try: args = parser.parse_args(argv) + except KeyboardInterrupt: + _emit_interrupted(console, as_json=as_json, to_stderr=raw_binary_stdout) + return 130 + except http.CloudError as exc: + _emit_error( + console, + exc, + as_json=as_json and not raw_binary_stdout, + to_stderr=raw_binary_stdout, + ) + return exc.exit_code except SystemExit as exc: return exc.code if isinstance(exc.code, int) else 2 @@ -84,24 +191,145 @@ def run(group: str, verb_label: str, cmd: Cmd, argv: list[str]) -> int: path = path.replace("{" + name + "}", value) as_json = json_mode(flag=bool(getattr(args, "json", False))) + raw_binary_stdout = _uses_raw_binary_stdout(cmd, args) token = getattr(args, "token", None) - http.configure(base_url=getattr(args, "app_url", None), timeout=getattr(args, "timeout", None)) - try: + http.configure( + base_url=getattr(args, "app_url", None), + timeout=getattr(args, "timeout", None), + token_override=bool(token), + ) query = _collect(args, cmd.query) body = _collect(args, cmd.body) data = getattr(args, "data", None) if data: - body.update(_load_data(data)) + _merge_extra_body(body, _load_data(data)) + _validate_body(cmd, body) if getattr(args, "no_monthly_cap", False): body["monthly_cap_credits"] = None return _execute(console, cmd, args, path, query, body, as_json=as_json, token=token) + except KeyboardInterrupt: + _emit_interrupted(console, as_json=as_json, to_stderr=raw_binary_stdout) + return 130 except http.CloudError as exc: - _emit_error(console, exc, as_json=as_json) + _emit_error( + console, + exc, + as_json=as_json and not raw_binary_stdout, + to_stderr=raw_binary_stdout, + ) return exc.exit_code -def _execute( # noqa: PLR0912 +def _uses_raw_binary_stdout(cmd: Cmd, args: argparse.Namespace) -> bool: + if bool(getattr(args, "json", False)) or getattr(args, "output", None): + return False + format_value = str(getattr(args, "format", "") or "").strip().lower() + audit_export = ( + cmd.method == "GET" + and cmd.path == "/audit" + and format_value + in { + "csv", + "ndjson", + "jsonl", + "snowflake", + "splunk", + } + ) + return not sys.stdout.isatty() and bool(cmd.binary or audit_export) + + +def _argv_uses_raw_binary_stdout(cmd: Cmd, argv: list[str]) -> bool: + """Choose the error channel before argparse can reject a binary command.""" + if "--json" in argv or any(arg.startswith("--json=") for arg in argv): + return False + has_output = any( + (arg.startswith("--output=") and bool(arg.partition("=")[2])) + or (arg == "--output" and index + 1 < len(argv) and not argv[index + 1].startswith("-")) + for index, arg in enumerate(argv) + ) + if has_output: + return False + format_value = "" + for index, arg in enumerate(argv): + if arg.startswith("--format="): + format_value = arg.partition("=")[2] + elif arg == "--format" and index + 1 < len(argv): + format_value = argv[index + 1] + audit_export = ( + cmd.method == "GET" + and cmd.path == "/audit" + and format_value.lower() + in { + "csv", + "ndjson", + "jsonl", + "snowflake", + "splunk", + } + ) + return not sys.stdout.isatty() and bool(cmd.binary or audit_export) + + +def _request_with_idempotency( + cmd: Cmd, + path: str, + *, + token: str | None, + query: dict[str, Any], + body: dict[str, Any], + stream: bool, + idempotency_key: str | None, +) -> requests.Response: + """Retry only exact, caller-keyed mutations whose outcome may be ambiguous.""" + attempts = 1 + (len(_IDEMPOTENCY_RETRY_DELAYS_S) if idempotency_key else 0) + for attempt in range(attempts): + try: + response = http.request( + cmd.method, + path, + token=token, + query=query or None, + body=body if cmd.method in ("POST", "PUT", "PATCH") else None, + stream=stream, + idempotency_key=idempotency_key, + ) + except http.CloudTransportError: + if attempt + 1 >= attempts: + raise + else: + if attempt + 1 >= attempts or not _idempotency_response_is_retryable(response): + return response + response.close() + time.sleep(_IDEMPOTENCY_RETRY_DELAYS_S[attempt]) + raise AssertionError("idempotent request retry loop exhausted without returning") + + +def _idempotency_response_is_retryable(response: requests.Response) -> bool: + if 500 <= response.status_code < 600 or response.status_code == 429: + return True + if response.status_code != 409: + return False + payload = http.parsed(response) + if not isinstance(payload, dict): + return False + fields = cast("dict[str, Any]", payload) + return fields.get("retry_safe") is True or fields.get("code") in _RETRYABLE_IDEMPOTENCY_CODES + + +def _scan_rejection_is_definitive(response: requests.Response) -> bool: + payload = http.parsed(response) + if isinstance(payload, dict): + fields = cast("dict[str, Any]", payload) + if fields.get("retry_safe") is True: + return False + if fields.get("terminal") is True: + return True + return response.status_code in _DEFINITIVE_SCAN_REJECTION_STATUSES + + +def _execute( # noqa: PLR0912, PLR0915 console: Console, cmd: Cmd, args: argparse.Namespace, @@ -113,85 +341,146 @@ def _execute( # noqa: PLR0912 token: str | None, ) -> int: if cmd.path == "/billing/topup": - return _topup(console, args, body, as_json=as_json, token=token) - source_bundle: SourceBundle | None = None - source_upload_id: str | None = None + return run_topup(console, args, body, as_json=as_json, token=token) + audit_export = _audit_export_format(cmd, query) + output_path = getattr(args, "output", None) + explicit_json = bool(getattr(args, "json", False)) + binary_response = bool(cmd.binary or audit_export) + if binary_response and explicit_json and not output_path: + raise http.CloudError( + "--json for a binary response requires --output FILE; omit --json only when " + "intentionally redirecting the raw bytes.", + exit_code=http.EXIT_USAGE, + ) + binary_json_metadata = explicit_json or bool(output_path and not sys.stdout.isatty()) + if cmd.path == "/audit" and getattr(args, "output", None) and not audit_export: + raise http.CloudError( + "--output requires --format csv, ndjson, jsonl, snowflake, or splunk.", + exit_code=http.EXIT_USAGE, + ) + idempotency_key = _resolve_idempotency_key(cmd, args) + source_workflow = LocalSourceScan(idempotency_key=idempotency_key) + scan_request_started = False try: if cmd.path == "/scans" and cmd.method == "POST": _set_default_scan_engagement( body, has_local_source=getattr(args, "source", None) is not None, ) - source_bundle = _prepare_scan_source(console, args, as_json=as_json) - if source_bundle is not None and getattr(args, "dry_run", False): - emit( - console, - { - "source": source_bundle.summary( - show_files=getattr(args, "show_files", False) - ) - }, - as_json=as_json, - ) + if source_workflow.prepare_and_attach( + console, + args, + body, + as_json=as_json, + token=token, + ): return http.EXIT_OK - if source_bundle is not None: - source_upload_id = _upload_scan_source(source_bundle, token=token) - existing = body.get("upload_ids") - body["upload_ids"] = [ - *(existing if isinstance(existing, list) else []), - source_upload_id, - ] - response = http.request( - cmd.method, + source_workflow.mark_launch_started() + scan_request_started = cmd.idempotent + response = _request_with_idempotency( + cmd, path, token=token, - query=query or None, - body=body if cmd.method in ("POST", "PUT", "PATCH") else None, + query=query, + body=body, + stream=bool(cmd.binary or audit_export), + idempotency_key=idempotency_key, ) - except BaseException: - if source_upload_id is not None: - _delete_upload(source_upload_id, token=token) + except BaseException as exc: + source_workflow.handle_request_failure(exc, token=token) + if scan_request_started: + if isinstance(exc, KeyboardInterrupt): + raise _interrupted_scan_launch_error(idempotency_key) from None + if isinstance(exc, Exception): + raise _ambiguous_scan_launch_error(exc, idempotency_key) from exc raise finally: - if source_bundle is not None: - remove_bundle(source_bundle) + source_workflow.close() + if audit_export: + return _emit_binary( + console, + response, + output_path, + force=bool(getattr(args, "force", False)), + json_metadata=binary_json_metadata, + ) if cmd.binary: - return _emit_binary(console, response, getattr(args, "output", None)) + return _emit_binary( + console, + response, + output_path, + force=bool(getattr(args, "force", False)), + json_metadata=binary_json_metadata, + ) try: result = http.check(response) - except BaseException: - if source_upload_id is not None: - _delete_upload(source_upload_id, token=token) + except BaseException as exc: + source_workflow.handle_response_failure( + exc, + definitive=_scan_rejection_is_definitive(response), + token=token, + ) + if ( + source_workflow.upload_id is None + and scan_request_started + and not _scan_rejection_is_definitive(response) + and isinstance(exc, Exception) + ): + raise _ambiguous_scan_launch_error(exc, idempotency_key) from exc raise if getattr(args, "wait", False): - if cmd.wait_self: - result = _poll(console, path, token=token, as_json=as_json) - elif cmd.wait_path: - result = _wait(console, cmd, result, token=token, as_json=as_json) - if source_bundle is not None: - result = { - "source": source_bundle.summary(show_files=getattr(args, "show_files", False)), - "upload_id": source_upload_id, - "scan": result, - } + wait_timeout = cast("float", getattr(args, "wait_timeout", _DEFAULT_WAIT_TIMEOUT_S)) + try: + if cmd.wait_self: + result = _poll( + console, + path, + token=token, + as_json=as_json, + wait_timeout=wait_timeout, + ) + elif cmd.wait_path: + result = _wait( + console, + cmd, + result, + token=token, + as_json=as_json, + wait_timeout=wait_timeout, + ) + except KeyboardInterrupt: + raise _wait_status_error(result, interrupted=True) from None + except http.CloudError as exc: + raise _wait_status_error(result, error=exc) from exc + result = source_workflow.wrap_result(result, args) if cmd.link: return _handoff_link(console, cmd, args, result, as_json=as_json) workspace_list = cmd.method == "GET" and cmd.path == "/workspaces" + integration_list = cmd.method == "GET" and cmd.path == "/integrations" emit( console, result, as_json=as_json, - row_numbers=workspace_list, + row_numbers=workspace_list or integration_list, omit_columns=frozenset({"id"}) if workspace_list else frozenset(), - hint=("Switch with `strix cloud workspaces use NUMBER`." if workspace_list else None), + hint=( + "Switch with `strix cloud workspaces use NUMBER`." + if workspace_list + else ( + "For Git providers, disconnect with `strix cloud integrations disconnect " + "PROVIDER --installation-id INSTALLATION_ID`; omit the ID for Slack." + if integration_list + else None + ) + ), + view=f"{cmd.method} {cmd.path}", + warning=_one_time_secret_warning(cmd, args, result), ) return http.EXIT_OK -def _set_default_scan_engagement( - body: dict[str, Any], *, has_local_source: bool = False -) -> None: +def _set_default_scan_engagement(body: dict[str, Any], *, has_local_source: bool = False) -> None: """Infer the scan type from its targets when the caller did not choose one.""" if body.get("engagement_type"): return @@ -203,6 +492,31 @@ def _set_default_scan_engagement( body["engagement_type"] = "code_review" +def _validate_body(cmd: Cmd, body: dict[str, Any]) -> None: + missing = [ + "--" + (param.flag or param.name.replace("_", "-")) + for param in cmd.body + if param.required and body.get(param.name) is None + ] + if missing: + raise http.CloudError( + "missing required request field(s): " + + ", ".join(missing) + + ". Supply them as options or with --data @file/-.", + exit_code=http.EXIT_USAGE, + ) + if ( + cmd.method == "POST" + and cmd.path == "/tokens" + and body.get("expires_at") is not None + and body.get("expires_in_days") is not None + ): + raise http.CloudError( + "--expires-at and --expires-in-days are mutually exclusive.", + exit_code=http.EXIT_USAGE, + ) + + def _handoff_link( console: Console, cmd: Cmd, args: argparse.Namespace, result: Any, *, as_json: bool ) -> int: @@ -210,13 +524,22 @@ def _handoff_link( fields = cast("dict[str, Any]", result) if isinstance(result, dict) else {} url = fields.get(cmd.link) if cmd.link else None if not isinstance(url, str) or not url: - emit(console, result, as_json=as_json) - return http.EXIT_OK - interactive = sys.stdout.isatty() and not getattr(args, "no_browser", False) + raise http.CloudError( + f"the platform response did not include the expected {cmd.link or 'continuation'} URL." + ) + if not is_safe_web_url(url, trusted_origin=http.app_url()): + raise http.CloudError("the platform returned an invalid continuation URL.") + interactive = ( + not as_json + and sys.stdin.isatty() + and sys.stdout.isatty() + and not getattr(args, "no_browser", False) + ) if as_json: emit(console, result, as_json=True) else: - console.print(f"Open this URL to continue:\n [bold]{url}[/]") + console.print("Open this URL to continue:") + console.print(f" {sanitize_terminal_text(url)}", markup=False, soft_wrap=True) if interactive: webbrowser.open(url) return http.EXIT_OK @@ -237,25 +560,45 @@ def _load_data(value: str) -> dict[str, Any]: try: parsed_value = json.loads(text) except ValueError as exc: - raise http.CloudError("--data must be a JSON object.") from exc + raise http.CloudError("--data must be a JSON object.", exit_code=http.EXIT_USAGE) from exc if not isinstance(parsed_value, dict): - raise http.CloudError("--data must be a JSON object.") + raise http.CloudError("--data must be a JSON object.", exit_code=http.EXIT_USAGE) return cast("dict[str, Any]", parsed_value) +def _merge_extra_body(body: dict[str, Any], extra_body: dict[str, Any]) -> None: + collisions = sorted(body.keys() & extra_body.keys()) + if collisions: + flags = ", ".join(f"--{name.replace('_', '-')}" for name in collisions) + raise http.CloudError( + f"--data cannot override explicit option(s): {flags}", + exit_code=http.EXIT_USAGE, + ) + body.update(extra_body) + + def _build_parser(group: str, verb_label: str, cmd: Cmd) -> argparse.ArgumentParser: - parser = argparse.ArgumentParser(prog=f"strix cloud {group} {verb_label}", description=cmd.help) + parser = CloudArgumentParser(prog=f"strix cloud {group} {verb_label}", description=cmd.help) for name in _PLACEHOLDER.findall(cmd.path): parser.add_argument(_dest(name), metavar=_metavar(name)) - for param in cmd.query + cmd.body: - _add_option(parser, param) - parser.add_argument("--json", action="store_true", help="Print the raw JSON response.") + for param in cmd.query: + _add_option(parser, param, required=param.required) + for param in cmd.body: + # Required body fields may be supplied securely through --data @file/-; + # validate them only after the two body sources are merged. + _add_option(parser, param, required=False) + json_help = "Print the raw JSON response." + if cmd.binary: + json_help = "With --output, print structured download metadata as JSON." + elif cmd.path == "/audit": + json_help = "Print JSON results, or download metadata when exporting with --output." + parser.add_argument("--json", action="store_true", help=json_help) parser.add_argument("--token", default=None, help="API token override.") parser.add_argument("--app-url", default=None, metavar="URL", help="Platform URL override.") parser.add_argument( "--timeout", default=None, - type=float, + type=_positive_seconds, metavar="SECONDS", help="Request timeout in seconds.", ) @@ -266,14 +609,25 @@ def _build_parser(group: str, verb_label: str, cmd: Cmd) -> argparse.ArgumentPar metavar="JSON", help="JSON object with extra request fields. Use @file to read a file, or - for stdin.", ) + _add_idempotency_option(parser, cmd) if cmd.path == "/billing/auto-topup" and cmd.method == "PUT": parser.add_argument( "--no-monthly-cap", action="store_true", help="Remove the monthly cap. Omit this flag to keep the stored cap.", ) - if cmd.binary: - parser.add_argument("--output", default=None, metavar="FILE", help="Write to this file.") + if cmd.binary or cmd.path == "/audit": + output_help = ( + "Write the CSV or NDJSON-compatible export to this file." + if cmd.path == "/audit" and not cmd.binary + else "Write to this file." + ) + parser.add_argument("--output", default=None, metavar="FILE", help=output_help) + parser.add_argument( + "--force", + action="store_true", + help="Replace --output if it already exists.", + ) if cmd.link: parser.add_argument( "--no-browser", @@ -284,11 +638,27 @@ def _build_parser(group: str, verb_label: str, cmd: Cmd) -> argparse.ArgumentPar parser.add_argument( "--wait", action="store_true", help="Wait until the operation reaches a final state." ) - if cmd.path == "/billing/topup": parser.add_argument( - "--yes", action="store_true", help="Do not ask for confirmation before payment." + "--wait-timeout", + type=_positive_seconds, + default=float(_DEFAULT_WAIT_TIMEOUT_S), + metavar="SECONDS", + help=( + "Maximum total time to wait before returning an error " + f"(default: {_DEFAULT_WAIT_TIMEOUT_S})." + ), ) - parser.add_argument( + if cmd.path == "/billing/topup": + payment_mode = parser.add_mutually_exclusive_group() + payment_mode.add_argument( + "--yes", + action="store_true", + help=( + "Explicitly authorize payment without a TTY prompt. Required in " + "non-interactive mode." + ), + ) + payment_mode.add_argument( "--no-pay", action="store_true", help="Print the payment challenge instead of paying it.", @@ -314,10 +684,20 @@ def _build_parser(group: str, verb_label: str, cmd: Cmd) -> argparse.ArgumentPar action="store_true", help="Build and print the source manifest without uploading or starting a scan.", ) - parser.add_argument( + source_approval = parser.add_mutually_exclusive_group() + source_approval.add_argument( "--yes", action="store_true", - help="Approve the displayed source upload without an interactive prompt.", + help="Approve the source snapshot built by this invocation without a prompt.", + ) + source_approval.add_argument( + "--approve-sha256", + default=None, + metavar="SHA256", + help=( + "Upload only if the archive exactly matches this --dry-run SHA-256 digest. " + "Best for agent and CI approval handoffs." + ), ) parser.add_argument( "--show-files", @@ -349,105 +729,114 @@ def _build_parser(group: str, verb_label: str, cmd: Cmd) -> argparse.ArgumentPar return parser -def _prepare_scan_source( - console: Console, args: argparse.Namespace, *, as_json: bool -) -> SourceBundle | None: - source = getattr(args, "source", None) - source_flags = ( - "dry_run", - "show_files", - "include_hidden", - "include_sensitive", - "include_archives", +def _add_idempotency_option(parser: argparse.ArgumentParser, cmd: Cmd) -> None: + if not cmd.idempotent: + return + parser.add_argument( + "--idempotency-key", + default=None, + metavar="KEY", + help=( + "Stable key for an exact retry after an ambiguous response. A fresh UUID is " + "generated when omitted; never reuse a key for a different request." + ), ) - 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", [])), + + +def _wait_status_error( + result: Any, + *, + error: http.CloudError | None = None, + interrupted: bool = False, +) -> http.CloudError: + operation_id = _created_id(result) + suffix = f" Operation ID: {operation_id}." if operation_id else "" + prefix = ( + "Interrupted while waiting" + if interrupted + else f"Waiting for the remote operation failed: {error}" ) - if getattr(args, "dry_run", False): - return bundle - if getattr(args, "yes", False): - return bundle - if as_json or not (sys.stdin.isatty() and sys.stdout.isatty()): - remove_bundle(bundle) - raise http.CloudError( - "source upload requires explicit approval in non-interactive mode. " - "Review with --dry-run --show-files, then rerun with --yes." - ) - 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." + message = ( + f"{prefix}; the remote operation may still be running.{suffix} " + "Check its status before retrying." ) - answer = console.input("Upload this source and start the scan? [y/N]: ").strip().lower() - if answer not in ("y", "yes"): - remove_bundle(bundle) - raise http.CloudError("source upload cancelled.") - return bundle - - -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", - }, - ) + payload: dict[str, Any] = { + "error": message, + "status_unknown": True, + } + if interrupted: + payload["interrupted"] = True + if operation_id: + payload["operation_id"] = operation_id + return http.CloudError( + message, + exit_code=130 if interrupted else (error.exit_code if error else http.EXIT_ERROR), + payload=payload, ) - 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)): - raise http.CloudError("the platform did not return complete source upload credentials.") - try: - http.upload_file(cast("str", signed_url), cast("str", upload_token), bundle.archive_path) - http.check( - http.request( - "POST", - "/uploads/complete", - token=token, - body={"upload_id": upload_id}, - ) - ) - except BaseException: - _delete_upload(cast("str", upload_id), token=token) - raise - return cast("str", upload_id) -def _delete_upload(upload_id: str, *, token: str | None) -> None: - with suppress(http.CloudError): - http.request("DELETE", f"/uploads/{quote(upload_id, safe='')}", token=token) +def _interrupted_scan_launch_error(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; check " + f"`strix cloud scans list` before retrying.{retry_note}" + ) + payload: dict[str, Any] = { + "error": message, + "interrupted": True, + "launch_outcome_unknown": True, + } + _attach_idempotency_recovery(payload, idempotency_key) + return http.CloudError( + message, + exit_code=130, + payload=payload, + ) -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" +def _ambiguous_scan_launch_error( + error: Exception, idempotency_key: str | None = None +) -> http.CloudError: + retry_note = _idempotency_retry_note(idempotency_key) + message = ( + f"{error} The scan launch outcome is unknown; check `strix cloud scans list` before " + f"retrying to avoid a duplicate scan or charge.{retry_note}" + ) + payload: dict[str, Any] = { + "error": message, + "launch_outcome_unknown": True, + } + exit_code = http.EXIT_ERROR + if isinstance(error, http.CloudError): + exit_code = error.exit_code + if isinstance(error.payload, dict): + payload.update(cast("dict[str, Any]", error.payload)) + payload["error"] = message + _attach_idempotency_recovery(payload, idempotency_key) + return http.CloudError(message, exit_code=exit_code, payload=payload) -def _add_option(parser: argparse.ArgumentParser, param: P) -> None: +def _idempotency_retry_note(idempotency_key: str | None) -> str: + if not idempotency_key: + return "" + return ( + f" An exact retry is safe with the same request and `--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 _add_option(parser: argparse.ArgumentParser, param: P, *, required: bool) -> None: flag = "--" + (param.flag or param.name.replace("_", "-")) if param.kind == "bool": parser.add_argument( @@ -455,7 +844,7 @@ def _add_option(parser: argparse.ArgumentParser, param: P) -> None: dest=param.name, action=argparse.BooleanOptionalAction, default=None, - required=param.required, + required=required, help=param.help, ) elif param.kind == "list": @@ -464,7 +853,7 @@ def _add_option(parser: argparse.ArgumentParser, param: P) -> None: dest=param.name, nargs="+", default=None, - required=param.required, + required=required, help=param.help, ) elif param.kind in ("int", "float"): @@ -473,13 +862,11 @@ def _add_option(parser: argparse.ArgumentParser, param: P) -> None: dest=param.name, type=int if param.kind == "int" else float, default=None, - required=param.required, + required=required, help=param.help, ) else: - parser.add_argument( - flag, dest=param.name, default=None, required=param.required, help=param.help - ) + parser.add_argument(flag, dest=param.name, default=None, required=required, help=param.help) def _collect(args: argparse.Namespace, params: tuple[P, ...]) -> dict[str, Any]: @@ -488,34 +875,164 @@ def _collect(args: argparse.Namespace, params: tuple[P, ...]) -> dict[str, Any]: value = getattr(args, param.name, None) if value is None: continue - if param.kind == "json" and isinstance(value, str): + if param.kind in ("json", "json-list") and isinstance(value, str): try: value = json.loads(value) except ValueError as exc: - raise http.CloudError(f"--{param.name.replace('_', '-')} must be JSON") from exc + raise http.CloudError( + f"--{param.name.replace('_', '-')} must be JSON", + exit_code=http.EXIT_USAGE, + ) from exc + if param.kind == "json-list" and not isinstance(value, list): + raise http.CloudError( + f"--{param.name.replace('_', '-')} must be a JSON array", + exit_code=http.EXIT_USAGE, + ) values[param.name] = value return values -def _emit_binary(console: Console, response: Any, output: str | None) -> int: - if not response.ok: - http.check(response) - if output: - Path(output).write_bytes(response.content) - console.print(f"Saved to [bold]{output}[/]") +def _emit_binary( + console: Console, + response: Any, + output: str | None, + *, + force: bool = False, + json_metadata: bool = False, +) -> int: + try: + if not 200 <= response.status_code < 300: + http.check(response) + if output: + return _write_binary_file( + console, + response, + Path(output).expanduser(), + force=force, + as_json=json_metadata, + ) + if sys.stdout.isatty(): + raise http.CloudError( + "binary responses require --output FILE when stdout is a terminal; " + "redirect stdout only when intentionally piping the bytes.", + exit_code=http.EXIT_USAGE, + ) + output_stream: Any = getattr(sys.stdout, "buffer", None) + try: + for chunk in _response_chunks(response): + if output_stream is not None: + output_stream.write(chunk) + else: + sys.stdout.write(chunk.decode("utf-8")) + except (OSError, UnicodeDecodeError, requests.RequestException) as exc: + raise http.CloudError(f"could not write the response to stdout: {exc}") from exc return http.EXIT_OK - sys.stdout.buffer.write(response.content) - return http.EXIT_OK + finally: + close = getattr(response, "close", None) + if callable(close): + with suppress(Exception): + close() -def _emit_error(console: Console, exc: http.CloudError, *, as_json: bool) -> None: +def _write_binary_file( + console: Console, response: Any, path: Path, *, force: bool, as_json: bool +) -> int: + if path.exists() and not force: + raise http.CloudError( + f"refusing to replace existing file {path}; pass --force to overwrite it." + ) + temporary: Path | None = None + bytes_written = 0 + try: + try: + path.parent.mkdir(parents=True, exist_ok=True) + with tempfile.NamedTemporaryFile( + mode="wb", + prefix=f".{path.name}.", + suffix=".tmp", + dir=path.parent, + delete=False, + ) as stream: + temporary = Path(stream.name) + for chunk in _response_chunks(response): + stream.write(chunk) + bytes_written += len(chunk) + except (OSError, requests.RequestException) as exc: + raise http.CloudError(f"could not write {path}: {exc}") from exc + + try: + if force: + temporary.replace(path) + else: + os.link(temporary, path) + temporary.unlink() + except FileExistsError as exc: + raise http.CloudError( + f"refusing to replace existing file {path}; pass --force to overwrite it." + ) from exc + except OSError as exc: + raise http.CloudError(f"could not write {path}: {exc}") from exc + + if as_json: + content_type = str(getattr(response, "headers", {}).get("content-type", "")) + emit( + console, + { + "output": str(path), + "bytes": bytes_written, + **({"content_type": content_type} if content_type else {}), + }, + as_json=True, + view="binary_download", + ) + else: + console.print("Saved to:") + console.print(sanitize_terminal_text(path), markup=False, soft_wrap=True) + return http.EXIT_OK + finally: + if temporary is not None: + temporary.unlink(missing_ok=True) + + +def _response_chunks(response: Any) -> Iterator[bytes]: + iter_content = getattr(response, "iter_content", None) + if callable(iter_content): + chunks = cast("Iterator[bytes]", iter_content(chunk_size=1024 * 1024)) + for chunk in chunks: + if chunk: + yield bytes(chunk) + return + content = getattr(response, "content", b"") + if content: + yield bytes(content) + + +def _emit_error( + console: Console, exc: http.CloudError, *, as_json: bool, to_stderr: bool = False +) -> None: if as_json: - payload = {"error": str(exc)} - if exc.payload is not None: - payload["detail"] = exc.payload + if isinstance(exc.payload, dict): + error_payload = cast("dict[str, Any]", exc.payload) + payload = dict(error_payload) + payload.setdefault("error", str(exc)) + if payload.get("detail") == payload.get("error"): + payload.pop("detail", None) + else: + payload = {"error": str(exc)} + if exc.payload is not None: + payload["detail"] = exc.payload sys.stdout.write(json.dumps(payload, indent=2, default=str) + "\n") return - console.print(f"[red]Error:[/] {exc}") + target = Console(stderr=True) if to_stderr else console + target.print(f"[red]Error:[/] {escape(sanitize_terminal_text(exc))}") + + +def _emit_interrupted(console: Console, *, as_json: bool, to_stderr: bool) -> None: + if as_json and not to_stderr: + sys.stdout.write(json.dumps({"error": "Interrupted.", "interrupted": True}) + "\n") + return + target = Console(stderr=True) if to_stderr else console + target.print("[yellow]Interrupted.[/]") def _created_id(created: Any) -> str | None: @@ -529,90 +1046,56 @@ def _created_id(created: Any) -> str | None: return None -def _wait(console: Console, cmd: Cmd, created: Any, *, token: str | None, as_json: bool) -> Any: +def _wait( + console: Console, + cmd: Cmd, + created: Any, + *, + token: str | None, + as_json: bool, + wait_timeout: float, +) -> Any: item_id = _created_id(created) if not item_id or not cmd.wait_path: return created path = cmd.wait_path.replace("{id}", str(item_id)) if not as_json: - console.print(f"[dim]Waiting for {item_id} to reach a final state…[/]") - return _poll(console, path, token=token, as_json=as_json) + console.print( + f"[dim]Waiting for {escape(sanitize_terminal_text(item_id))} to reach a final state…[/]" + ) + return _poll( + console, + path, + token=token, + as_json=as_json, + wait_timeout=wait_timeout, + ) -def _poll(console: Console, path: str, *, token: str | None, as_json: bool) -> Any: +def _poll( + console: Console, + path: str, + *, + token: str | None, + as_json: bool, + wait_timeout: float, +) -> Any: """Poll a GET path until its status is final. Returns the last response.""" + deadline = time.monotonic() + wait_timeout while True: - time.sleep(_WAIT_POLL_S) current: Any = http.check(http.request("GET", path, token=token)) fields = cast("dict[str, Any]", current) if isinstance(current, dict) else {} status = str(fields.get("status", "")) if status.lower() in _TERMINAL_STATUSES: return fields if isinstance(current, dict) else current if not as_json: - console.print(f"[dim] status: {status or 'unknown'}[/]") - - -def _topup( - console: Console, - args: argparse.Namespace, - body: dict[str, Any], - *, - as_json: bool, - token: str | None, -) -> int: - 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 - - npx = shutil.which("npx") - if npx is None: - emit(console, challenge, as_json=as_json) - console.print( - "[yellow]Payment required.[/] Install Node.js and run the command again, " - "or pay the challenge above with an MPP wallet client." - ) - return http.EXIT_PAYMENT - - credit_count = body.get("credits") - if not getattr(args, "yes", False) and sys.stdin.isatty(): - 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 - - url = f"{http.app_url()}/api/v1/billing/topup" - auth_header = f"X-Strix-Authorization: Bearer {http.api_token(token)}" - command = [ - npx, - "--yes", - "mppx", - url, - "-J", - json.dumps(body), - "-H", - auth_header, - ] - payment_method = getattr(args, "payment_method", None) or os.environ.get( - "MPPX_STRIPE_PAYMENT_METHOD" - ) - if payment_method: - command += ["-M", f"paymentMethod={payment_method}"] - elif not os.environ.get("MPPX_ACCOUNT") and not os.environ.get("MPPX_STRIPE_SECRET_KEY"): - console.print( - "[dim]Tip: payments need a wallet. Set up a Stripe agent wallet at " - "https://link.com/agents, and the user approves each payment in the Link app. " - "If the user does not want a wallet, run " - "`strix cloud billing subscribe --plan strix_top_up` for a hosted checkout link.[/]" - ) - result = subprocess.run(command, check=False) # noqa: S603 - return http.EXIT_OK if result.returncode == 0 else http.EXIT_PAYMENT + console.print( + f"[dim] status: {escape(sanitize_terminal_text(status or 'unknown'))}[/]" + ) + remaining = deadline - time.monotonic() + if remaining <= 0: + raise http.CloudError( + f"wait timed out after {wait_timeout:g} seconds; the remote operation is still " + "running. Re-run its get command to check the status." + ) + time.sleep(min(_WAIT_POLL_S, remaining)) diff --git a/strix/interface/cloud/source_scan.py b/strix/interface/cloud/source_scan.py new file mode 100644 index 00000000..408c10d1 --- /dev/null +++ b/strix/interface/cloud/source_scan.py @@ -0,0 +1,395 @@ +"""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 ; 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) + http.check( + http.request( + "POST", + "/uploads/complete", + token=token, + body={"upload_id": 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 _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 + if isinstance(error.payload, dict): + error_payload = cast("dict[str, Any]", error.payload) + payload.update(error_payload) + elif error.payload is not None: + payload["detail"] = error.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" diff --git a/strix/interface/cloud/source_upload.py b/strix/interface/cloud/source_upload.py index dc9e801c..bda3a0c8 100644 --- a/strix/interface/cloud/source_upload.py +++ b/strix/interface/cloud/source_upload.py @@ -13,14 +13,27 @@ import zipfile from collections import Counter from dataclasses import dataclass from pathlib import Path, PurePosixPath +from typing import TYPE_CHECKING -from strix.interface.cloud import http +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( { @@ -30,11 +43,20 @@ _ALWAYS_EXCLUDED_DIRS = frozenset( "node_modules", "vendor", "venv", + ".venv", + "env", "__pycache__", + ".tox", + ".pytest_cache", + ".mypy_cache", + ".ruff_cache", "dist", "build", "coverage", "target", + ".next", + ".nuxt", + ".gradle", } ) _SENSITIVE_NAMES = frozenset( @@ -50,6 +72,8 @@ _SENSITIVE_NAMES = frozenset( ".npmrc", ".pypirc", ".netrc", + ".git-credentials", + "application_default_credentials.json", } ) _SENSITIVE_PATTERNS = ( @@ -63,6 +87,15 @@ _SENSITIVE_PATTERNS = ( "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", @@ -75,6 +108,22 @@ _ARCHIVE_SUFFIXES = ( ".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", ) @@ -83,6 +132,10 @@ class SelectedFile: path: Path archive_name: str size: int + device: int + inode: int + mtime_ns: int + ctime_ns: int @dataclass(frozen=True) @@ -190,8 +243,14 @@ def select_source( 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): + for relative in _candidate_paths( + source, + include_hidden=include_hidden, + patterns=patterns, + excluded=excluded, + ): archive_name = relative.as_posix() reason = _exclusion_reason( relative, @@ -212,11 +271,24 @@ def select_source( 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, archive_name, info.st_size)) + 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( @@ -242,41 +314,171 @@ def remove_bundle(bundle: SourceBundle) -> None: bundle.archive_path.unlink(missing_ok=True) -def _candidate_paths(source: Path) -> list[Path]: +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 None: - return [path.relative_to(source) for path in source.rglob("*")] - 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()) - result = subprocess.run( # noqa: S603 # nosec B603 - command, check=False, capture_output=True + 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, ) - if result.returncode == 0: - paths: list[Path] = [] - for raw in result.stdout.split(b"\0"): - if not raw: + 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 - repo_relative = Path(os.fsdecode(raw)) try: - paths.append(repo_relative.relative_to(relative_source)) - except ValueError: - continue - return paths - return [path.relative_to(source) for path in source.rglob("*")] + 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: @@ -306,22 +508,24 @@ def _exclusion_reason( # noqa: PLR0911 patterns: list[str], ) -> str | None: parts = relative.parts - if any(part == ".git" for part in 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 parts[:-1]): + 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" - posix = PurePosixPath(relative.as_posix()) - if any( - posix.match(pattern) or fnmatch.fnmatch(relative.as_posix(), pattern) - for pattern in patterns - ): + 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): @@ -329,6 +533,27 @@ def _exclusion_reason( # noqa: PLR0911 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 @@ -341,7 +566,14 @@ def _write_archive(destination: Path, files: tuple[SelectedFile, ...]) -> None: 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: + 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." @@ -350,7 +582,30 @@ def _write_archive(destination: Path, files: tuple[SelectedFile, ...]) -> None: info.compress_type = zipfile.ZIP_DEFLATED info.external_attr = 0o100644 << 16 with archive.open(info, "w", force_zip64=True) as target: - shutil.copyfileobj(source_file, target, length=1024 * 1024) + 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: @@ -361,14 +616,29 @@ def _sha256(path: Path) -> str: 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" - try: - lines = path.read_text(encoding="utf-8").splitlines() - except FileNotFoundError: + raw_text = _read_ignore_file(path) + if raw_text is None: return [] - except OSError as exc: - raise http.CloudError(f"could not read {path}: {exc}") from exc + 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() @@ -379,4 +649,55 @@ def _load_ignore_patterns(source: Path) -> list[str]: 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.") diff --git a/strix/interface/cloud/spec.py b/strix/interface/cloud/spec.py index fbbabbf1..1ec85c36 100644 --- a/strix/interface/cloud/spec.py +++ b/strix/interface/cloud/spec.py @@ -14,7 +14,8 @@ from dataclasses import dataclass class P: """One command parameter. - ``kind`` is one of ``str``, ``int``, ``float``, ``bool``, ``list``, or ``json``. + ``kind`` is one of ``str``, ``int``, ``float``, ``bool``, ``list``, ``json``, + or ``json-list``. """ name: str @@ -40,6 +41,9 @@ class Cmd: # checkout page. The runner opens the browser for an interactive terminal # and always prints the URL. link: str | None = None + # Caller retries for this mutation must carry one stable opaque key. The + # platform binds it to the authenticated actor and exact request body. + idempotent: bool = False def _q(*names: str) -> tuple[P, ...]: @@ -47,14 +51,21 @@ def _q(*names: str) -> tuple[P, ...]: _SCAN_START_BODY = ( - P("engagement_type", help="Test category, for example live_test or code_review."), + P( + "engagement_type", + help=("Test category: code_review, live_test, internal_infra, or compliance_pentest."), + ), P("domain_ids", "list", help="Domain asset IDs to test."), P("domain_paths", "json", help="JSON map of domain ID to start paths."), P("repository_ids", "list", help="Repository asset IDs to test."), P("repository_branches", "json", help="JSON map of repository ID to branch."), P("credentials", "json", help="JSON list of credential objects."), - P("headers", "json", help="JSON map of extra HTTP headers for the target."), - P("concerns", "list", help="Vulnerability classes to focus on."), + P( + "headers", + "json", + help='JSON array of target header objects: [{"name":"...","value":"...","notes":"..."}].', + ), + P("concerns", help="Free-form security concerns to investigate."), P("focus", help="Free-form focus instructions for the agents."), P("context", help="Extra context about the target."), P("upload_ids", "list", help="Upload IDs to attach to the scan."), @@ -63,9 +74,15 @@ _SCAN_START_BODY = ( P("org_knowledge_enabled", "bool", help="Use the organization knowledge base."), P("notify_on_completion", "bool", help="Send an email when the scan completes."), P("notification_emails", "list", help="Extra notification email addresses."), - P("scan_tier", help="Scan tier, for example lite, pro, or max."), - P("model_config_id", help="Model configuration ID to run with."), - P("max_budget_usd", "float", help="Budget limit for the scan in USD."), + P( + "scan_tier", + help=( + "Scan tier: lite, standard, or ultra (default). " + "Not used for Enterprise or self-hosted scans." + ), + ), + P("model_config_id", help="Self-hosted only: model configuration ID to run with."), + P("max_budget_usd", "float", help="Self-hosted only: budget limit for the scan in USD."), ) _TEST_USER_ADD_BODY = ( @@ -105,15 +122,21 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/scans", "List scans.", - query=_q( - "status", - "scan_type", - "date_from", - "date_to", - "domain_id", - "repository_id", - "search", - "include_retests", + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + *_q( + "status", + "scan_type", + "date_from", + "date_to", + "domain_id", + "repository_id", + "search", + ), + P("include_retests", "bool", help="Include per-finding retest scans."), + P("sort_by", help="Sort key; currently created_at."), + P("sort_order", help="Sort order: asc or desc."), ), ), "start": Cmd( @@ -122,6 +145,7 @@ SPEC: dict[str, dict[str, Cmd]] = { "Start a scan.", body=_SCAN_START_BODY, wait_path="/scans/{id}", + idempotent=True, ), "get": Cmd("GET", "/scans/{scanId}", "Get one scan."), "delete": Cmd("DELETE", "/scans/{scanId}", "Delete a scan."), @@ -132,8 +156,11 @@ SPEC: dict[str, dict[str, Cmd]] = { "/scans/{scanId}/message", "Send a message to the scan agents.", body=( - P("message", required=True, help="Message text for the agents."), - P("cancel_current", "bool", help="Stop the current task first."), + P( + "message", + help="Message text for the agents. Required unless --cancel-current is used.", + ), + P("cancel_current", "bool", help="Cancel the current task before delivery."), P("agent_id", help="Target one agent instead of the root agent."), ), ), @@ -141,10 +168,53 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/scans/{scanId}/report", "Download the scan report.", - query=_q("format", "type"), + query=( + P( + "format", + help=( + "Report content: technical (default), retest, attestation, or " + "executive_summary. Advanced formats require Enterprise." + ), + ), + P( + "type", + help="Rendered file type: pdf (default) or docx. DOCX requires Enterprise.", + ), + P( + "providerName", + flag="provider-name", + help="Enterprise report-cover provider name (up to 80 characters).", + ), + P( + "memberName0", + flag="member-name-0", + help="First Enterprise report preparer's name (up to 120 characters).", + ), + P( + "memberEmail0", + flag="member-email-0", + help="First Enterprise report preparer's email address.", + ), + P( + "memberName1", + flag="member-name-1", + help="Second Enterprise report preparer's name (up to 120 characters).", + ), + P( + "memberEmail1", + flag="member-email-1", + help="Second Enterprise report preparer's email address.", + ), + ), binary=True, ), - "rerun": Cmd("POST", "/scans/{scanId}/rerun", "Run the scan again.", wait_path=None), + "rerun": Cmd( + "POST", + "/scans/{scanId}/rerun", + "Run the scan again.", + wait_path="/scans/{id}", + idempotent=True, + ), "retest-all": Cmd( "POST", "/scans/{scanId}/retest-all", @@ -195,19 +265,24 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/vulnerabilities", "List vulnerabilities.", - query=_q( - "scan_id", - "severity", - "status", - "search", - "from", - "to", - "domain_id", - "repository_id", - "finding_type", - "dependency_relation", - "reachability", - "sort_by", + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + *_q( + "scan_id", + "severity", + "status", + "search", + "from", + "to", + "domain_id", + "repository_id", + "finding_type", + "dependency_relation", + "reachability", + "sort_by", + ), + P("sort_order", help="Sort order: asc or desc."), ), ), "get": Cmd("GET", "/vulnerabilities/{vulnerabilityId}", "Get one vulnerability."), @@ -232,6 +307,7 @@ SPEC: dict[str, dict[str, Cmd]] = { "/vulnerabilities/{vulnerabilityId}/retest", "Retest one vulnerability.", body=(P("upload_ids", "list", help="Upload IDs with updated code."),), + wait_path="/scans/{id}", ), "fix-pr": Cmd( "POST", @@ -263,7 +339,12 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/domains", "List domain assets.", - query=_q("limit", "search", "verified", "business_unit", "tags", "sort_by"), + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + *_q("search", "verified", "business_unit", "tags", "sort_by"), + P("sort_order", help="Sort order: asc or desc."), + ), ), "add": Cmd( "POST", @@ -316,8 +397,11 @@ SPEC: dict[str, dict[str, Cmd]] = { "test-users provision-inbox": Cmd( "POST", "/domains/{domainId}/test-users/provision-inbox", - "Create a test user with a managed email inbox.", - body=(P("label", help="Display label for the test user."),), + ( + "Provision a Strix-managed inbox for email OTP or magic-link MFA. " + "Returns an address; it does not create a test user." + ), + body=(P("label", help="Optional display label for the managed inbox."),), ), "test-users inbox": Cmd( "GET", @@ -348,7 +432,12 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/repositories", "List repository assets.", - query=_q("limit", "search", "business_unit", "tags", "sort_by"), + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + *_q("search", "business_unit", "tags", "sort_by"), + P("sort_order", help="Sort order: asc or desc."), + ), ), "add": Cmd("POST", "/repositories", "Add a repository asset. Use --data for the fields."), "update": Cmd( @@ -386,18 +475,20 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/repositories/{repositoryId}/supply-chain/components", "List the dependency components of a repository.", - query=_q( - "job_id", - "snapshot_id", - "component_id", - "ecosystem", - "status", - "relationship", - "source_file", - "q", - "changed", - "limit", - "offset", + query=( + *_q( + "job_id", + "snapshot_id", + "component_id", + "ecosystem", + "status", + "relationship", + "source_file", + "q", + "changed", + ), + P("limit", "int", help="Maximum components to return."), + P("offset", "int", help="Number of components to skip."), ), ), "supply-chain sbom": Cmd( @@ -425,7 +516,12 @@ SPEC: dict[str, dict[str, Cmd]] = { }, "schedules": { "list": Cmd("GET", "/schedules", "List scan schedules."), - "create": Cmd("POST", "/schedules", "Create a scan schedule. Use --data for the fields."), + "create": Cmd( + "POST", + "/schedules", + "Create a scan schedule. Use --data for the fields.", + idempotent=True, + ), "get": Cmd("GET", "/schedules/{scheduleId}", "Get one schedule."), "update": Cmd( "PATCH", @@ -436,32 +532,48 @@ SPEC: dict[str, dict[str, Cmd]] = { P("cron_expression", help="Cron expression for the schedule."), P("timezone", help="Time zone for the cron expression."), P("name", help="Display name of the schedule."), - P("max_budget_usd", "int", help="Budget limit per run in USD."), - P("scan_tier", help="Scan tier for scheduled runs."), + P( + "max_budget_usd", + "float", + help=( + "Self-hosted only: budget limit per run in USD. " + "Use --data to set null and clear it." + ), + ), + P("scan_tier", help="Scan tier: lite, standard, or ultra."), ), ), "delete": Cmd("DELETE", "/schedules/{scheduleId}", "Delete a schedule."), "template": Cmd( "GET", "/schedules/{scheduleId}/template", "Get the schedule configuration template." ), - "trigger": Cmd("POST", "/schedules/{scheduleId}/trigger", "Run a schedule now."), + "trigger": Cmd( + "POST", + "/schedules/{scheduleId}/trigger", + "Run a schedule now.", + idempotent=True, + ), }, "pr-reviews": { "list": Cmd( "GET", "/pr-reviews", "List PR reviews.", - query=_q( - "search", - "status", - "group", - "pr_state", - "repository_full_name", - "date_from", - "date_to", - "sort_by", - "sort_order", - "include_counts", + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + *_q( + "search", + "status", + "group", + "pr_state", + "repository_full_name", + "date_from", + "date_to", + "sort_by", + "sort_order", + ), + P("include_counts", "bool", help="Include exact disposition counts."), ), ), "get": Cmd("GET", "/pr-reviews/{prReviewId}", "Get one PR review."), @@ -469,7 +581,12 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/pr-reviews/findings", "List PR review findings.", - query=_q("severity", "pr_state", "search", "repository_full_name", "include_stats"), + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + *_q("severity", "pr_state", "search", "repository_full_name"), + P("include_stats", "bool", help="Include all-time impact statistics."), + ), ), "start": Cmd( "POST", @@ -545,7 +662,11 @@ SPEC: dict[str, dict[str, Cmd]] = { "Start a chat session.", body=( P("message", required=True, help="First message of the session."), - P("repos", "list", help="Repository full names for context."), + P( + "repos", + "json", + help='JSON array of repository refs: [{"repoId":"...","branch":"main"}].', + ), P("domain_ids", "list", help="Domain asset IDs for context."), ), ), @@ -555,11 +676,18 @@ SPEC: dict[str, dict[str, Cmd]] = { "/chat/{chatId}/message", "Send a message in a chat session.", body=( - P("message", required=True, help="Message text."), - P("cancel_current", "bool", help="Stop the current task first."), - P("stop_agent", "bool", help="Stop the agent."), - P("repos", "list", help="Repository full names for context."), - P("agent_id", help="Target one agent."), + P( + "message", + help="Message text. Required unless --cancel-current or --stop-agent is used.", + ), + P("cancel_current", "bool", help="Cancel the in-flight agent turn first."), + P("stop_agent", "bool", help="Park the target agent and its descendants."), + P( + "repos", + "json", + help='JSON array of repository refs: [{"repoId":"...","branch":"main"}].', + ), + P("agent_id", help="Target one subagent instead of the root agent."), ), ), "findings": Cmd("GET", "/chat/{chatId}/findings", "List the findings of a chat session."), @@ -626,7 +754,10 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/knowledge", "List knowledge documents.", - query=_q("source_type", "search", "limit"), + query=( + *_q("source_type", "search"), + P("limit", "int", help="Maximum documents to return."), + ), ), "add": Cmd( "POST", @@ -731,7 +862,21 @@ SPEC: dict[str, dict[str, Cmd]] = { "Create an installation link. The provider is github or slack. A person approves it.", link="url", ), - "disconnect": Cmd("DELETE", "/integrations/{provider}", "Disconnect an integration."), + "disconnect": Cmd( + "DELETE", + "/integrations/{provider}", + "Disconnect an integration.", + query=( + P( + "installation_id", + "int", + help=( + "Installation ID. Required for github, gitlab, and bitbucket; " + "unsupported for other providers." + ), + ), + ), + ), }, "connectors": { "list": Cmd("GET", "/connectors", "List network connectors."), @@ -745,7 +890,16 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/connectors/{connectorId}", "Get one network connector.", - query=_q("include_command"), + query=( + P( + "include_command", + "bool", + help=( + "Include the one-time Docker enrollment command. " + "The command contains sensitive connector credentials." + ), + ), + ), ), "status": Cmd( "GET", "/connectors/{connectorId}/status", "Get the status of a network connector." @@ -780,7 +934,13 @@ SPEC: dict[str, dict[str, Cmd]] = { ), "delete": Cmd("DELETE", "/webhooks/{webhookId}", "Delete a webhook."), "deliveries": Cmd( - "GET", "/webhooks/{webhookId}/deliveries", "List the deliveries of a webhook." + "GET", + "/webhooks/{webhookId}/deliveries", + "List the deliveries of a webhook.", + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-100)."), + ), ), }, "analytics": { @@ -800,23 +960,37 @@ SPEC: dict[str, dict[str, Cmd]] = { "GET", "/audit", "List audit log entries.", - query=_q( - "action", "resource_type", "actor_id", "date_from", "date_to", "format", "all" + query=( + P("page", "int", help="Results page (starts at 1)."), + P("limit", "int", help="Results per page (1-1000)."), + *_q("action", "resource_type", "actor_id", "date_from", "date_to"), + P( + "format", + help="Output format: json, csv, ndjson, jsonl, snowflake, or splunk.", + ), + P("all", "bool", help="Stream all matches when exporting instead of one page."), ), ), }, "costs": { "overview": Cmd( - "GET", "/llm-costs", "Show the LLM cost overview.", query=_q("range", "from", "to") + "GET", + "/llm-costs", + "Self-hosted only: show the LLM cost overview.", + query=_q("range", "from", "to"), + ), + "run": Cmd( + "GET", + "/llm-costs/runs/{runType}/{runId}", + "Self-hosted only: get the LLM costs of one run.", ), - "run": Cmd("GET", "/llm-costs/runs/{runType}/{runId}", "Get the LLM costs of one run."), }, "llm-settings": { - "get": Cmd("GET", "/llm-settings", "Get the LLM settings."), + "get": Cmd("GET", "/llm-settings", "Self-hosted only: get the LLM settings."), "update": Cmd( "PUT", "/llm-settings", - "Update the LLM settings.", + "Self-hosted only: update the LLM settings.", body=( P( "modelConfigs", @@ -856,6 +1030,21 @@ SPEC: dict[str, dict[str, Cmd]] = { P("type", required=True, help="Token type, personal or service."), P("name", required=True, help="Token name."), P("scopes", "list", help="API scopes for the token."), + P( + "rbac_scopes", + "json-list", + help=( + "JSON array of resource restrictions; each item has type " + "target, tag, or business_unit and a value." + ), + ), + P( + "expires_at", + help=( + "Absolute expiration date/time (ISO 8601; mutually exclusive " + "with --expires-in-days)." + ), + ), P("expires_in_days", "int", help="Days until the token expires."), ), ), @@ -920,8 +1109,8 @@ GROUP_HELP: dict[str, str] = { "webhooks": "Manage webhooks", "analytics": "Read analytics data", "audit": "Read the audit log", - "costs": "Read LLM cost data", - "llm-settings": "Manage LLM model settings", + "costs": "Self-hosted only: read LLM cost data", + "llm-settings": "Self-hosted only: manage LLM model settings", "settings": "Manage notification settings", "license": "Read license information", "tokens": "Manage API tokens", diff --git a/strix/interface/cloud/workspaces.py b/strix/interface/cloud/workspaces.py index 6df14a5a..133ec92a 100644 --- a/strix/interface/cloud/workspaces.py +++ b/strix/interface/cloud/workspaces.py @@ -1,27 +1,34 @@ """`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 for a token in that workspace, and -stores the token in the credential file. The role of the account in the new -workspace limits the granted scopes. +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 argparse -from typing import Any, cast +import os +from typing import TYPE_CHECKING, Any, cast from rich.console import Console +from rich.markup import escape -from strix.interface.cloud import http +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.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 = argparse.ArgumentParser( + parser = CloudArgumentParser( prog="strix cloud workspaces use", description="Switch the stored API token to another workspace.", ) @@ -36,7 +43,7 @@ def run_workspace_use(argv: list[str]) -> int: metavar="SCOPE", default=None, help=( - "API scopes for the new token. Without this option, preserve the stored token's scopes." + "API scopes after switching. Without this option, preserve the stored request scopes." ), ) parser.add_argument("--json", action="store_true", help="Print the raw JSON response.") @@ -45,61 +52,102 @@ def run_workspace_use(argv: list[str]) -> int: 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)) - http.configure(base_url=args.app_url, timeout=args.timeout) try: + http.configure( + base_url=args.app_url, + timeout=args.timeout, + token_override=bool(args.token), + ) return _use(console, args, as_json=as_json) except http.CloudError as exc: - if as_json: - emit(console, {"error": str(exc)}, as_json=True) - else: - console.print(f"[red]Error:[/] {exc}") + _emit_cloud_error(console, exc, as_json=as_json) return exc.exit_code def _use(console: Console, args: argparse.Namespace, *, as_json: bool) -> int: workspace = _find_workspace(args.workspace, token=args.token) - record = read_record() or {} + 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 - elif args.token is None: - stored_scopes = record.get("scopes") + elif not external_token: + stored_scopes = stored_record.get("requested_scopes", stored_record.get("scopes")) if ( isinstance(stored_scopes, list) and stored_scopes and all(isinstance(scope, str) for scope in stored_scopes) ): - body["scopes"] = stored_scopes - minted = http.check( - http.request( - "POST", - f"/workspaces/{workspace['id']}/token", - token=args.token, - body=body or None, - ) + body["scopes"] = [scope for scope in stored_scopes if isinstance(scope, str)] + switched = _switch_workspace_token( + str(workspace["id"]), + token=args.token, + body=body or None, ) - if not isinstance(minted, dict) or not minted.get("api_token"): - raise http.CloudError("the platform did not return a token.") - minted_record = cast("dict[str, Any]", minted) + 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") + if not isinstance(switched_scopes, list) or not all( + isinstance(scope, str) for scope in switched_scopes + ): + raise _workspace_switch_unknown("the platform response contained invalid scopes") record.update( { - "api_token": minted_record["api_token"], - "organization_id": minted_record.get("organization_id", workspace["id"]), - "organization_name": minted_record.get("organization_name", workspace.get("name", "")), - "expires_at": minted_record.get("expires_at"), - "scopes": minted_record.get("scopes", []), + "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"), + "scopes": switched_scopes, + "requested_scopes": ( + list(args.scopes) + if args.scopes + else ( + stored_record.get( + "requested_scopes", + stored_record.get("scopes", []), + ) + if not external_token + else switched_scopes + ) + ), + "app_url": http.app_url(), } ) - if minted_record.get("email"): - record["email"] = minted_record["email"] - save_record(record) + if switched_record.get("email"): + record["email"] = switched_record["email"] + 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"], @@ -109,17 +157,71 @@ def _use(console: Console, args: argparse.Namespace, *, as_json: bool) -> int: if as_json: emit(console, result, as_json=True) return http.EXIT_OK - console.print(f"[green]✓ Switched to workspace [bold]{record['organization_name']}[/].[/]") + 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: - console.print(f" Scopes: [dim]{' '.join(str(s) for s in scopes)}[/]") - console.print(f" Token: stored in [dim]{AUTH_PATH}[/]") + scope_names = [scope for scope in scopes if isinstance(scope, str)] + if scope_names: + rendered_scopes = escape(sanitize_terminal_text(" ".join(scope_names))) + console.print(f" Scopes: [dim]{rendered_scopes}[/]") + 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: + payload = dict(error.payload) if isinstance(error.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)) - items = listed.get("workspaces") if isinstance(listed, dict) else None + listed_record = cast("dict[str, Any]", listed) if isinstance(listed, dict) else {} + items = listed_record.get("workspaces") workspaces = [cast("dict[str, Any]", item) for item in (items or []) if isinstance(item, dict)] if not workspaces: raise http.CloudError("no workspaces found for this account.") diff --git a/strix/interface/completions.py b/strix/interface/completions.py index 83a53026..d1ebdfce 100644 --- a/strix/interface/completions.py +++ b/strix/interface/completions.py @@ -3,14 +3,18 @@ from __future__ import annotations import sys +from pathlib import Path from typing import Any -from strix.interface.cloud.spec import SPEC, Cmd +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", "credits") -_COMMON_FLAGS = ("--json", "--token", "--app-url", "--timeout", "--help") +_COMMON_FLAGS = ("--json", "--token", "--app-url", "--timeout", "-h", "--help") +_COMMON_VALUE_FLAGS = frozenset({"--token", "--app-url", "--timeout"}) +_WORKSPACE_USE_FLAGS = (*_COMMON_FLAGS, "--scopes") def run_completions(argv: list[str]) -> int: @@ -32,7 +36,9 @@ def run_completions(argv: list[str]) -> int: scripts = {"zsh": _zsh_script, "bash": _bash_script, "fish": _fish_script} generator = scripts.get(shell) if generator is None: - sys.stderr.write(f"Unknown shell: {shell}. Choose zsh, bash, or fish.\n") + sys.stderr.write( + f"Unknown shell: {sanitize_terminal_text(shell)}. Choose zsh, bash, or fish.\n" + ) return 2 sys.stdout.write(generator()) return 0 @@ -42,10 +48,14 @@ def completion_candidates(words: list[str]) -> list[str]: """Return candidates for words after the ``strix`` executable.""" prior, current = _split_cursor(words) if not prior: - return _matching(_ROOT_COMMANDS, current) - if prior[0] != "cloud": - return [] - return _cloud_candidates(prior[1:], current) + 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]: @@ -54,41 +64,100 @@ def _split_cursor(words: list[str]) -> tuple[list[str], str]: return words[:-1], words[-1] -def _cloud_candidates(prior: list[str], current: str) -> list[str]: +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 _matching(_session_flags(group), current) + 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) - verb_paths = [verb.split() for verb in commands] + 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)}) - exact_verbs = [" ".join(path) for path in matching_paths if len(path) == len(rest)] - candidates: list[str] = list(next_words) - for verb in exact_verbs: - if verb == "use" and group == "workspaces": - candidates.extend(_COMMON_FLAGS) - else: - candidates.extend(_command_flags(commands[verb])) - return _matching(candidates, current) + return _matching(next_words, current) + + +def _session_candidates(group: str, prior: list[str], current: str) -> list[str]: + flags = _session_flags(group) + value_flags: frozenset[str] = frozenset() + if group == "login": + value_flags = frozenset({"--scopes", "--workspace"}) + 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", "--workspace", "--help") + return ("--no-browser", "--scopes", "--workspace", "-h", "--help") if group == "whoami": - return ("--json", "--help") - return ("--help",) + return ("--json", "-h", "--help") + if group == "logout": + return ("--json", "-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, ...]: @@ -100,18 +169,21 @@ def _command_flags(cmd: Cmd) -> tuple[str, ...]: flags.append("--no-" + flag.removeprefix("--")) if cmd.method in ("POST", "PUT", "PATCH"): flags.append("--data") - if cmd.binary: - flags.append("--output") + 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.append("--wait") + 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", @@ -126,6 +198,93 @@ def _command_flags(cmd: Cmd) -> tuple[str, ...]: 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: @@ -154,8 +313,19 @@ compdef _strix strix def _bash_script() -> str: return r"""_strix_completion() { local -a candidates - mapfile -t candidates < <(strix completions --candidates "${COMP_WORDS[@]:1:$COMP_CWORD}") - COMPREPLY=( $(compgen -W "${candidates[*]}" -- "${COMP_WORDS[$COMP_CWORD]}") ) + 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 """ diff --git a/strix/interface/main.py b/strix/interface/main.py index e7dfc4ab..1f3d21ad 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -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] = [] @@ -416,6 +424,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 []` 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": diff --git a/strix/interface/platform_cli.py b/strix/interface/platform_cli.py index 65dd067c..7ae8b739 100644 --- a/strix/interface/platform_cli.py +++ b/strix/interface/platform_cli.py @@ -15,15 +15,18 @@ import sys import time import webbrowser from pathlib import Path -from typing import Any, cast -from urllib.parse import urlparse +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.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 @@ -34,12 +37,6 @@ _DEFAULT_POLL_INTERVAL_S = 5 _MAX_POLL_INTERVAL_S = 60 _MAX_EXPIRES_IN_S = 30 * 60 -_LOGIN_USAGE = ( - "Usage:\n" - " strix cloud login [--no-browser] [--scopes SCOPE ...] [--workspace WORKSPACE]\n" - " strix cloud whoami\n strix cloud logout" -) - _ROLE_RANK = {"viewer": 0, "analyst": 1, "admin": 2} @@ -47,6 +44,19 @@ 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("/") @@ -83,13 +93,10 @@ def run_login(argv: list[str]) -> int: console = Console() subcommand = argv[0] if argv else None - if subcommand in ("-h", "--help", "help"): - console.print(_LOGIN_USAGE) - return 0 if subcommand == "status": return _status(console, argv[1:]) if subcommand == "logout": - return _logout(console) + return _logout(console, argv[1:]) return _login(console, argv) @@ -128,7 +135,7 @@ def _login(console: Console, argv: list[str]) -> int: console.print() host = urlparse(_app_url()).netloc or _app_url() - console.print(f"[bold]Signing in to the Strix platform[/] [dim]({host})[/]") + 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.[/]" ) @@ -142,7 +149,7 @@ def _login(console: Console, argv: list[str]) -> int: workspace=args.workspace, ) except PlatformAuthError as exc: - console.print(f"[red]Sign-in failed:[/] {exc}") + console.print(f"[red]Sign-in failed:[/] {_terminal_markup(exc)}") return 1 except KeyboardInterrupt: console.print("\n[yellow]Sign-in cancelled.[/]") @@ -151,9 +158,11 @@ def _login(console: Console, argv: list[str]) -> int: try: save_record(record) except OSError as exc: - console.print(f"[red]Sign-in succeeded, but the token could not be stored:[/] {exc}") console.print( - f"[dim]Check that {AUTH_PATH.parent} is writable, " + 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 @@ -172,10 +181,14 @@ def _run_device_flow( interactive = workspace is not None or (sys.stdin.isatty() and scopes is None) try: - response = requests.post(f"{app_url}/api/v1/cli/login", timeout=_HTTP_TIMEOUT_S) + 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 response.ok: + if not 200 <= response.status_code < 300: raise PlatformAuthError(_error_detail(response)) authorization = _json_object(response) @@ -196,18 +209,20 @@ def _run_device_flow( ) 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"), - (user_code, "bold cyan"), - ("\n\nOpen this URL in your browser and confirm the code:\n", "dim"), - (verification_uri, "underline"), + (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): @@ -223,21 +238,25 @@ def _run_device_flow( deadline = time.monotonic() + expires_in while time.monotonic() < deadline: - time.sleep(interval) + 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 poll.ok: + if 200 <= poll.status_code < 300: return _finish_login(console, app_url, poll, scopes=scopes, workspace=workspace) delta = _handle_poll_error(poll) if delta is None: break - interval += delta + interval = min(interval + delta, _MAX_POLL_INTERVAL_S) raise PlatformAuthError("the sign-in request expired. Run `strix cloud login` again.") @@ -269,11 +288,20 @@ def _finish_login( result = _json_object(poll) if result.get("selection_required"): return _complete_selection(console, app_url, result, scopes=scopes, workspace=workspace) - return _require_api_token(result) + return _bind_login_record(_require_api_token(result), app_url, scopes) -def _signed_in_record(response: requests.Response) -> dict[str, Any]: - return _require_api_token(_json_object(response)) +def _signed_in_record( + response: requests.Response, + *, + app_url: str, + requested_scopes: list[str] | None, +) -> dict[str, Any]: + return _bind_login_record( + _require_api_token(_json_object(response)), + app_url, + requested_scopes, + ) def _require_api_token(record: dict[str, Any]) -> dict[str, Any]: @@ -283,6 +311,33 @@ def _require_api_token(record: dict[str, Any]) -> dict[str, Any]: return record +def _bind_login_record( + record: dict[str, Any], app_url: str, requested_scopes: list[str] | None +) -> 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 = requested_scopes if requested_scopes is not None else record.get("scopes") + if isinstance(preference, list) and all(isinstance(scope, str) for scope in preference): + bound["requested_scopes"] = list(dict.fromkeys(cast("list[str]", preference))) + return bound + + def _complete_selection( console: Console, app_url: str, @@ -314,12 +369,17 @@ def _complete_selection( 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 response.ok: + if not 200 <= response.status_code < 300: raise PlatformAuthError(_error_detail(response)) - return _signed_in_record(response) + return _signed_in_record( + response, + app_url=app_url, + requested_scopes=chosen_scopes, + ) def _dict_items(value: Any) -> list[dict[str, Any]]: @@ -332,19 +392,38 @@ def _choose_workspace( console: Console, organizations: list[dict[str, Any]], workspace: str | None ) -> dict[str, Any]: if workspace is not None: - wanted = workspace.strip().lower() - for org in organizations: - if wanted in (str(org.get("id", "")).lower(), str(org.get("name", "")).lower()): - return org + 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): - console.print(f" [cyan]{index}[/]. {org.get('name', '')} [dim]({org.get('role', '')})[/]") + 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): @@ -364,10 +443,11 @@ def _choose_scopes(console: Console, catalog: list[dict[str, Any]], role: str) - console.print() console.print("[bold]Select token scopes:[/]") console.print( - " [cyan]1[/]. Recommended [dim](scans, vulnerabilities, schedules, assets, billing)[/]" + " [cyan]1[/]. Recommended [dim](scans, findings, schedules, assets, uploads, " + "workspace switching, billing)[/]" ) console.print(" [cyan]2[/]. Full access [dim](every scope your role allows)[/]") - console.print(" [cyan]3[/]. Minimal [dim](scans and billing read only)[/]") + 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" @@ -396,9 +476,11 @@ def _choose_custom_scopes(console: Console, allowed: list[dict[str, Any]]) -> li 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}[/]. {scope}{required}" - f"\n [dim]{item.get('description', '')}[/]" + 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: " @@ -407,12 +489,14 @@ def _choose_custom_scopes(console: Console, allowed: list[dict[str, Any]]) -> li 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 {part!r}: not a number from the list.[/]") + 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]{scope} is always included.[/]") + console.print(f"[yellow]{_terminal_markup(scope)} is always included.[/]") continue if scope in selected: selected.discard(scope) @@ -454,13 +538,14 @@ def _print_success(console: Console, record: dict[str, Any]) -> None: console.print() console.print("[green]✓ Signed in to the Strix platform.[/]") if email: - console.print(f" Account: [bold]{email}[/]") + console.print(f" Account: [bold]{_terminal_markup(email)}[/]") if organization: - console.print(f" Workspace: [bold]{organization}[/]") + console.print(f" Workspace: [bold]{_terminal_markup(organization)}[/]") scopes = record.get("scopes") if isinstance(scopes, list) and scopes: - console.print(f" Scopes: [dim]{' '.join(str(s) for s in scopes)}[/]") - console.print(f" Token: stored in [dim]{AUTH_PATH}[/]") + rendered_scopes = _terminal_markup(" ".join(str(s) for s in scopes)) + console.print(f" Scopes: [dim]{rendered_scopes}[/]") + 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. " @@ -469,16 +554,27 @@ def _print_success(console: Console, record: dict[str, Any]) -> None: def _status(console: Console, argv: list[str]) -> int: - parser = argparse.ArgumentParser(prog="strix cloud whoami") + 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.") + 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 args.json: + 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.") @@ -486,7 +582,7 @@ def _status(console: Console, argv: list[str]) -> int: email = record.get("email", "unknown") organization = record.get("organization_name") or record.get("organization_id", "") expires_at = record.get("expires_at", "") - if args.json: + if as_json: payload = { "signed_in": True, "email": email, @@ -494,25 +590,67 @@ def _status(console: Console, argv: list[str]) -> int: "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]{email}[/]") + console.print(f"[green]Signed in[/] as [bold]{_terminal_markup(email)}[/]") if organization: - console.print(f" Workspace: {organization}") + console.print(f" Workspace: {_terminal_markup(organization)}") if expires_at: - console.print(f" Token expires: {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: + console.print(f" Scopes: {_terminal_markup(' '.join(str(scope) for scope in scopes))}") return 0 -def _logout(console: Console) -> int: +def _logout(console: Console, argv: list[str]) -> int: # noqa: PLR0911 + parser = _SessionArgumentParser( + prog="strix cloud logout", + description="Remove the managed-platform API token stored on this machine.", + ) + parser.add_argument("--json", action="store_true", help="Print the result as JSON.") + 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 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 {AUTH_PATH} manually." + 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}) + "\n") + return 0 console.print("[green]Signed out.[/] The stored API token was removed from this machine.") return 0 diff --git a/strix/interface/terminal_text.py b/strix/interface/terminal_text.py new file mode 100644 index 00000000..b0cc2b4b --- /dev/null +++ b/strix/interface/terminal_text.py @@ -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), + ) diff --git a/strix/interface/url_safety.py b/strix/interface/url_safety.py new file mode 100644 index 00000000..b9a77165 --- /dev/null +++ b/strix/interface/url_safety.py @@ -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) diff --git a/tests/test_cloud_cli.py b/tests/test_cloud_cli.py index 46dc88c6..43184bf5 100644 --- a/tests/test_cloud_cli.py +++ b/tests/test_cloud_cli.py @@ -6,19 +6,17 @@ import io import json import shutil import subprocess +import urllib.request import webbrowser -from typing import TYPE_CHECKING, Any +from pathlib import Path +from typing import Any import pytest import requests from strix.interface import cloud, platform_cli -from strix.interface.cloud import http, render, runner, workspaces -from strix.interface.cloud.spec import SPEC - - -if TYPE_CHECKING: - from pathlib import Path +from strix.interface.cloud import billing, http, payment_proxy, render, runner, workspaces +from strix.interface.cloud.spec import GROUP_HELP, SPEC class FakeResponse: @@ -41,6 +39,13 @@ class FakeResponse: raise ValueError("no JSON") return self._payload + def iter_content(self, chunk_size: int) -> Any: + for index in range(0, len(self.content), chunk_size): + yield self.content[index : index + chunk_size] + + def close(self) -> None: + pass + @pytest.fixture(autouse=True) def _token_env(monkeypatch: pytest.MonkeyPatch) -> None: @@ -177,6 +182,83 @@ def test_data_merges_extra_fields(monkeypatch: pytest.MonkeyPatch) -> None: assert seen["body"] == {"engagement_type": "code_review"} +def test_token_create_accepts_expiry_and_rbac_scope_flags( + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["body"] = kwargs.get("body") + return FakeResponse(payload={"id": "token-1", "token": "strix_pat_once"}) + + monkeypatch.setattr(http, "request", fake_request) + code = cloud.run_cloud( + [ + "tokens", + "create", + "--type", + "service", + "--name", + "ci", + "--expires-at", + "2026-09-30T12:00:00Z", + "--rbac-scopes", + '[{"type":"tag","value":"staging"}]', + "--json", + ] + ) + + assert code == 0 + assert seen["body"] == { + "type": "service", + "name": "ci", + "expires_at": "2026-09-30T12:00:00Z", + "rbac_scopes": [{"type": "tag", "value": "staging"}], + } + + +def test_token_create_rejects_non_array_rbac_scopes(capsys: Any) -> None: + code = cloud.run_cloud( + [ + "tokens", + "create", + "--type", + "service", + "--name", + "ci", + "--rbac-scopes", + '{"type":"tag","value":"staging"}', + "--json", + ] + ) + + assert code == http.EXIT_USAGE + assert json.loads(capsys.readouterr().out)["error"] == ("--rbac-scopes must be a JSON array") + + +def test_token_create_rejects_two_expiration_modes(capsys: Any) -> None: + code = cloud.run_cloud( + [ + "tokens", + "create", + "--type", + "personal", + "--name", + "local", + "--expires-at", + "2026-09-30T12:00:00Z", + "--expires-in-days", + "30", + "--json", + ] + ) + + assert code == http.EXIT_USAGE + assert json.loads(capsys.readouterr().out)["error"] == ( + "--expires-at and --expires-in-days are mutually exclusive." + ) + + def test_data_reads_a_file(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: seen: dict[str, Any] = {} @@ -204,6 +286,27 @@ def test_data_reads_stdin(monkeypatch: pytest.MonkeyPatch) -> None: assert seen["body"] == {"context": "staging"} +def test_required_secret_body_field_can_come_from_stdin( + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["body"] = kwargs.get("body") + return FakeResponse(payload={"ok": True}) + + monkeypatch.setattr(http, "request", fake_request) + monkeypatch.setattr("sys.stdin", io.StringIO('{"token":"provider-secret"}')) + + assert cloud.run_cloud(["integrations", "connect", "gitlab", "--data", "-", "--json"]) == 0 + assert seen["body"] == {"token": "provider-secret"} + + +def test_required_body_field_is_validated_after_data_merge(capsys: Any) -> None: + assert cloud.run_cloud(["integrations", "connect", "gitlab", "--json"]) == http.EXIT_USAGE + assert "--provider-token" in json.loads(capsys.readouterr().out)["error"] + + def test_data_reports_a_missing_file(tmp_path: Path) -> None: assert cloud.run_cloud(["scans", "start", "--data", f"@{tmp_path / 'nope.json'}"]) == 1 @@ -271,6 +374,36 @@ def test_wait_polls_until_the_status_is_final(monkeypatch: pytest.MonkeyPatch) - assert next(statuses, None) is None +def test_wait_failure_keeps_the_created_operation_id( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + def fake_request(method: str, _path: str, **_kwargs: Any) -> FakeResponse: + if method == "POST": + return FakeResponse(payload={"id": "scan-created", "status": "running"}) + return FakeResponse(status_code=503, payload={"detail": "temporarily unavailable"}) + + monkeypatch.setattr(http, "request", fake_request) + assert cloud.run_cloud(["scans", "start", "--domain-ids", "d1", "--wait", "--json"]) == 1 + payload = json.loads(capsys.readouterr().out) + assert payload["operation_id"] == "scan-created" + assert payload["status_unknown"] is True + + +def test_ambiguous_scan_request_warns_before_retry( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_args, **_kwargs: (_ for _ in ()).throw(http.CloudError("connection reset")), + ) + + assert cloud.run_cloud(["scans", "start", "--domain-ids", "d1", "--json"]) == 1 + payload = json.loads(capsys.readouterr().out) + assert payload["launch_outcome_unknown"] is True + assert "scans list" in payload["error"] + + def test_insufficient_credits_exits_with_payment_code(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( http, "request", lambda *_a, **_k: FakeResponse(status_code=402, payload={}) @@ -279,8 +412,12 @@ def test_insufficient_credits_exits_with_payment_code(monkeypatch: pytest.Monkey def test_data_rejects_non_object() -> None: - assert cloud.run_cloud(["scans", "start", "--data", "[1,2]"]) == 1 - assert cloud.run_cloud(["scans", "start", "--data", "not json"]) == 1 + assert cloud.run_cloud(["scans", "start", "--data", "[1,2]"]) == http.EXIT_USAGE + assert cloud.run_cloud(["scans", "start", "--data", "not json"]) == http.EXIT_USAGE + + +def test_typed_json_flag_parse_error_is_usage_error() -> None: + assert cloud.run_cloud(["scans", "start", "--domain-paths", "not-json"]) == http.EXIT_USAGE def test_missing_token_exits_with_auth_code( @@ -291,6 +428,91 @@ def test_missing_token_exits_with_auth_code( assert cloud.run_cloud(["credits"]) == http.EXIT_AUTH +def test_stored_token_is_never_sent_to_a_different_platform_origin( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) + monkeypatch.setattr(http, "_app_url_override", "https://attacker.example") + monkeypatch.setattr( + http, + "read_record", + lambda: {"api_token": "stored-secret", "app_url": "https://app.strix.ai"}, + ) + monkeypatch.setattr( + http.requests, + "request", + lambda *_args, **_kwargs: pytest.fail("a mismatched origin must not receive the token"), + ) + + with pytest.raises(http.CloudError, match="different platform") as raised: + http.request("GET", "/billing/credits") + assert raised.value.exit_code == http.EXIT_AUTH + + +def test_stored_token_requires_an_issuer_binding(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) + monkeypatch.setattr(http, "_app_url_override", "https://app.strix.ai") + monkeypatch.setattr(http, "read_record", lambda: {"api_token": "legacy-secret"}) + monkeypatch.setattr( + http.requests, + "request", + lambda *_args, **_kwargs: pytest.fail("an unbound token must not be sent"), + ) + + with pytest.raises(http.CloudError, match="not bound") as raised: + http.request("GET", "/billing/credits") + assert raised.value.exit_code == http.EXIT_AUTH + + +def test_stored_token_is_sent_only_to_its_bound_platform( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) + monkeypatch.setattr(http, "_app_url_override", "https://preview.strix.ai") + monkeypatch.setattr( + http, + "read_record", + lambda: {"api_token": "stored-secret", "app_url": "https://preview.strix.ai"}, + ) + seen: dict[str, Any] = {} + + def request(_method: str, url: str, **kwargs: Any) -> FakeResponse: + seen.update(url=url, headers=kwargs["headers"]) + return FakeResponse(payload={"balance": 1}) + + monkeypatch.setattr(http.requests, "request", request) + response = http.request("GET", "/billing/credits") + + assert response.status_code == 200 + assert seen["url"] == "https://preview.strix.ai/api/v1/billing/credits" + assert seen["headers"]["Authorization"] == "Bearer stored-secret" + + +def test_explicit_token_can_target_an_explicit_platform( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) + monkeypatch.setattr(http, "_app_url_override", "https://preview.strix.ai") + monkeypatch.setattr( + http, + "read_record", + lambda: {"api_token": "stored-secret", "app_url": "https://app.strix.ai"}, + ) + seen: dict[str, Any] = {} + + def request(_method: str, url: str, **kwargs: Any) -> FakeResponse: + seen.update(url=url, headers=kwargs["headers"]) + return FakeResponse(payload={"balance": 1}) + + monkeypatch.setattr(http.requests, "request", request) + override_value = "explicit-preview-" + str(1) + response = http.request("GET", "/billing/credits", token=override_value) + + assert response.status_code == 200 + assert seen["url"] == "https://preview.strix.ai/api/v1/billing/credits" + assert seen["headers"]["Authorization"] == f"Bearer {override_value}" + + def test_http_error_exit_codes(monkeypatch: pytest.MonkeyPatch) -> None: for status, expected in ((401, http.EXIT_AUTH), (403, http.EXIT_AUTH), (500, http.EXIT_ERROR)): monkeypatch.setattr( @@ -355,6 +577,107 @@ def test_topup_no_pay_prints_challenge(monkeypatch: pytest.MonkeyPatch, capsys: } +def test_topup_noninteractive_requires_explicit_payment_approval( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + challenge = {"payment_requirements": [{"amount": 500}]} + monkeypatch.setattr( + http, "request", lambda *_a, **_k: FakeResponse(status_code=402, payload=challenge) + ) + monkeypatch.setattr(runner.sys.stdin, "isatty", lambda: False) + monkeypatch.setattr( + billing.subprocess, + "run", + lambda *_a, **_k: pytest.fail("wallet must not run without --yes"), + ) + + code = cloud.run_cloud(["billing", "topup", "--credits", "5", "--json"]) + + assert code == http.EXIT_PAYMENT + payload = json.loads(capsys.readouterr().out) + assert "requires explicit approval" in payload["error"] + assert payload["challenge"] == challenge + + +@pytest.mark.parametrize("explicit_json,stdout_tty", [(True, True), (False, False)]) +def test_topup_machine_output_never_prompts_even_with_terminal_stdin( + explicit_json: bool, + stdout_tty: bool, + monkeypatch: pytest.MonkeyPatch, + capsys: Any, +) -> None: + challenge = {"payment_requirements": [{"amount": 500}]} + monkeypatch.setattr( + http, "request", lambda *_a, **_k: FakeResponse(status_code=402, payload=challenge) + ) + monkeypatch.setattr(runner.sys.stdin, "isatty", lambda: True) + monkeypatch.setattr(runner.sys.stdout, "isatty", lambda: stdout_tty) + monkeypatch.setattr( + runner.Console, + "input", + lambda *_a, **_k: pytest.fail("machine-readable top-up must not prompt"), + ) + + argv = ["billing", "topup", "--credits", "5"] + if explicit_json: + argv.append("--json") + assert cloud.run_cloud(argv) == http.EXIT_PAYMENT + + payload = json.loads(capsys.readouterr().out) + assert "requires explicit approval" in payload["error"] + assert payload["challenge"] == challenge + + +def test_topup_payment_flags_are_mutually_exclusive( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(http, "request", lambda *_a, **_k: pytest.fail("must not request")) + assert ( + cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--no-pay", "--json"]) + == http.EXIT_USAGE + ) + + +def test_data_cannot_override_an_explicit_payment_amount( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(http, "request", lambda *_a, **_k: pytest.fail("must not request")) + + assert ( + cloud.run_cloud( + [ + "billing", + "topup", + "--credits", + "5", + "--data", + '{"credits": 500}', + "--yes", + "--json", + ] + ) + == http.EXIT_USAGE + ) + assert "cannot override explicit" in json.loads(capsys.readouterr().out)["error"] + + +def test_topup_missing_wallet_keeps_json_machine_readable( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + challenge = {"payment_requirements": [{"amount": 500}]} + monkeypatch.setattr( + http, "request", lambda *_a, **_k: FakeResponse(status_code=402, payload=challenge) + ) + monkeypatch.setattr(shutil, "which", lambda _name: None) + + code = cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--json"]) + + assert code == http.EXIT_PAYMENT + payload = json.loads(capsys.readouterr().out) + assert "wallet client" in payload["error"] + assert payload["challenge"] == challenge + + def test_topup_success_without_payment(monkeypatch: pytest.MonkeyPatch, capsys: Any) -> None: receipt = {"credits_granted": 5, "duplicate": False, "balance": 5} monkeypatch.setattr( @@ -365,30 +688,292 @@ def test_topup_success_without_payment(monkeypatch: pytest.MonkeyPatch, capsys: assert json.loads(capsys.readouterr().out) == receipt -def test_topup_passes_payment_method_to_wallet(monkeypatch: pytest.MonkeyPatch) -> None: +def test_topup_keeps_token_out_of_wallet_process_and_forwards_payment( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: challenge = {"payment_requirements": [{"amount": 500}]} + receipt = { + "credits_granted": 5, + "duplicate": False, + "reference": "pay_test_1", + "balance": 10, + } + api_credential = "opaque-test-api-credential-value" + monkeypatch.chdir(tmp_path) + (tmp_path / ".npmrc").write_text("registry=https://malicious.invalid\n", encoding="utf-8") + monkeypatch.setenv("UNRELATED_CODING_AGENT_SECRET", "must-not-reach-wallet") monkeypatch.setattr( http, "request", lambda *_a, **_k: FakeResponse(status_code=402, payload=challenge) ) - monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: api_credential) monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") commands: list[list[str]] = [] + child_envs: list[dict[str, str]] = [] + child_cwds: list[Path] = [] + upstream: dict[str, Any] = {} - def fake_run(command: list[str], **_kwargs: Any) -> Any: + def fake_upstream_request(method: str, url: str, **kwargs: Any) -> FakeResponse: + upstream.update(method=method, url=url, **kwargs) + return FakeResponse(payload=receipt, content=json.dumps(receipt).encode()) + + monkeypatch.setattr(payment_proxy.requests, "request", fake_upstream_request) + + def fake_run(command: list[str], **kwargs: Any) -> Any: commands.append(command) - return type("Result", (), {"returncode": 0})() + child_envs.append(kwargs["env"]) + child_cwds.append(Path(kwargs["cwd"])) + wallet_url = next( + argument for argument in command if argument.startswith("http://127.0.0.1:") + ) + request = urllib.request.Request( # noqa: S310 + wallet_url, + data=json.dumps({"credits": 5}).encode(), + headers={ + "Authorization": "Payment wallet-credential", + "Content-Type": "application/json", + }, + method="POST", + ) + with urllib.request.urlopen(request, timeout=2) as response: # noqa: S310 + stdout = response.read().decode() + return type( + "Result", + (), + { + "returncode": 0, + "stdout": stdout, + "stderr": "", + }, + )() monkeypatch.setattr(subprocess, "run", fake_run) code = cloud.run_cloud( ["billing", "topup", "--credits", "5", "--yes", "--payment-method", "pm_card_visa"] ) assert code == 0 - assert "X-Strix-Authorization: Bearer tok" in commands[0] - assert "Authorization: Bearer tok" not in commands[0] + assert all(api_credential not in argument for argument in commands[0]) + assert "mppx@0.8.17" in commands[0] + assert "--registry=https://registry.npmjs.org" in commands[0] + assert "--ignore-scripts" in commands[0] + assert "-H" not in commands[0] + assert "--fail" in commands[0] + assert child_envs[0].get("STRIX_API_TOKEN") is None + assert child_envs[0].get("UNRELATED_CODING_AGENT_SECRET") is None + for name in ("NO_PROXY", "no_proxy"): + bypasses = child_envs[0][name].split(",") + assert "127.0.0.1" in bypasses + assert "localhost" in bypasses + assert "::1" in bypasses + assert child_cwds[0] != tmp_path + assert upstream["method"] == "POST" + assert upstream["url"].endswith("/api/v1/billing/topup") + assert upstream["headers"]["X-Strix-Authorization"] == f"Bearer {api_credential}" + assert upstream["headers"]["Authorization"] == "Payment wallet-credential" + assert upstream["data"] == json.dumps({"credits": 5}).encode() assert "-M" in commands[0] assert "paymentMethod=pm_card_visa" in commands[0] +def test_topup_wallet_failure_is_one_redacted_json_object( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + challenge = {"payment_requirements": [{"amount": 500}]} + monkeypatch.setattr( + http, "request", lambda *_a, **_k: FakeResponse(status_code=402, payload=challenge) + ) + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") + monkeypatch.setattr( + subprocess, + "run", + lambda *_a, **_k: type( + "Result", + (), + { + "returncode": 9, + "stdout": "", + "stderr": ( + "failed with Bearer super-secret and " + "Authorization: Payment wallet-super-secret\x1b[2J" + ), + }, + )(), + ) + + assert cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--json"]) == 5 + payload = json.loads(capsys.readouterr().out) + assert payload["wallet_exit_code"] == 9 + assert payload["payment_outcome_unknown"] is True + assert "billing credits" in payload["error"] + assert "super-secret" not in payload["detail"] + assert "Bearer [redacted]" in payload["detail"] + assert "Payment [redacted]" in payload["detail"] + assert "\x1b" not in payload["detail"] + + +def test_topup_wallet_interruption_reports_unknown_payment_outcome( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + status_code=402, + payload={"payment_requirements": [{"amount": 500}]}, + ), + ) + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") + monkeypatch.setattr( + subprocess, "run", lambda *_a, **_k: (_ for _ in ()).throw(KeyboardInterrupt) + ) + + assert cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--json"]) == 130 + payload = json.loads(capsys.readouterr().out) + assert payload["interrupted"] is True + assert payload["payment_outcome_unknown"] is True + assert "billing credits" in payload["error"] + assert "before retrying" in payload["error"] + + +def test_topup_non_json_wallet_success_requires_balance_verification( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + status_code=402, + payload={"payment_requirements": [{"amount": 500}]}, + ), + ) + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") + monkeypatch.setattr( + subprocess, + "run", + lambda *_a, **_k: type("Result", (), {"returncode": 0, "stdout": "paid", "stderr": ""})(), + ) + + assert cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--json"]) == 5 + payload = json.loads(capsys.readouterr().out) + assert "did not return JSON" in payload["error"] + assert "before retrying" in payload["error"] + assert payload["payment_outcome_unknown"] is True + + +def test_topup_rejects_parseable_wallet_error_as_a_success( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + status_code=402, + payload={"payment_requirements": [{"amount": 500}]}, + ), + ) + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") + monkeypatch.setattr( + subprocess, + "run", + lambda *_a, **_k: type( + "Result", + (), + { + "returncode": 0, + "stdout": '{"detail":"Failed to process the top-up payment"}', + "stderr": "", + }, + )(), + ) + + assert cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--json"]) == 5 + payload = json.loads(capsys.readouterr().out) + assert "invalid top-up receipt" in payload["error"] + assert payload["payment_outcome_unknown"] is True + + +def test_topup_does_not_trust_an_unobserved_wallet_receipt( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + receipt = { + "credits_granted": 5, + "duplicate": False, + "reference": "untrusted-wallet-output", + "balance": 10, + } + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + status_code=402, + payload={"payment_requirements": [{"amount": 500}]}, + ), + ) + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") + monkeypatch.setattr( + subprocess, + "run", + lambda *_a, **_k: type( + "Result", + (), + {"returncode": 0, "stdout": json.dumps(receipt), "stderr": ""}, + )(), + ) + + assert cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes", "--json"]) == 5 + payload = json.loads(capsys.readouterr().out) + assert "did not confirm" in payload["error"] + assert payload["payment_outcome_unknown"] is True + + +def test_topup_human_mode_requires_a_bridge_confirmed_receipt( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(render.sys.stdout, "isatty", lambda: True) + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + status_code=402, + payload={"payment_requirements": [{"amount": 500}]}, + ), + ) + monkeypatch.setattr(http, "api_token", lambda *_a, **_k: "tok") + monkeypatch.setattr(shutil, "which", lambda _name: "/usr/bin/npx") + monkeypatch.setattr( + payment_proxy.requests, + "request", + lambda *_a, **_k: FakeResponse(status_code=200, content=b"not a receipt"), + ) + + def fake_run(command: list[str], **_kwargs: Any) -> Any: + wallet_url = next( + argument for argument in command if argument.startswith("http://127.0.0.1:") + ) + request = urllib.request.Request( # noqa: S310 + wallet_url, + data=json.dumps({"credits": 5}).encode(), + headers={"Authorization": "Payment wallet-credential"}, + method="POST", + ) + with urllib.request.urlopen(request, timeout=2) as response: # noqa: S310 + response.read() + return type("Result", (), {"returncode": 0, "stdout": None, "stderr": None})() + + monkeypatch.setattr(subprocess, "run", fake_run) + + assert cloud.run_cloud(["billing", "topup", "--credits", "5", "--yes"]) == 5 + output = capsys.readouterr().out + assert "without a confirmed receipt" in output + assert "outcome is unknown" in output + assert "billing credits" in output + + def test_render_json_mode_when_not_a_tty() -> None: assert render.json_mode(flag=True) is True # Under pytest, stdout is captured and is not a terminal. @@ -398,7 +983,7 @@ def test_render_json_mode_when_not_a_tty() -> None: def test_render_list_extraction() -> None: rows = render._list_of_dicts({"scans": [{"id": "a"}, {"id": "b"}]}) assert rows == [{"id": "a"}, {"id": "b"}] - assert render._list_of_dicts({"scans": [], "total": 1}) is None + assert render._list_of_dicts({"scans": [], "total": 1}) == [] assert render._list_of_dicts([{"id": "a"}, "x"]) is None @@ -409,9 +994,278 @@ def test_spec_paths_are_well_formed() -> None: assert cmd.method in ("GET", "POST", "PUT", "PATCH", "DELETE"), f"{group} {verb}" assert cmd.help, f"{group} {verb} has no help text" for param in cmd.query + cmd.body: - assert param.kind in ("str", "int", "float", "bool", "list", "json"), ( - f"{group} {verb} {param.name}" - ) + assert param.kind in ( + "str", + "int", + "float", + "bool", + "list", + "json", + "json-list", + ), f"{group} {verb} {param.name}" + + +@pytest.mark.parametrize( + ("group", "verb"), + [ + ("scans", "list"), + ("vulns", "list"), + ("domains", "list"), + ("repos", "list"), + ("pr-reviews", "list"), + ("pr-reviews", "findings"), + ("webhooks", "deliveries"), + ("audit", "list"), + ], +) +def test_paginated_commands_expose_integer_page_and_limit(group: str, verb: str) -> None: + params = {param.name: param for param in SPEC[group][verb].query} + assert params["page"].kind == "int" + assert params["limit"].kind == "int" + + +def test_list_query_types_match_the_api_contract() -> None: + scans = {param.name: param for param in SPEC["scans"]["list"].query} + assert scans["include_retests"].kind == "bool" + assert {"sort_by", "sort_order"} <= scans.keys() + + vulnerabilities = {param.name: param for param in SPEC["vulns"]["list"].query} + assert "sort_order" in vulnerabilities + + for group in ("domains", "repos"): + params = {param.name: param for param in SPEC[group]["list"].query} + assert params["limit"].kind == "int" + assert "sort_order" in params + + reviews = {param.name: param for param in SPEC["pr-reviews"]["list"].query} + findings = {param.name: param for param in SPEC["pr-reviews"]["findings"].query} + audit = {param.name: param for param in SPEC["audit"]["list"].query} + assert reviews["include_counts"].kind == "bool" + assert findings["include_stats"].kind == "bool" + assert audit["all"].kind == "bool" + + components = {param.name: param for param in SPEC["repos"]["supply-chain components"].query} + knowledge = {param.name: param for param in SPEC["knowledge"]["list"].query} + assert components["limit"].kind == "int" + assert components["offset"].kind == "int" + assert knowledge["limit"].kind == "int" + + +def test_scan_creating_replay_commands_support_bounded_waits() -> None: + assert SPEC["scans"]["rerun"].wait_path == "/scans/{id}" + assert SPEC["vulns"]["retest"].wait_path == "/scans/{id}" + + +def test_scan_start_parameter_contract_and_help() -> None: + params = {param.name: param for param in SPEC["scans"]["start"].body} + assert params["headers"].kind == "json" + assert "array" in params["headers"].help.lower() + assert params["concerns"].kind == "str" + assert all(tier in params["scan_tier"].help for tier in ("lite", "standard", "ultra")) + assert "pro" not in params["scan_tier"].help + assert "max" not in params["scan_tier"].help + assert "self-hosted" in params["model_config_id"].help.lower() + assert "self-hosted" in params["max_budget_usd"].help.lower() + + +def test_report_branding_flags_preserve_the_api_query_names( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["query"] = kwargs.get("query") + return FakeResponse(content=b"report") + + monkeypatch.setattr(http, "request", fake_request) + output = tmp_path / "report.pdf" + assert ( + cloud.run_cloud( + [ + "scans", + "report", + "scan-1", + "--provider-name", + "Strix Partner", + "--member-name-0", + "Alex", + "--member-email-0", + "alex@example.test", + "--member-name-1", + "Sam", + "--member-email-1", + "sam@example.test", + "--output", + str(output), + ] + ) + == 0 + ) + assert output.read_bytes() == b"report" + assert seen["query"] == { + "providerName": "Strix Partner", + "memberName0": "Alex", + "memberEmail0": "alex@example.test", + "memberName1": "Sam", + "memberEmail1": "sam@example.test", + } + + +def test_scan_start_collects_header_array_and_string_concerns( + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["body"] = kwargs.get("body") + return FakeResponse(payload={"id": "scan-1"}) + + monkeypatch.setattr(http, "request", fake_request) + assert ( + cloud.run_cloud( + [ + "scans", + "start", + "--headers", + '[{"name":"X-Test","value":"one"}]', + "--concerns", + "authorization boundaries", + "--scan-tier", + "standard", + "--json", + ] + ) + == 0 + ) + assert seen["body"] == { + "headers": [{"name": "X-Test", "value": "one"}], + "concerns": "authorization boundaries", + "scan_tier": "standard", + } + + +def test_control_only_scan_and_chat_messages_do_not_require_message( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[tuple[str, dict[str, Any] | None]] = [] + + def fake_request(_method: str, path: str, **kwargs: Any) -> FakeResponse: + calls.append((path, kwargs.get("body"))) + return FakeResponse(payload={"success": True}) + + monkeypatch.setattr(http, "request", fake_request) + assert cloud.run_cloud(["scans", "message", "scan-1", "--cancel-current", "--json"]) == 0 + assert ( + cloud.run_cloud( + ["chat", "send", "chat-1", "--stop-agent", "--agent-id", "agent-1", "--json"] + ) + == 0 + ) + assert calls == [ + ("/scans/scan-1/message", {"cancel_current": True}), + ("/chat/chat-1/message", {"stop_agent": True, "agent_id": "agent-1"}), + ] + + +def test_chat_repositories_use_the_api_object_shape(monkeypatch: pytest.MonkeyPatch) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["body"] = kwargs.get("body") + return FakeResponse(payload={"id": "chat-1"}) + + monkeypatch.setattr(http, "request", fake_request) + assert ( + cloud.run_cloud( + [ + "chat", + "start", + "--message", + "Review this repository", + "--repos", + '[{"repoId":"repo-1","branch":"main"}]', + "--json", + ] + ) + == 0 + ) + assert seen["body"] == { + "message": "Review this repository", + "repos": [{"repoId": "repo-1", "branch": "main"}], + } + + +def test_schedule_budget_accepts_fractional_usd(monkeypatch: pytest.MonkeyPatch) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["body"] = kwargs.get("body") + return FakeResponse(payload={"id": "schedule-1"}) + + monkeypatch.setattr(http, "request", fake_request) + assert ( + cloud.run_cloud(["schedules", "update", "schedule-1", "--max-budget-usd", "1.5", "--json"]) + == 0 + ) + assert seen["body"] == {"max_budget_usd": 1.5} + + +def test_integration_disconnect_sends_installation_id_query( + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: dict[str, Any] = {} + + def fake_request(method: str, path: str, **kwargs: Any) -> FakeResponse: + seen.update(method=method, path=path, query=kwargs.get("query")) + return FakeResponse(payload={"success": True}) + + monkeypatch.setattr(http, "request", fake_request) + assert ( + cloud.run_cloud( + ["integrations", "disconnect", "github", "--installation-id", "42", "--json"] + ) + == 0 + ) + assert seen == { + "method": "DELETE", + "path": "/integrations/github", + "query": {"installation_id": 42}, + } + + +def test_connector_command_flag_is_boolean_and_warns_that_it_is_sensitive( + monkeypatch: pytest.MonkeyPatch, +) -> None: + param = next( + param for param in SPEC["connectors"]["get"].query if param.name == "include_command" + ) + assert param.kind == "bool" + assert "sensitive" in param.help.lower() + + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["query"] = kwargs.get("query") + return FakeResponse(payload={"id": "connector-1"}) + + monkeypatch.setattr(http, "request", fake_request) + assert cloud.run_cloud(["connectors", "get", "connector-1", "--include-command", "--json"]) == 0 + assert seen["query"] == {"include_command": True} + + +def test_corrected_help_distinguishes_inboxes_reports_and_self_hosted_commands() -> None: + inbox = SPEC["domains"]["test-users provision-inbox"] + assert "does not create a test user" in inbox.help + + report = {param.name: param.help for param in SPEC["scans"]["report"].query} + assert "Report content" in report["format"] + assert "file type" in report["type"] + + for command in (*SPEC["costs"].values(), *SPEC["llm-settings"].values()): + assert "self-hosted only" in command.help.lower() + assert "self-hosted only" in GROUP_HELP["costs"].lower() + assert "self-hosted only" in GROUP_HELP["llm-settings"].lower() def test_every_command_builds_a_parser() -> None: @@ -581,9 +1435,59 @@ def test_integration_install_url_does_not_open_browser( assert "https://github.test/app" in capsys.readouterr().out +@pytest.mark.parametrize( + "argv,payload", + [ + ( + ["billing", "subscribe", "--plan", "strix_cloud"], + {"checkout_url": "file:///tmp/not-a-checkout"}, + ), + ( + ["integrations", "install", "github"], + {"url": "javascript:alert(1)"}, + ), + ], +) +def test_handoff_links_reject_non_http_schemes( + argv: list[str], + payload: dict[str, str], + monkeypatch: pytest.MonkeyPatch, + capsys: Any, +) -> None: + monkeypatch.setattr(runner.sys.stdout, "isatty", lambda: True) + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse(status_code=200, payload=payload), + ) + monkeypatch.setattr( + webbrowser, + "open", + lambda _url: pytest.fail("an untrusted URL must never be opened"), + ) + + assert cloud.run_cloud(argv) == http.EXIT_ERROR + output = capsys.readouterr().out + assert "invalid continuation URL" in output + assert next(iter(payload.values())) not in output + + +def test_handoff_missing_expected_url_is_an_error( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse(status_code=200, payload={"status": "created"}), + ) + assert cloud.run_cloud(["integrations", "install", "github", "--json"]) == 1 + assert "expected url URL" in json.loads(capsys.readouterr().out)["error"] + + def test_workspaces_use_switches_stored_token( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, capsys: Any ) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) auth_path = tmp_path / "platform-auth.json" monkeypatch.setattr(platform_cli, "AUTH_PATH", auth_path) monkeypatch.setattr(workspaces, "AUTH_PATH", auth_path) @@ -592,6 +1496,12 @@ def test_workspaces_use_switches_stored_token( "api_token": "old", "email": "a@b.test", "scopes": ["scans:read", "organizations:read", "tokens:write"], + "requested_scopes": [ + "scans:read", + "scans:write", + "organizations:read", + "tokens:write", + ], } ) @@ -608,9 +1518,9 @@ def test_workspaces_use_switches_stored_token( ) token_body = kwargs.get("body") return FakeResponse( - status_code=201, + status_code=200, payload={ - "api_token": "new-token", + "api_token": "old", "organization_id": "org_1", "organization_name": "Team One", "scopes": ["scans:read"], @@ -621,15 +1531,114 @@ def test_workspaces_use_switches_stored_token( code = cloud.run_cloud(["workspaces", "use", "team one", "--json"]) assert code == 0 assert calls == [("GET", "/workspaces"), ("POST", "/workspaces/org_1/token")] - assert token_body == {"scopes": ["scans:read", "organizations:read", "tokens:write"]} + assert token_body == { + "scopes": ["scans:read", "scans:write", "organizations:read", "tokens:write"] + } record = platform_cli.read_record() assert record is not None - assert record["api_token"] == "new-token" + assert record["api_token"] == "old" assert record["organization_name"] == "Team One" assert record["email"] == "a@b.test" assert "org_1" in capsys.readouterr().out +def test_workspace_use_explicit_token_starts_with_fresh_account_state( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + auth_path = tmp_path / "platform-auth.json" + monkeypatch.setattr(platform_cli, "AUTH_PATH", auth_path) + monkeypatch.setattr(workspaces, "AUTH_PATH", auth_path) + platform_cli.save_record( + { + "api_token": "account-a-token", + "email": "account-a@example.test", + "organization_id": "org_a", + "organization_name": "Account A", + "scopes": ["scans:read"], + "requested_scopes": ["scans:read", "tokens:write"], + } + ) + switch_body: dict[str, Any] | None = None + + def fake_request(method: str, path: str, **kwargs: Any) -> FakeResponse: + nonlocal switch_body + assert kwargs.get("token") == "account-b-token" + if method == "GET": + return FakeResponse(payload={"workspaces": [{"id": "org_b", "name": "Account B"}]}) + assert path == "/workspaces/org_b/token" + switch_body = kwargs.get("body") + return FakeResponse( + payload={ + "api_token": "account-b-token", + "organization_id": "org_b", + "organization_name": "Account B", + "scopes": ["scans:read", "organizations:read"], + } + ) + + monkeypatch.setattr(http, "request", fake_request) + assert ( + cloud.run_cloud(["workspaces", "use", "Account B", "--token", "account-b-token", "--json"]) + == 0 + ) + assert switch_body is None + record = platform_cli.read_record() + assert record is not None + assert record["api_token"] == "account-b-token" + assert record["organization_id"] == "org_b" + assert record["requested_scopes"] == ["scans:read", "organizations:read"] + assert "email" not in record + assert "account-a@example.test" not in auth_path.read_text(encoding="utf-8") + + +def test_workspace_use_environment_token_starts_with_fresh_account_state( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + auth_path = tmp_path / "platform-auth.json" + monkeypatch.setattr(platform_cli, "AUTH_PATH", auth_path) + monkeypatch.setattr(workspaces, "AUTH_PATH", auth_path) + monkeypatch.setenv("STRIX_API_TOKEN", "account-b-token") + platform_cli.save_record( + { + "api_token": "account-a-token", + "email": "account-a@example.test", + "organization_id": "org_a", + "organization_name": "Account A", + "scopes": ["scans:read"], + "requested_scopes": ["scans:read", "tokens:write"], + } + ) + switch_body: dict[str, Any] | None = None + + def fake_request(method: str, path: str, **kwargs: Any) -> FakeResponse: + nonlocal switch_body + assert kwargs.get("token") is None + if method == "GET": + return FakeResponse(payload={"workspaces": [{"id": "org_b", "name": "Account B"}]}) + assert path == "/workspaces/org_b/token" + switch_body = kwargs.get("body") + return FakeResponse( + payload={ + "api_token": "account-b-token", + "organization_id": "org_b", + "organization_name": "Account B", + "email": "account-b@example.test", + "scopes": ["scans:read", "organizations:read"], + } + ) + + monkeypatch.setattr(http, "request", fake_request) + assert cloud.run_cloud(["workspaces", "use", "Account B", "--json"]) == 0 + assert switch_body is None + record = platform_cli.read_record() + assert record is not None + assert record["api_token"] == "account-b-token" + assert record["organization_id"] == "org_b" + assert record["email"] == "account-b@example.test" + assert record["requested_scopes"] == ["scans:read", "organizations:read"] + assert "account-a@example.test" not in auth_path.read_text(encoding="utf-8") + + def test_workspaces_use_reports_unknown_workspace(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( http, @@ -641,7 +1650,89 @@ def test_workspaces_use_reports_unknown_workspace(monkeypatch: pytest.MonkeyPatc assert cloud.run_cloud(["workspaces", "use", "missing", "--json"]) == 1 -def test_group_help_lists_all_verbs_instead_of_default_verb_help(capsys: Any) -> None: +def test_workspaces_use_reports_auth_storage_failure( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse: + if method == "GET": + return FakeResponse(payload={"workspaces": [{"id": "org_1", "name": "Team One"}]}) + assert path == "/workspaces/org_1/token" + return FakeResponse( + payload={ + "api_token": "test-token", + "organization_id": "org_1", + "organization_name": "Team One", + "scopes": ["scans:read"], + } + ) + + monkeypatch.setattr(http, "request", fake_request) + monkeypatch.setattr( + workspaces, "save_record", lambda _record: (_ for _ in ()).throw(OSError("disk full")) + ) + + assert cloud.run_cloud(["workspaces", "use", "1", "--json"]) == http.EXIT_ERROR + payload = json.loads(capsys.readouterr().out) + assert "could not be stored" in payload["error"] + assert payload["workspace_switched"] is True + assert payload["local_record_updated"] is False + assert payload["retry_safe"] is True + + +@pytest.mark.parametrize( + "failure", + [ + requests.ConnectionError("connection reset"), + FakeResponse(status_code=503, text="temporarily unavailable"), + FakeResponse(status_code=200, text="not JSON"), + FakeResponse(status_code=200, payload={"organization_id": "org_1"}), + ], +) +def test_workspace_use_reports_retry_safe_unknown_outcomes( + failure: Exception | FakeResponse, + monkeypatch: pytest.MonkeyPatch, + capsys: Any, +) -> None: + def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse: + if method == "GET": + return FakeResponse(payload={"workspaces": [{"id": "org_1", "name": "Team One"}]}) + assert path == "/workspaces/org_1/token" + if isinstance(failure, Exception): + raise http.CloudError(str(failure)) from failure + return failure + + monkeypatch.setattr(http, "request", fake_request) + + assert cloud.run_cloud(["workspaces", "use", "1", "--json"]) == http.EXIT_ERROR + payload = json.loads(capsys.readouterr().out) + assert payload["switch_outcome_unknown"] is True + assert payload["retry_safe"] is True + assert "safely rerun" in payload["error"] + + +def test_workspace_use_preserves_definitive_conflict( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + def fake_request(method: str, path: str, **_kwargs: Any) -> FakeResponse: + if method == "GET": + return FakeResponse(payload={"workspaces": [{"id": "org_1", "name": "Team One"}]}) + return FakeResponse( + status_code=409, + payload={"error": {"code": "token_conflict", "message": "token changed"}}, + ) + + monkeypatch.setattr(http, "request", fake_request) + + assert cloud.run_cloud(["workspaces", "use", "1", "--json"]) == http.EXIT_ERROR + payload = json.loads(capsys.readouterr().out) + assert "token changed" in payload["error"] + assert "switch_outcome_unknown" not in payload + + +def test_group_help_lists_all_verbs_instead_of_default_verb_help( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(render.sys.stdout, "isatty", lambda: True) assert cloud.run_cloud(["workspaces", "-h"]) == 0 output = capsys.readouterr().out assert "workspaces verbs" in output @@ -687,6 +1778,47 @@ def test_workspace_human_list_is_numbered_and_hides_ids( assert "workspaces use NUMBER" in output +def test_integrations_human_list_exposes_installation_id_and_json_stays_full( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + payload = { + "integrations": [ + { + "id": "integration-uuid", + "organization_id": "org-secret", + "connected_by": "user-secret", + "provider": "github", + "installation_id": 154419799, + "account_login": "usestrix", + "repository_selection": "selected", + "connected_at": "2026-08-27T12:00:00Z", + } + ], + "merge_accounts": [ + { + "id": "merge-uuid", + "provider": "jira", + "status": "linked", + "default_collection_name": "Security", + } + ], + "bitbucket_oauth_enabled": True, + } + monkeypatch.setattr(render.sys.stdout, "isatty", lambda: True) + monkeypatch.setattr(http, "request", lambda *_a, **_k: FakeResponse(payload=payload)) + + assert cloud.run_cloud(["integrations", "list"]) == 0 + output = capsys.readouterr().out + for value in ("1.", "2.", "github", "usestrix", "154419799", "jira", "Security"): + assert value in output + for value in ("integration-uuid", "merge-uuid", "org-secret", "user-secret"): + assert value not in output + assert "--installation-id INSTALLATION_ID" in output + + assert cloud.run_cloud(["integrations", "list", "--json"]) == 0 + assert json.loads(capsys.readouterr().out) == payload + + def test_pr_review_human_list_prioritizes_actionable_fields( monkeypatch: pytest.MonkeyPatch, capsys: Any ) -> None: @@ -710,6 +1842,7 @@ def test_pr_review_human_list_prioritizes_actionable_fields( "verdict": "pass", "status": "posted", "findings_count": 0, + "open_findings_count": 0, } ], "meta": {"total": 1}, @@ -719,7 +1852,17 @@ def test_pr_review_human_list_prioritizes_actionable_fields( assert cloud.run_cloud(["pr-reviews", "list"]) == 0 output = capsys.readouterr().out - for value in ("usestrix/strix", "1177", "Improve cloud CLI", "feature", "main", "pass"): + for value in ( + "usestrix/strix", + "1177", + "Improve cloud CLI", + "feature", + "main", + "posted", + "pass", + "0 open / 0 total", + "review-id", + ): assert value in output for value in ("org-id", "user-id", "installation_id"): assert value not in output @@ -774,9 +1917,9 @@ def test_workspace_use_accepts_list_number(monkeypatch: pytest.MonkeyPatch, tmp_ } ) return FakeResponse( - status_code=201, + status_code=200, payload={ - "api_token": "new", + "api_token": "old", "organization_id": "org_2", "organization_name": "Two", "scopes": ["organizations:read", "tokens:write"], @@ -786,6 +1929,9 @@ def test_workspace_use_accepts_list_number(monkeypatch: pytest.MonkeyPatch, tmp_ monkeypatch.setattr(http, "request", fake_request) assert cloud.run_cloud(["workspaces", "use", "2", "--json"]) == 0 assert called_paths == ["/workspaces", "/workspaces/org_2/token"] + record = platform_cli.read_record() + assert record is not None + assert record["api_token"] == "old" def test_logout_help_does_not_remove_stored_auth( @@ -797,7 +1943,7 @@ def test_logout_help_does_not_remove_stored_auth( assert cloud.run_cloud(["logout", "--help"]) == 0 assert platform_cli.read_record() == {"api_token": "keep-me"} - assert "Usage:" in capsys.readouterr().out + assert "usage: strix cloud logout" in capsys.readouterr().out def test_logout_rejects_unknown_arguments_without_removing_stored_auth( diff --git a/tests/test_cloud_cli_runtime.py b/tests/test_cloud_cli_runtime.py new file mode 100644 index 00000000..842aaa25 --- /dev/null +++ b/tests/test_cloud_cli_runtime.py @@ -0,0 +1,1111 @@ +"""Focused regressions for managed-cloud CLI rendering and runtime safety.""" + +from __future__ import annotations + +import argparse +import io +import json +import sys +from typing import TYPE_CHECKING, Any + +import pytest +from rich.console import Console + +from strix.interface import cloud, platform_cli +from strix.interface.cloud import http, render, runner, source_scan +from strix.interface.main import main as interface_main + + +if TYPE_CHECKING: + from pathlib import Path + + +class FakeResponse: + def __init__( + self, + payload: Any = None, + *, + status_code: int = 200, + content: bytes | None = None, + content_type: str | None = None, + ) -> None: + self._payload = payload + self.status_code = status_code + self.ok = 200 <= status_code < 400 + self.content = content if content is not None else json.dumps(payload).encode() + self.text = self.content.decode("utf-8", errors="replace") + self.headers = { + "content-type": content_type + or ("application/json" if payload is not None else "application/octet-stream") + } + self.closed = False + + def json(self) -> Any: + if self._payload is None: + raise ValueError("not JSON") + 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", "runtime-test-token") + + +def _console_output(data: Any, *, view: str | None = None, width: int = 160) -> str: + output = io.StringIO() + render.emit(Console(file=output, width=width), data, as_json=False, view=view) + return output.getvalue() + + +def test_paginated_envelope_with_scalar_metadata_is_a_compact_list() -> None: + output = _console_output( + { + "items": [{"id": "scan-1", "status": "running", "title": "Acceptance"}], + "scansThisMonth": 42, + "total": 1, + }, + view="GET /scans", + ) + assert "Acceptance" in output + assert "running" in output + assert "scansThisMonth" not in output + assert "42" not in output + assert len(output.splitlines()) < 10 + + +@pytest.mark.parametrize("width", [80, 160]) +def test_scan_lists_flatten_actionable_targets_findings_and_ids(width: int) -> None: + output = _console_output( + { + "items": [ + { + "id": "scan-code-uuid", + "title": "Source review", + "engagement_type": "code_review", + "scan_type": "ultra", + "status": "running", + "repositories": [ + { + "url": "https://github.com/usestrix/strix", + "branch": "feature", + "provider": "github", + } + ], + "findings": {"total": 4, "critical": 1, "high": 3}, + "created_at": "2026-08-28T12:00:00Z", + }, + { + "id": "scan-live-uuid", + "title": "Staging pentest", + "engagement_type": "live_test", + "scan_type": "ultra", + "status": "completed", + "urls": ["https://staging.example.test"], + "findings": {"total": 2, "high": 2}, + "created_at": "2026-08-28T13:00:00Z", + }, + ], + "total": 2, + }, + view="GET /scans", + width=width, + ) + for value in ( + "https://github.com/usestrix/strix @ feature", + "https://staging.example.test", + "running", + "completed", + "scan-code-uuid", + "scan-live-uuid", + ): + assert value in output + assert "findings" in output + assert "4" in output + assert "2" in output + assert "scans get ID" in output + + +@pytest.mark.parametrize("width", [80, 160]) +def test_vulnerability_lists_always_include_the_actionable_uuid(width: int) -> None: + vulnerability_id = "12345678-1234-4321-8765-123456789abc" + output = _console_output( + { + "items": [ + { + "id": vulnerability_id, + "scan_id": "internal-scan-id", + "title": "SQL injection", + "target": "https://example.test/search", + "severity": "critical", + "cve": "CVE-2026-0001", + "cvss": 9.8, + "status": "open", + "created_at": "2026-08-28T12:00:00Z", + "dependency_metadata": {"package": "example"}, + "display_number": "VULN-42", + "finding_type": "dast", + } + ] + }, + view="GET /vulnerabilities", + width=width, + ) + assert vulnerability_id in output + assert "internal-scan-id" not in output + assert "SQL injection" in output + assert "critical" in output + assert "vulns get ID" in output + + +def test_connector_enrollment_command_is_complete_multiline_and_terminal_safe( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + command = ( + "docker run --rm \\\n" + " -e TS_AUTHKEY=tskey-" + "a" * 180 + " \\\n" + " -e LABEL=before\x1b]52;c;copied\x07after \\\n" + " ghcr.io/usestrix/connector:latest" + ) + monkeypatch.setattr(render.sys.stdout, "isatty", lambda: True) + monkeypatch.setattr( + http, + "request", + lambda *_args, **_kwargs: FakeResponse( + {"id": "connector-1", "name": "Private network", "docker_command": command} + ), + ) + + assert cloud.run_cloud(["connectors", "get", "connector-1", "--include-command"]) == 0 + output = capsys.readouterr().out + assert "TS_AUTHKEY=tskey-" + "a" * 180 in output + assert "ghcr.io/usestrix/connector:latest" in output + assert "\\x0a" not in output + assert "\x1b" not in output + assert "\\x1b]52;c;copied\\x07" in output + + +def test_detail_field_cap_counts_only_populated_values() -> None: + payload = {**{f"unused_{index}": None for index in range(40)}, "result": "visible"} + output = _console_output(payload) + assert "visible" in output + assert "additional field" not in output + + +def test_empty_envelope_has_a_clear_empty_state() -> None: + output = _console_output({"webhooks": [], "total": 0}, view="GET /webhooks") + assert "No items." in output + assert "{}" not in output + + +def test_pr_detail_summarizes_nested_findings() -> None: + findings = [{"severity": "high", "title": f"Finding {index}"} for index in range(12)] + output = _console_output( + { + "id": "review-1", + "repository_full_name": "usestrix/strix", + "pr_number": 1177, + "findings": findings, + }, + view="GET /pr-reviews/{reviewId}", + ) + assert "usestrix/strix" in output + assert "Finding 0" in output + assert "7 more; use --json" in output + assert "Finding 11" not in output + assert len(output.splitlines()) < 25 + + +def test_analytics_views_are_bounded_and_frequency_prefers_activity() -> None: + overview = { + f"section_{index}": {f"metric_{inner}": inner for inner in range(10)} for index in range(20) + } + overview_output = _console_output(overview, view="GET /analytics/overview") + assert "Showing 36 of 200 summary metrics" in overview_output + assert len(overview_output.splitlines()) < 45 + + points = [{"date": f"day-{index}", "count": 0} for index in range(300)] + points[100]["count"] = 3 + points[250]["count"] = 7 + frequency_output = _console_output( + {"items": points, "total": 300}, view="GET /analytics/scan-frequency" + ) + assert "day-100" in frequency_output + assert "day-250" in frequency_output + assert "day-299" not in frequency_output + assert "2 non-zero point(s) from 300 total" in frequency_output + assert len(frequency_output.splitlines()) < 15 + + +def test_nested_webhook_envelope_and_events_render_cleanly() -> None: + output = _console_output( + { + "data": { + "webhooks": [ + { + "id": "hook-1", + "url": "https://example.test/hook", + "events": ["scan.completed", "finding.created"], + "is_active": True, + } + ], + "pagination": {"page": 1, "total": 1}, + } + }, + view="GET /webhooks", + ) + assert "https://example.test/hook" in output + assert "scan.completed, finding.created" in output + assert "pagination" not in output + + +def test_human_rendering_neutralizes_osc_and_csi_control_sequences() -> None: + dangerous = "before\x1b]52;c;copied\x07after\x1b[2J\x9b31m" + outputs = ( + _console_output(dangerous), + _console_output([{"name": dangerous}], view="GET /scans"), + _console_output({f"field{dangerous}": dangerous}), + _console_output({"source": {"files": [dangerous]}}, view="source_manifest"), + ) + + for output in outputs: + assert "\x1b" not in output + assert "\x07" not in output + assert "\x9b" not in output + assert "\\x1b]52;c;copied\\x07" in output + assert "\\x1b[2J\\x9b31m" in output + + +def test_source_prompt_shows_paths_and_literal_confirmation( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + (tmp_path / "app.py").write_text("print('ok')\n", encoding="utf-8") + dangerous_name = "visible\x1b]52;c;copied\x07\x1b[2J.py" + (tmp_path / dangerous_name).write_text("print('safe')\n", encoding="utf-8") + output = io.StringIO() + console = Console(file=output, width=100) + prompts: list[tuple[str, bool]] = [] + + def answer(prompt: str, *, markup: bool = True, **_kwargs: Any) -> str: + prompts.append((prompt, markup)) + return "n" + + monkeypatch.setattr(console, "input", answer) + monkeypatch.setattr(source_scan.sys.stdin, "isatty", lambda: True) + monkeypatch.setattr(source_scan.sys.stdout, "isatty", lambda: True) + args = argparse.Namespace( + source=str(tmp_path), + dry_run=False, + show_files=True, + include_hidden=False, + include_sensitive=False, + include_archives=False, + exclude=[], + yes=False, + ) + with pytest.raises(http.CloudError, match="cancelled"): + source_scan.prepare_scan_source(console, args, as_json=False) + rendered = output.getvalue() + assert "app.py" in rendered + assert "\x1b" not in rendered + assert "\x07" not in rendered + assert "visible\\x1b]52;c;copied\\x07\\x1b[2J.py" in rendered + assert prompts == [("Upload this source and start the scan? [y/N]: ", False)] + + +@pytest.mark.parametrize("failure", [KeyboardInterrupt(), EOFError()]) +def test_source_prompt_interruption_removes_temporary_archive( + failure: BaseException, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + (tmp_path / "app.py").write_text("print('ok')\n", encoding="utf-8") + console = Console(file=io.StringIO(), width=100) + archive_paths: list[Path] = [] + original_prepare = source_scan.prepare_source + + def capture_bundle(*args: Any, **kwargs: Any) -> Any: + bundle = original_prepare(*args, **kwargs) + archive_paths.append(bundle.archive_path) + return bundle + + def interrupt(*_args: Any, **_kwargs: Any) -> str: + raise failure + + monkeypatch.setattr(source_scan, "prepare_source", capture_bundle) + monkeypatch.setattr(console, "input", interrupt) + monkeypatch.setattr(source_scan.sys.stdin, "isatty", lambda: True) + monkeypatch.setattr(source_scan.sys.stdout, "isatty", lambda: True) + args = argparse.Namespace( + source=str(tmp_path), + dry_run=False, + show_files=False, + include_hidden=False, + include_sensitive=False, + include_archives=False, + exclude=[], + yes=False, + ) + + with pytest.raises(type(failure)): + source_scan.prepare_scan_source(console, args, as_json=False) + + assert archive_paths + assert all(not path.exists() for path in archive_paths) + + +def test_human_error_neutralizes_terminal_control_sequences() -> None: + output = io.StringIO() + console = Console(file=output, width=100) + runner._emit_error( + console, + http.CloudError("failed\x1b]52;c;copied\x07\x1b[2J"), + as_json=False, + ) + rendered = output.getvalue() + assert "\x1b" not in rendered + assert "\x07" not in rendered + assert "failed\\x1b]52;c;copied\\x07\\x1b[2J" in rendered + + +def test_session_human_output_neutralizes_server_control_sequences() -> None: + dangerous = "value\x1b]52;c;copied\x07\x1b[2J" + output = io.StringIO() + console = Console(file=output, width=100) + + platform_cli._print_success( + console, + { + "email": dangerous, + "organization_name": dangerous, + "scopes": [dangerous], + }, + ) + + rendered = output.getvalue() + assert "\x1b" not in rendered + assert "\x07" not in rendered + assert rendered.count("\\x1b]52;c;copied\\x07\\x1b[2J") == 3 + + +def test_device_login_rejects_non_http_verification_url( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + platform_cli.requests, + "post", + lambda *_a, **_k: FakeResponse( + { + "device_code": "device-1", + "user_code": "ABCD", + "verification_uri": "javascript:alert(1)", + "expires_in": 300, + } + ), + ) + + with pytest.raises(platform_cli.PlatformAuthError, match="invalid verification URL"): + platform_cli._run_device_flow(Console(file=io.StringIO()), open_browser=False) + + +@pytest.mark.parametrize("value", ["0", "-1", "nan", "inf"]) +def test_invalid_timeout_is_a_usage_error_without_request( + value: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(http, "request", lambda *_a, **_k: pytest.fail("must not request")) + assert cloud.run_cloud(["scans", "list", "--timeout", value, "--json"]) == 2 + + +def test_boolean_query_values_are_lowercase_for_url_search_params( + monkeypatch: pytest.MonkeyPatch, +) -> None: + seen: dict[str, Any] = {} + + def fake_request(_method: str, _url: str, **kwargs: Any) -> FakeResponse: + seen["params"] = kwargs.get("params") + return FakeResponse({"items": []}) + + monkeypatch.setattr(http.requests, "request", fake_request) + http.request("GET", "/test", query={"enabled": True, "disabled": False}) + assert seen["params"] == {"enabled": "true", "disabled": "false"} + + +@pytest.mark.parametrize("value", ["0", "-1", "nan", "inf"]) +def test_workspace_use_invalid_timeout_is_clean( + value: str, monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(http, "request", lambda *_a, **_k: pytest.fail("must not request")) + assert cloud.run_cloud(["workspaces", "use", "1", "--timeout", value, "--json"]) == 2 + assert "greater than 0" in json.loads(capsys.readouterr().out)["error"] + + +def test_json_argument_errors_do_not_leak_argparse_prose( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(http, "request", lambda *_a, **_k: pytest.fail("must not request")) + + assert cloud.run_cloud(["scans", "get", "--json"]) == http.EXIT_USAGE + captured = capsys.readouterr() + payload = json.loads(captured.out) + assert "SCAN_ID" in payload["error"] + assert captured.err == "" + + assert cloud.run_cloud(["workspaces", "use", "--json"]) == http.EXIT_USAGE + captured = capsys.readouterr() + payload = json.loads(captured.out) + assert "WORKSPACE" in payload["error"] + assert captured.err == "" + + +@pytest.mark.parametrize("command", ["whoami", "logout"]) +def test_redirected_session_argument_errors_are_json_only(command: str, capsys: Any) -> None: + assert cloud.run_cloud([command, "--bogus"]) == http.EXIT_USAGE + captured = capsys.readouterr() + payload = json.loads(captured.out) + assert f"strix cloud {command}" in payload["error"] + assert captured.err == "" + + +@pytest.mark.parametrize( + "content_type,payload", + [ + ("application/pdf", b"%PDF-1.7\n\x1b]52;c;copied\x07"), + ("application/zip", b"PK\x03\x04\x1b[2J"), + ], +) +def test_binary_response_refuses_to_write_to_a_terminal( + content_type: str, + payload: bytes, + monkeypatch: pytest.MonkeyPatch, + capsys: Any, +) -> None: + monkeypatch.setattr(runner.sys.stdout, "isatty", lambda: True) + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse(content=payload, content_type=content_type), + ) + + assert cloud.run_cloud(["scans", "report", "scan-1"]) == http.EXIT_USAGE + output = capsys.readouterr().out + assert "binary responses" in output + assert "--output FILE" in output + assert "\x1b" not in output + + +def test_binary_response_can_be_intentionally_redirected( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class RedirectedStdout: + def __init__(self) -> None: + self.buffer = io.BytesIO() + + @staticmethod + def isatty() -> bool: + return False + + def write(self, value: str) -> int: + return len(value) + + def flush(self) -> None: + return None + + redirected = RedirectedStdout() + monkeypatch.setattr(runner.sys, "stdout", redirected) + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + content=b"PK\x03\x04archive", content_type="application/zip" + ), + ) + + assert cloud.run_cloud(["chat", "files", "archive", "chat-1"]) == 0 + assert redirected.buffer.getvalue() == b"PK\x03\x04archive" + + +@pytest.mark.parametrize("failure", ["rejected", "interrupted"]) +def test_redirected_binary_errors_never_append_diagnostics_to_stdout( + failure: str, monkeypatch: pytest.MonkeyPatch, capsysbinary: Any +) -> None: + class InterruptedResponse(FakeResponse): + def iter_content(self, *, chunk_size: int) -> Any: + assert chunk_size == 1024 * 1024 + yield b"%PDF-partial" + raise http.requests.ConnectionError("connection lost") + + response = ( + FakeResponse({"detail": "report rejected"}, status_code=500) + if failure == "rejected" + else InterruptedResponse(content=b"unused", content_type="application/pdf") + ) + monkeypatch.setattr(http, "request", lambda *_a, **_k: response) + + assert cloud.run_cloud(["scans", "report", "scan-1"]) == http.EXIT_ERROR + captured = capsysbinary.readouterr() + assert captured.out == (b"" if failure == "rejected" else b"%PDF-partial") + assert b"Error:" in captured.err + expected = b"report rejected" if failure == "rejected" else b"connection lost" + assert expected in captured.err + + +def test_redirected_binary_parse_errors_go_only_to_stderr( + monkeypatch: pytest.MonkeyPatch, capsysbinary: Any +) -> None: + monkeypatch.setattr( + http, "request", lambda *_a, **_k: pytest.fail("invalid usage must not request a report") + ) + + assert cloud.run_cloud(["scans", "report", "scan-1", "--not-an-option"]) == http.EXIT_USAGE + captured = capsysbinary.readouterr() + assert captured.out == b"" + assert b"invalid arguments" in captured.err + + +def test_explicit_json_binary_response_requires_an_output_file( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: pytest.fail("usage must be rejected before downloading"), + ) + + assert cloud.run_cloud(["scans", "report", "scan-1", "--json"]) == http.EXIT_USAGE + payload = json.loads(capsys.readouterr().out) + assert "requires --output FILE" in payload["error"] + + +def test_binary_output_can_return_json_download_metadata( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse(content=b"%PDF-report", content_type="application/pdf"), + ) + target = tmp_path / "report.pdf" + + assert ( + cloud.run_cloud(["scans", "report", "scan-1", "--output", str(target), "--json"]) + == http.EXIT_OK + ) + payload = json.loads(capsys.readouterr().out) + assert payload == { + "output": str(target), + "bytes": len(b"%PDF-report"), + "content_type": "application/pdf", + } + assert target.read_bytes() == b"%PDF-report" + + +def test_binary_download_creates_parents_and_requires_force( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + contents = iter((b"first", b"second", b"third")) + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse(content=next(contents), content_type="application/pdf"), + ) + target = tmp_path / "nested" / "report.pdf" + command = ["scans", "report", "scan-1", "--output", str(target)] + assert cloud.run_cloud(command) == 0 + assert target.read_bytes() == b"first" + assert cloud.run_cloud(command) == 1 + assert target.read_bytes() == b"first" + assert cloud.run_cloud([*command, "--force"]) == 0 + assert target.read_bytes() == b"third" + + +def test_binary_download_bad_parent_is_a_clean_error( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse(content=b"report", content_type="application/pdf"), + ) + blocker = tmp_path / "not-a-directory" + blocker.write_text("x", encoding="utf-8") + assert ( + cloud.run_cloud(["scans", "report", "scan-1", "--output", str(blocker / "report.pdf")]) == 1 + ) + assert "could not write" in capsys.readouterr().out + + +def test_binary_download_streams_and_preserves_existing_file_on_failure( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + class InterruptedResponse(FakeResponse): + closed = False + + def iter_content(self, *, chunk_size: int) -> Any: + assert chunk_size == 1024 * 1024 + yield b"partial" + raise http.requests.ConnectionError("connection lost") + + def close(self) -> None: + self.closed = True + + response = InterruptedResponse(content=b"must not be buffered", content_type="application/pdf") + seen: dict[str, Any] = {} + + def fake_request(_method: str, _path: str, **kwargs: Any) -> FakeResponse: + seen["stream"] = kwargs.get("stream") + return response + + monkeypatch.setattr(http, "request", fake_request) + target = tmp_path / "report.pdf" + target.write_bytes(b"original") + assert ( + cloud.run_cloud(["scans", "report", "scan-1", "--output", str(target), "--force", "--json"]) + == 1 + ) + assert seen["stream"] is True + assert target.read_bytes() == b"original" + assert response.closed is True + assert list(tmp_path.iterdir()) == [target] + + +def test_audit_csv_downloads_while_json_remains_parsed( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + responses = iter( + ( + FakeResponse(content=b"action,actor\nlogin,alex\n", content_type="text/csv"), + FakeResponse(payload={"items": [{"action": "login"}], "total": 1}), + ) + ) + monkeypatch.setattr(http, "request", lambda *_a, **_k: next(responses)) + target = tmp_path / "exports" / "audit.csv" + assert cloud.run_cloud(["audit", "list", "--format", "csv", "--output", str(target)]) == 0 + assert target.read_text(encoding="utf-8") == "action,actor\nlogin,alex\n" + capsys.readouterr() + assert cloud.run_cloud(["audit", "list", "--format", "json", "--json"]) == 0 + assert json.loads(capsys.readouterr().out)["items"][0]["action"] == "login" + + +@pytest.mark.parametrize("format_name", ["ndjson", "jsonl", "snowflake", "splunk"]) +def test_audit_ndjson_compatible_formats_download_raw( + format_name: str, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse( + content=b'{"action":"login"}\n', content_type="application/x-ndjson" + ), + ) + target = tmp_path / f"audit-{format_name}.ndjson" + assert cloud.run_cloud(["audit", "list", "--format", format_name, "--output", str(target)]) == 0 + assert target.read_bytes() == b'{"action":"login"}\n' + + +def test_credit_limit_payload_code_maps_to_payment_exit() -> None: + response = FakeResponse( + {"code": "scan_credit_limit_reached", "detail": "monthly scan credits exhausted"}, + status_code=403, + ) + with pytest.raises(http.CloudError) as raised: + http.check(response) # type: ignore[arg-type] + assert raised.value.exit_code == http.EXIT_PAYMENT + + +def test_wait_timeout_is_bounded(monkeypatch: pytest.MonkeyPatch, capsys: Any) -> None: + def response(method: str, _path: str, **_kwargs: Any) -> FakeResponse: + if method == "POST": + return FakeResponse({"id": "scan-1", "status": "pending"}) + return FakeResponse({"id": "scan-1", "status": "running"}) + + monkeypatch.setattr(http, "request", response) + assert ( + cloud.run_cloud( + [ + "scans", + "start", + "--domain-ids", + "domain-1", + "--wait", + "--wait-timeout", + "0.000001", + "--json", + ] + ) + == 1 + ) + assert "wait timed out" in json.loads(capsys.readouterr().out)["error"] + + +def test_wait_interruption_returns_130_with_remote_operation_id( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + def response(method: str, _path: str, **_kwargs: Any) -> FakeResponse: + if method == "POST": + return FakeResponse({"id": "scan-interrupted", "status": "pending"}) + raise KeyboardInterrupt + + monkeypatch.setattr(http, "request", response) + assert ( + cloud.run_cloud(["scans", "start", "--domain-ids", "domain-1", "--wait", "--json"]) == 130 + ) + payload = json.loads(capsys.readouterr().out) + assert payload["interrupted"] is True + assert payload["status_unknown"] is True + assert payload["operation_id"] == "scan-interrupted" + + +def test_session_help_is_specific_and_human_whoami_shows_scopes( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) + monkeypatch.setattr(platform_cli, "AUTH_PATH", tmp_path / "auth.json") + platform_cli.save_record( + { + "api_token": "secret", + "email": "alex@example.test", + "organization_name": "Demo", + "scopes": ["scans:read", "organizations:read"], + } + ) + monkeypatch.setattr(platform_cli.sys.stdout, "isatty", lambda: True) + assert cloud.run_cloud(["whoami", "--help"]) == 0 + who_help = capsys.readouterr().out + assert "strix cloud whoami" in who_help + assert "--no-browser" not in who_help + assert cloud.run_cloud(["whoami"]) == 0 + assert "scans:read organizations:read" in capsys.readouterr().out + + +def test_non_tty_whoami_and_logout_emit_json( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.delenv("STRIX_API_TOKEN", raising=False) + monkeypatch.setattr(platform_cli, "AUTH_PATH", tmp_path / "auth.json") + monkeypatch.setattr(platform_cli.sys.stdout, "isatty", lambda: False) + platform_cli.save_record( + { + "api_token": "secret", + "email": "agent@example.test", + "organization_name": "Demo", + "scopes": ["scans:read"], + } + ) + + assert cloud.run_cloud(["whoami"]) == 0 + assert json.loads(capsys.readouterr().out)["email"] == "agent@example.test" + + assert cloud.run_cloud(["logout"]) == 0 + assert json.loads(capsys.readouterr().out) == {"signed_in": False, "removed": True} + assert not platform_cli.AUTH_PATH.exists() + + +def test_scope_picker_labels_match_the_server_presets() -> None: + output = io.StringIO() + console = Console(file=output, width=120) + console.input = lambda *_args, **_kwargs: "1" # type: ignore[method-assign] + assert ( + platform_cli._choose_scopes( + console, + [{"scope": "scans:read", "min_role": "viewer", "minimum": True}], + "admin", + ) + is None + ) + rendered = output.getvalue() + assert "uploads" in rendered + assert "workspace switching" in rendered + assert "scan read/write and billing read" in rendered + + +def test_noninteractive_login_never_prompts_for_workspace( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(platform_cli.sys.stdin, "isatty", lambda: False) + console = Console(file=io.StringIO()) + console.input = lambda *_args, **_kwargs: pytest.fail("must not prompt") # type: ignore[method-assign] + + with pytest.raises(platform_cli.PlatformAuthError, match="--workspace NAME_OR_ID"): + platform_cli._choose_workspace( + console, + [ + {"id": "org-1", "name": "One"}, + {"id": "org-2", "name": "Two"}, + ], + None, + ) + + +def test_login_workspace_selector_prefers_ids_and_rejects_duplicate_names() -> None: + console = Console(file=io.StringIO()) + organizations = [ + {"id": "org_1", "name": "org_2"}, + {"id": "org_2", "name": "Strix"}, + {"id": "org_3", "name": "Strix"}, + ] + + assert platform_cli._choose_workspace(console, organizations, "org_2")["id"] == "org_2" + with pytest.raises(platform_cli.PlatformAuthError, match="org_2, org_3"): + platform_cli._choose_workspace(console, organizations, "Strix") + + +def test_device_flow_slow_down_never_exceeds_the_poll_interval_cap( + monkeypatch: pytest.MonkeyPatch, +) -> None: + authorization = FakeResponse( + { + "user_code": "ABCD-EFGH", + "verification_uri": "https://example.test/device", + "device_code": "device-1", + "expires_in": 1000, + "interval": 60, + } + ) + polls = iter( + [ + FakeResponse({"error": "slow_down"}, status_code=400), + FakeResponse({"error": "slow_down"}, status_code=400), + FakeResponse({"error": "expired_token"}, status_code=400), + ] + ) + calls = 0 + + def post(*_args: Any, **_kwargs: Any) -> FakeResponse: + nonlocal calls + calls += 1 + return authorization if calls == 1 else next(polls) + + now = 0.0 + sleeps: list[float] = [] + + def monotonic() -> float: + return now + + def sleep(seconds: float) -> None: + nonlocal now + sleeps.append(seconds) + now += seconds + + monkeypatch.setattr(platform_cli.requests, "post", post) + monkeypatch.setattr(platform_cli, "_app_url", lambda: "https://example.test") + monkeypatch.setattr(platform_cli.time, "monotonic", monotonic) + monkeypatch.setattr(platform_cli.time, "sleep", sleep) + + with pytest.raises(platform_cli.PlatformAuthError, match="expired"): + platform_cli._run_device_flow( + Console(file=io.StringIO()), + open_browser=False, + scopes=["scans:read"], + ) + assert sleeps == [60, 60, 60] + + +def test_device_flow_accepts_external_authkit_url_and_binds_token_origin( + monkeypatch: pytest.MonkeyPatch, +) -> None: + responses = iter( + [ + FakeResponse( + { + "user_code": "ABCD-EFGH", + "verification_uri": "https://auth.example-workos.com/device?code=ABCD", + "device_code": "device-1", + "expires_in": 300, + "interval": 1, + } + ), + FakeResponse( + { + "api_token": "strix_pat_test", + "organization_id": "org-1", + "scopes": ["scans:read"], + } + ), + ] + ) + monkeypatch.setattr(platform_cli.requests, "post", lambda *_a, **_k: next(responses)) + monkeypatch.setattr(platform_cli, "_app_url", lambda: "https://preview.strix.ai") + monkeypatch.setattr(platform_cli.time, "sleep", lambda _seconds: None) + + record = platform_cli._run_device_flow( + Console(file=io.StringIO()), + open_browser=False, + scopes=["scans:read"], + ) + + assert record["app_url"] == "https://preview.strix.ai" + assert record["requested_scopes"] == ["scans:read"] + + +def test_missing_verb_json_is_structured(capsys: Any) -> None: + assert cloud.run_cloud(["scans", "--json"]) == 0 + payload = json.loads(capsys.readouterr().out) + assert payload["command"] == "strix cloud scans" + assert any(item["name"] == "start" for item in payload["verbs"]) + + +@pytest.mark.parametrize( + "argv", + [ + ["pr-reviews", "--json", "-h"], + ["pr-reviews", "--help", "--json"], + ["workspaces", "--json", "help"], + ], +) +def test_group_help_accepts_json_and_help_in_either_order(argv: list[str], capsys: Any) -> None: + assert cloud.run_cloud(argv) == 0 + payload = json.loads(capsys.readouterr().out) + assert payload["verbs"] + assert "error" not in payload + + +def test_root_help_accepts_json_before_help_and_leaf_help_stays_specific( + capsys: Any, +) -> None: + assert cloud.run_cloud(["--json", "--help"]) == 0 + assert json.loads(capsys.readouterr().out)["command"] == "strix cloud" + + assert cloud.run_cloud(["scans", "get", "scan-1", "-h"]) == 0 + leaf_help = capsys.readouterr().out + assert "strix cloud scans get" in leaf_help + assert "scans verbs" not in leaf_help + + +def test_non_tty_dispatcher_always_emits_structured_json( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(render.sys.stdout, "isatty", lambda: False) + + assert cloud.run_cloud([]) == 0 + assert json.loads(capsys.readouterr().out)["command"] == "strix cloud" + + assert cloud.run_cloud(["scans"]) == 0 + assert json.loads(capsys.readouterr().out)["command"] == "strix cloud scans" + + assert cloud.run_cloud(["does-not-exist"]) == http.EXIT_USAGE + assert json.loads(capsys.readouterr().out) == {"error": "unknown command: does-not-exist"} + + assert cloud.run_cloud(["scans", "does-not-exist"]) == http.EXIT_USAGE + payload = json.loads(capsys.readouterr().out) + assert payload["command"] == "strix cloud scans" + assert payload["error"] == "unknown verb" + + +@pytest.mark.parametrize( + "signed_url", + [ + "http://project.supabase.co/storage/v1/object/upload/sign/bucket/file", + "https://127.0.0.1/storage/v1/object/upload/sign/bucket/file", + "https://10.0.0.1/storage/v1/object/upload/sign/bucket/file", + "https://app.strix.ai@127.0.0.1/storage/v1/object/upload/sign/bucket/file", + "https://evil.example/storage/v1/object/upload/sign/bucket/file", + "https://project.supabase.co.evil/storage/v1/object/upload/sign/bucket/file", + "https://project.supabase.co/not-storage/file", + ], +) +def test_source_upload_rejects_untrusted_destinations_before_reading_file( + signed_url: str, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + source = tmp_path / "approved.zip" + source.write_bytes(b"approved source") + monkeypatch.setattr(http, "_app_url_override", "https://app.strix.ai") + monkeypatch.setattr( + http.requests, + "put", + lambda *_args, **_kwargs: pytest.fail("an untrusted URL must not receive source bytes"), + ) + + with pytest.raises(http.CloudError, match=r"(untrusted storage origin|storage API|invalid)"): + http.upload_file(signed_url, "upload-token", source) + + +@pytest.mark.parametrize( + "app_url,signed_url", + [ + ( + "https://app.strix.ai", + "https://project-ref.supabase.co/storage/v1/object/upload/sign/bucket/file", + ), + ( + "https://strix.corp.internal", + "https://strix.corp.internal/storage/v1/object/upload/sign/bucket/file", + ), + ( + "http://127.0.0.1:3000", + "http://127.0.0.1:3000/storage/v1/object/upload/sign/bucket/file", + ), + ], +) +def test_source_upload_allows_only_managed_or_same_origin_storage( + app_url: str, signed_url: str, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + source = tmp_path / "approved.zip" + source.write_bytes(b"approved source") + response = FakeResponse({"ok": True}) + request_options: dict[str, Any] = {} + + def put(*_args: Any, **kwargs: Any) -> FakeResponse: + request_options.update(kwargs) + return response + + monkeypatch.setattr(http, "_app_url_override", app_url) + monkeypatch.setattr(http.requests, "put", put) + http.upload_file(signed_url, "upload-token", source) + + assert request_options["allow_redirects"] is False + assert response.closed is True + + +def test_source_upload_refuses_redirects_without_following_them( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + source = tmp_path / "approved.zip" + source.write_bytes(b"approved source") + response = FakeResponse(None, status_code=307) + response.headers["location"] = "http://169.254.169.254/latest/meta-data" + request_options: dict[str, Any] = {} + + def put(*_args: Any, **kwargs: Any) -> FakeResponse: + request_options.update(kwargs) + return response + + monkeypatch.setattr(http, "_app_url_override", "https://app.strix.ai") + monkeypatch.setattr(http.requests, "put", put) + + with pytest.raises(http.CloudError, match="unexpected redirect"): + http.upload_file( + "https://project-ref.supabase.co/storage/v1/object/upload/sign/bucket/file", + "upload-token", + source, + ) + assert request_options["allow_redirects"] is False + assert response.closed is True + + +def test_one_time_api_token_has_save_now_warning( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(render.sys.stdout, "isatty", lambda: True) + monkeypatch.setattr( + http, + "request", + lambda *_a, **_k: FakeResponse({"id": "token-1", "token": "strix_pat_once"}), + ) + assert cloud.run_cloud(["tokens", "create", "--type", "personal", "--name", "test"]) == 0 + output = capsys.readouterr().out + assert "Save this now" in output + assert "shown only once" in output + + +def test_root_help_advertises_cloud_and_completions( + monkeypatch: pytest.MonkeyPatch, capsys: Any +) -> None: + monkeypatch.setattr(sys, "argv", ["strix", "--help"]) + with pytest.raises(SystemExit) as raised: + interface_main() + assert raised.value.code == 0 + output = capsys.readouterr().out + assert "strix cloud" in output + assert "strix completions" in output diff --git a/tests/test_cloud_idempotency.py b/tests/test_cloud_idempotency.py new file mode 100644 index 00000000..80246912 --- /dev/null +++ b/tests/test_cloud_idempotency.py @@ -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" diff --git a/tests/test_cloud_payment_proxy.py b/tests/test_cloud_payment_proxy.py new file mode 100644 index 00000000..c6fb6a79 --- /dev/null +++ b/tests/test_cloud_payment_proxy.py @@ -0,0 +1,136 @@ +"""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 + 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", + }, + ) + + assert result == b'{"ok":true}' + headers = captured[0]["headers"] + assert headers["Authorization"] == "Payment wallet-proof" + assert headers["X-Strix-Authorization"] == "Bearer strix-secret" + 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_allows_only_two_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"{}" + with pytest.raises(urllib.error.HTTPError) as third_request: + _post(wallet_url, b"{}") + + assert third_request.value.code == 429 + assert calls == 2 diff --git a/tests/test_cloud_source_upload.py b/tests/test_cloud_source_upload.py index eeba69a6..341fe05b 100644 --- a/tests/test_cloud_source_upload.py +++ b/tests/test_cloud_source_upload.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +import os import shutil import subprocess import zipfile @@ -101,6 +102,58 @@ def test_hidden_and_sensitive_files_need_separate_opt_ins(tmp_path: Path) -> Non ) +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") @@ -112,6 +165,22 @@ def test_strixignore_and_cli_excludes_are_applied(tmp_path: Path) -> None: 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: @@ -121,6 +190,85 @@ def test_source_limits_expanded_bytes_before_compression( 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: @@ -151,7 +299,44 @@ def test_noninteractive_source_upload_requires_yes( lambda *_args, **_kwargs: pytest.fail("approval must happen before any API request"), ) assert cloud.run_cloud(["scans", "start", "--source", str(tmp_path), "--json"]) == 1 - assert "requires explicit approval" in capsys.readouterr().out + output = capsys.readouterr().out + assert "requires explicit approval" in output + assert "--approve-sha256 " 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( @@ -280,3 +465,150 @@ def test_failed_scan_deletes_completed_source_upload( 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"]) +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) + 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_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 diff --git a/tests/test_completions.py b/tests/test_completions.py index 67c67c6f..69b1a082 100644 --- a/tests/test_completions.py +++ b/tests/test_completions.py @@ -29,6 +29,7 @@ def test_cloud_leaf_flag_candidates_come_from_command_spec() -> None: 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 @@ -40,6 +41,96 @@ def test_boolean_completion_includes_positive_and_negative_flags() -> None: 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: + 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 @@ -47,6 +138,17 @@ def test_completion_scripts_cover_supported_shells(capsys: Any) -> None: 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"]) == 2 - assert "Choose zsh, bash, or fish" in capsys.readouterr().err + 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 diff --git a/tests/test_pricing.py b/tests/test_pricing.py index abff873d..c9e9e5df 100644 --- a/tests/test_pricing.py +++ b/tests/test_pricing.py @@ -14,7 +14,10 @@ def test_resolves_common_bare_model_names() -> None: assert resolve_litellm_model("deepseek-v4-flash") == "deepseek/deepseek-v4-flash" assert resolve_litellm_model("openai/deepseek-v4-flash") == "deepseek/deepseek-v4-flash" assert resolve_litellm_model("grok-4.5") == "xai/grok-4.5" - assert resolve_litellm_model("MiniMax-M3") == "minimax/MiniMax-M3" + # MiniMax-M3 is sold by several LiteLLM providers at different prices, so + # the resolver must not guess from its bare name. A provider-qualified + # model remains deterministic. + assert resolve_litellm_model("minimax/MiniMax-M3") == "minimax/MiniMax-M3" def test_resolver_returns_none_for_unresolvable_model() -> None: