Compare commits

...
Author SHA1 Message Date
yoni 21e719a2eb fix(grok): close the token-endpoint response like codex does
The response was left for the garbage collector, which on Python 3.14 surfaces as "Exception ignored while finalizing" urllib3 noise at interpreter shutdown. codex._post_form already context-manages its response.
2026-08-14 17:33:39 +00:00
yoni 21afdaea9e fix(llm): resolve grok/ models to xai/ for LiteLLM metadata
LiteLLM maps xAI models only provider-qualified, so neither "grok/grok-4" nor bare "grok-4" resolves: subscription runs fell back to the generic 200k context window and an 8k output cap instead of Grok's 256k/256k.

subscription.litellm_model_name() now owns the routing-prefix -> LiteLLM name mapping (grok/ -> xai/, chatgpt/ -> bare) and context_budget uses it.
2026-08-14 17:09:39 +00:00
yoni d4573a3197 fix(tui): show the actual subscription provider in the Go TUI
The Go TUI rewrite hardcoded "ChatGPT subscription" in the stats view and its
snapshot only carried a boolean, so a Grok run was labeled ChatGPT — the same
bug previously fixed on the Python side. Carry the provider label through the
snapshot protocol and render it, falling back to a generic "Subscription".

utils._subscription_label becomes public subscription_label since the TUI
backend now needs it too.
2026-08-13 03:26:22 +00:00
yoni cd3250576c Merge remote-tracking branch 'origin/main' into grok-subscription-oauth 2026-08-13 03:19:17 +00:00
Alex SchapiroandAhmed Allam 8ca0c4a9b8 Fix LiteLLM cost model resolution 2026-08-12 17:26:00 +03:00
Ahmed AllamandAhmed Allam 7cc9fa9faa chore: release v1.5.3 2026-08-10 21:28:52 +03:00
devin-ai-integration[bot]andGitHub 174c16fa26 fix(llm): send OpenRouter app attribution on the request itself (#1045) 2026-08-10 11:24:02 -07:00
Ahmed AllamandAhmed Allam 94a2586aaa fix(container): write the browser profile as root 2026-08-10 10:08:17 +03:00
Ahmed AllamandAhmed Allam 372e27fa17 chore(container): drop explanatory comment 2026-08-10 09:54:49 +03:00
Ahmed AllamandAhmed Allam ad727edd66 fix(container): keep the browser env alive where image ENV is dropped 2026-08-10 09:54:49 +03:00
7b3c8f9b74 fix(container): reclaim abandoned browser sessions (#1034)
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
2026-08-09 16:57:51 -07:00
Ahmed AllamandAhmed Allam ae07af6159 chore: drop explanatory comment 2026-08-09 15:44:16 +03:00
Ahmed AllamandAhmed Allam 649a2e2140 fix(llm): omit parallel_tool_calls on tool-less requests 2026-08-09 15:44:16 +03:00
Ahmed AllamandAhmed Allam 597aae6715 chore: release v1.5.2 2026-08-09 04:29:34 +03:00
yoni 7bdc2424f2 fix(update): strip PyInstaller bootloader env vars before re-exec after self-update 2026-08-07 21:40:39 +00:00
yoni a22c686626 chore(viewer): rebuild static viewer bundle 2026-08-07 21:12:50 +00:00
yoni ab4d6ffa4a Merge origin/main into grok-subscription-oauth 2026-08-07 21:12:00 +00:00
yoni edb0a607bf fix(llm): open the store lock file with O_NOFOLLOW
The predictable lock path was opened with Path.open("w"), following (and
truncating through) a pre-positioned symlink. Open it via os.open with
O_NOFOLLOW and no O_TRUNC, raising StoreLockError on a symlinked lock path.
2026-07-29 18:05:58 +00:00
yoni 42df95b681 fix(llm): write credential store via mkstemp to defeat symlink attacks
Third review pass (security): the store temp file used a predictable
subscription-auth.json.tmp name, so a local attacker could pre-plant a symlink
there and divert the OAuth token write. Create it with tempfile.mkstemp
(random name, mode 0600, no symlink following) in the same directory, then
atomically rename over the target.
2026-07-29 18:00:36 +00:00
yoni 48db7f4d0e fix(llm): fail loudly when the store lock is unavailable
Third review pass: the shared credential store no longer proceeds with an
unlocked read-modify-write when fcntl is missing or flock fails. It now retries
on EINTR and otherwise raises StoreLockError, so concurrent provider
login/refresh/logout can never race by silently skipping the cross-process lock.
2026-07-29 17:55:18 +00:00
yoni 7289153f9b fix(llm): make logout-all atomic and persist provider in run record
Second review pass:
- strix auth logout (all providers) now holds the shared store lock across every provider removal, so a concurrent save/refresh cannot leave one credential behind while reporting all removed.
- _persist_run_record writes subscription_provider alongside auth_mode, matching ReportState, so resumed runs keep their original provider label.
2026-07-29 17:48:33 +00:00
yoni 9fd11eedec fix(llm): harden subscription credential store and provider labeling
Addresses PR review:
- Extract the shared ~/.strix/subscription-auth.json handling into subscription_store, writing tokens owner-only (0600) from creation via os.open instead of chmod-after-write, closing the window where credentials were briefly world/group-readable.
- Serialize read-modify-write across providers and processes with a reentrant lock, so overlapping ChatGPT/Grok save/logout/refresh operations no longer clobber each other.
- Live/TUI stats label the subscription from the persisted run record (falling back to provider-aware settings), so resumed runs no longer mislabel the provider when STRIX_LLM changes.
2026-07-29 17:39:26 +00:00
yoni 5c94872186 fix(viewer): label historical subscription runs by provider
read_run_summary backfills subscription_provider from the recorded provider/model slug (reusing subscription.provider_label) when the field is absent, so runs recorded before it existed still label correctly without a rescan. The viewer no longer defaults to "ChatGPT" when the provider is unknown. Rebuilds the committed viewer bundle.
2026-07-29 17:04:38 +00:00
yoni 218470f14d fix(viewer): show the actual subscription provider name
The Run details panel hardcoded "ChatGPT subscription" for any subscription run. Emit subscription_provider (ChatGPT/Grok) in the run record and render it in the viewer, so Grok runs read "Grok subscription". Rebuilds the committed viewer bundle.
2026-07-29 16:40:06 +00:00
yoni bfceb65a4c feat(llm): opt-in Grok/SuperGrok subscription OAuth
Add Sign in with Grok mirroring the merged ChatGPT/Codex integration: a strix/config/grok.py OAuth module (PKCE loopback flow against auth.x.ai, refresh with cross-process locking, secure ~/.strix/subscription-auth.json store) plus grok/<model> routing through the OpenAI-compatible api.x.ai/v1 endpoint via a bearer-stamping chat-completions client.

strix auth login grok / status / logout become provider-aware; subscription.py shares provider-agnostic auth-mode detection; Grok runs are marked zero-cost like other subscriptions.

Uses a consumer subscription outside xAI own products, which xAI does not officially support; opt-in and experimental.
2026-07-28 21:44:43 +00:00
41 changed files with 1829 additions and 220 deletions
+15
View File
@@ -117,6 +117,21 @@ ENV AGENT_BROWSER_EXECUTABLE_PATH=/usr/bin/chromium
ENV AGENT_BROWSER_USER_AGENT="Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36" ENV AGENT_BROWSER_USER_AGENT="Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
ENV AGENT_BROWSER_ARGS="--disable-blink-features=AutomationControlled,--no-first-run,--no-default-browser-check,--lang=en-US" ENV AGENT_BROWSER_ARGS="--disable-blink-features=AutomationControlled,--no-first-run,--no-default-browser-check,--lang=en-US"
ENV AGENT_BROWSER_SCREENSHOT_DIR=/workspace/.agent-browser-screenshots ENV AGENT_BROWSER_SCREENSHOT_DIR=/workspace/.agent-browser-screenshots
ENV AGENT_BROWSER_IDLE_TIMEOUT_MS=180000
USER root
RUN set -eu; \
{ \
for var in AGENT_BROWSER_EXECUTABLE_PATH AGENT_BROWSER_USER_AGENT \
AGENT_BROWSER_ARGS AGENT_BROWSER_SCREENSHOT_DIR \
AGENT_BROWSER_IDLE_TIMEOUT_MS; do \
eval "value=\${$var}"; \
printf 'export %s="${%s:-%s}"\n' "$var" "$var" "$value"; \
done; \
} > /tmp/agent-browser.sh; \
install -m 0644 /tmp/agent-browser.sh /etc/profile.d/agent-browser.sh; \
rm /tmp/agent-browser.sh; \
env -i bash -lc 'test "${AGENT_BROWSER_IDLE_TIMEOUT_MS}" = "180000"'
USER pentester
RUN /home/pentester/.npm-global/bin/agent-browser doctor --offline --quick RUN /home/pentester/.npm-global/bin/agent-browser doctor --offline --quick
RUN set -eux; \ RUN set -eux; \
+8 -1
View File
@@ -1,6 +1,6 @@
[project] [project]
name = "strix-agent" name = "strix-agent"
version = "1.5.1" version = "1.5.3"
description = "Open-source AI Hackers for your apps" description = "Open-source AI Hackers for your apps"
readme = "README.md" readme = "README.md"
license = "Apache-2.0" license = "Apache-2.0"
@@ -230,6 +230,7 @@ ignore = [
# args they intentionally ignore. # args they intentionally ignore.
"tests/test_viewer_auth.py" = ["S105", "S106", "ARG001"] "tests/test_viewer_auth.py" = ["S105", "S106", "ARG001"]
"tests/test_codex_auth.py" = ["S105", "S106", "SLF001"] "tests/test_codex_auth.py" = ["S105", "S106", "SLF001"]
"tests/test_grok_auth.py" = ["S105", "S106", "SLF001"]
# Hatchling loads the build hook by path, not as an importable package. # Hatchling loads the build hook by path, not as an importable package.
"scripts/tui_sidecar_hook.py" = ["INP001"] "scripts/tui_sidecar_hook.py" = ["INP001"]
# Stdlib HTTP handler overrides (do_GET/do_POST). # Stdlib HTTP handler overrides (do_GET/do_POST).
@@ -244,6 +245,9 @@ ignore = [
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a # Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
# circular dependency with strix.telemetry / strix.interface.viewer.report_pdf. # circular dependency with strix.telemetry / strix.interface.viewer.report_pdf.
"strix/interface/viewer/server.py" = ["N802", "PLC0415"] "strix/interface/viewer/server.py" = ["N802", "PLC0415"]
# Lazy import of the TUI live-view projection so importing the viewer does not
# eagerly pull in the Textual TUI.
"strix/interface/viewer/transcript.py" = ["PLC0415"]
# Lazy telemetry import to avoid importing PostHog before the viewer starts. # Lazy telemetry import to avoid importing PostHog before the viewer starts.
"strix/interface/viewer/cli.py" = ["PLC0415"] "strix/interface/viewer/cli.py" = ["PLC0415"]
# Lazy imports inside functions to avoid circular dependency with # Lazy imports inside functions to avoid circular dependency with
@@ -288,6 +292,9 @@ ignore = [
# Heavy inference deps (httpx, openai) imported lazily so auth-status checks # Heavy inference deps (httpx, openai) imported lazily so auth-status checks
# don't pull them in. # don't pull them in.
"strix/config/codex.py" = ["PLC0415"] "strix/config/codex.py" = ["PLC0415"]
"strix/config/grok.py" = ["PLC0415"]
# Lazy ``import fcntl`` so the module imports on non-POSIX platforms.
"strix/config/subscription_store.py" = ["PLC0415"]
# Interface utility branches per scope-mode / target-type combination; # Interface utility branches per scope-mode / target-type combination;
# splitting would obscure the decision tree without simplifying it. # splitting would obscure the decision tree without simplifying it.
"strix/interface/utils.py" = ["PLR0912", "BLE001", "PLC0415"] "strix/interface/utils.py" = ["PLR0912", "BLE001", "PLC0415"]
+7 -1
View File
@@ -263,7 +263,13 @@ Remember: A single well-validated high-impact vulnerability is worth more than d
<multi_agent_system> <multi_agent_system>
AGENT ISOLATION & SANDBOXING: AGENT ISOLATION & SANDBOXING:
- All agents run in the same shared Docker container for efficiency - All agents run in the same shared Docker container for efficiency
- Each agent has its own: browser sessions, terminal sessions - Each agent has its own terminal sessions
- Browsers are NOT per-agent by default: `agent-browser` with no `--session` is one
shared browser, so a concurrent agent's navigation invalidates your page and refs.
Pass `--session <your-agent-name>` for any browser work of your own — then it is
yours alone. Each session is a full Chromium (~340 MB) on this shared box, so keep
one, not several, and `agent-browser --session <name> close` when you're done with
the target; an idle browser is reclaimed automatically after 3 minutes
- All agents share the same /workspace directory and proxy history - All agents share the same /workspace directory and proxy history
- Agents can see each other's files and proxy traffic for better collaboration - Agents can see each other's files and proxy traffic for better collaboration
+18 -47
View File
@@ -16,7 +16,6 @@ import hashlib
import json import json
import logging import logging
import secrets import secrets
import threading
import time import time
import urllib.parse import urllib.parse
from pathlib import Path from pathlib import Path
@@ -24,7 +23,7 @@ from typing import TYPE_CHECKING, Any
import requests import requests
from strix.utils.secret_files import write_secret_text from strix.config import subscription_store
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -54,26 +53,12 @@ _ACCOUNT_CLAIM = "https://api.openai.com/auth"
_TOKEN_TIMEOUT = 30 _TOKEN_TIMEOUT = 30
_EXPIRY_SKEW_S = 300 _EXPIRY_SKEW_S = 300
_refresh_lock = threading.Lock()
# Kept separate from cli-config.json so OAuth tokens never land in the env-var config. # Kept separate from cli-config.json so OAuth tokens never land in the env-var config.
AUTH_PATH = Path.home() / ".strix" / "subscription-auth.json" AUTH_PATH = Path.home() / ".strix" / "subscription-auth.json"
def _read_store() -> dict[str, Any]:
try:
data = json.loads(AUTH_PATH.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return {}
return data if isinstance(data, dict) else {}
def _write_store(data: dict[str, Any]) -> None:
write_secret_text(AUTH_PATH, json.dumps(data, indent=2))
def read_record() -> dict[str, Any] | None: def read_record() -> dict[str, Any] | None:
record = _read_store().get(PROVIDER) record = subscription_store.read(AUTH_PATH).get(PROVIDER)
if not isinstance(record, dict) or record.get("type") != "oauth": if not isinstance(record, dict) or record.get("type") != "oauth":
return None return None
if not (record.get("access") and record.get("refresh") and record.get("account_id")): if not (record.get("access") and record.get("refresh") and record.get("account_id")):
@@ -86,45 +71,31 @@ def is_authenticated() -> bool:
def save_record(record: dict[str, Any]) -> None: def save_record(record: dict[str, Any]) -> None:
data = _read_store() with subscription_store.guard(AUTH_PATH):
data[PROVIDER] = record data = subscription_store.read(AUTH_PATH)
_write_store(data) data[PROVIDER] = record
subscription_store.write(AUTH_PATH, data)
def logout() -> None: def logout() -> None:
data = _read_store() with subscription_store.guard(AUTH_PATH):
if PROVIDER not in data: data = subscription_store.read(AUTH_PATH)
return if PROVIDER not in data:
del data[PROVIDER] return
if data: del data[PROVIDER]
_write_store(data) if data:
return subscription_store.write(AUTH_PATH, data)
with contextlib.suppress(OSError): return
AUTH_PATH.unlink() with contextlib.suppress(OSError):
AUTH_PATH.unlink()
@contextlib.contextmanager @contextlib.contextmanager
def _refresh_guard() -> Iterator[None]: def _refresh_guard() -> Iterator[None]:
"""Serialize token refresh within (lock) and across (flock) Strix processes, """Serialize token refresh within (lock) and across (flock) Strix processes,
so concurrent runs can't both spend the single-use refresh token.""" so concurrent runs can't both spend the single-use refresh token."""
with _refresh_lock: with subscription_store.guard(AUTH_PATH):
try: yield
import fcntl
lock_path = AUTH_PATH.with_suffix(".lock")
lock_path.parent.mkdir(parents=True, exist_ok=True)
handle = lock_path.open("w")
except (ImportError, OSError):
yield
return
try:
with contextlib.suppress(OSError):
fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
yield
finally:
with contextlib.suppress(OSError):
fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
handle.close()
class CodexAuthError(Exception): class CodexAuthError(Exception):
+314
View File
@@ -0,0 +1,314 @@
"""Grok (xAI) subscription auth: OAuth login, token refresh, and the OpenAI
client that routes inference through xAI's API.
Mirrors xAI's Grok CLI: OAuth 2.0 + PKCE against ``auth.x.ai``, with the access
token sent as a ``Bearer`` token to ``api.x.ai/v1`` (OpenAI-compatible, so the
subscription and a metered API key share one endpoint — only the bearer differs).
Using a Grok/SuperGrok subscription outside xAI's own products is not officially
supported by xAI; the user chooses this path knowingly. The OAuth constants are
xAI's own Grok CLI values (the backend only accepts that client).
"""
from __future__ import annotations
import base64
import contextlib
import hashlib
import json
import logging
import secrets
import time
import urllib.parse
from pathlib import Path
from typing import TYPE_CHECKING, Any
import requests
from strix.config import subscription_store
if TYPE_CHECKING:
from collections.abc import Iterator
from openai import AsyncOpenAI
logger = logging.getLogger(__name__)
PROVIDER = "grok"
CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828"
AUTHORIZE_URL = "https://auth.x.ai/oauth2/authorize"
TOKEN_URL = "https://auth.x.ai/oauth2/token" # noqa: S105 # nosec B105 - URL, not a secret
CALLBACK_HOST = "127.0.0.1"
CALLBACK_PORT = 56121
CALLBACK_PATH = "/callback"
REDIRECT_URI = f"http://{CALLBACK_HOST}:{CALLBACK_PORT}{CALLBACK_PATH}"
SCOPE = "openid profile email offline_access grok-cli:access api:access"
XAI_BASE_URL = "https://api.x.ai/v1"
_TOKEN_TIMEOUT = 30
_EXPIRY_SKEW_S = 300
# Shared with the other subscription providers; kept separate from cli-config.json
# so OAuth tokens never land in the env-var config.
AUTH_PATH = Path.home() / ".strix" / "subscription-auth.json"
def read_record() -> dict[str, Any] | None:
record = subscription_store.read(AUTH_PATH).get(PROVIDER)
if not isinstance(record, dict) or record.get("type") != "oauth":
return None
if not (record.get("access") and record.get("refresh")):
return None
return record
def is_authenticated() -> bool:
return read_record() is not None
def save_record(record: dict[str, Any]) -> None:
with subscription_store.guard(AUTH_PATH):
data = subscription_store.read(AUTH_PATH)
data[PROVIDER] = record
subscription_store.write(AUTH_PATH, data)
def logout() -> None:
with subscription_store.guard(AUTH_PATH):
data = subscription_store.read(AUTH_PATH)
if PROVIDER not in data:
return
del data[PROVIDER]
if data:
subscription_store.write(AUTH_PATH, data)
return
with contextlib.suppress(OSError):
AUTH_PATH.unlink()
@contextlib.contextmanager
def _refresh_guard() -> Iterator[None]:
"""Serialize token refresh within (lock) and across (flock) Strix processes,
so concurrent runs can't both spend the single-use refresh token."""
with subscription_store.guard(AUTH_PATH):
yield
class GrokAuthError(Exception):
def __init__(self, code: str, message: str | None = None) -> None:
self.code = code
super().__init__(message or code)
def _b64url(raw: bytes) -> str:
return base64.urlsafe_b64encode(raw).rstrip(b"=").decode("ascii")
def generate_pkce() -> tuple[str, str]:
verifier = _b64url(secrets.token_bytes(64))
challenge = _b64url(hashlib.sha256(verifier.encode("ascii")).digest())
return verifier, challenge
def create_state() -> str:
return secrets.token_hex(16)
def build_authorize_url(challenge: str, state: str) -> str:
params = {
"response_type": "code",
"client_id": CLIENT_ID,
"redirect_uri": REDIRECT_URI,
"scope": SCOPE,
"code_challenge": challenge,
"code_challenge_method": "S256",
"state": state,
}
return f"{AUTHORIZE_URL}?{urllib.parse.urlencode(params)}"
def parse_redirect_input(value: str) -> tuple[str | None, str | None]:
"""Extract ``(code, state)`` from a pasted redirect URL, ``code#state``,
query string, or bare code."""
value = (value or "").strip()
if not value:
return None, None
with contextlib.suppress(ValueError):
parsed = urllib.parse.urlparse(value)
if parsed.scheme and parsed.query:
query = urllib.parse.parse_qs(parsed.query)
return _first(query, "code"), _first(query, "state")
if "#" in value:
code, _, state = value.partition("#")
return code or None, state or None
if "code=" in value:
query = urllib.parse.parse_qs(value)
return _first(query, "code"), _first(query, "state")
return value, None
def _first(query: dict[str, list[str]], key: str) -> str | None:
values = query.get(key)
return values[0] if values else None
def _post_form(payload: dict[str, str]) -> dict[str, Any]:
detail = ""
try:
with requests.post(
TOKEN_URL,
data=payload,
headers={"Accept": "application/json"},
timeout=_TOKEN_TIMEOUT,
) as response:
status_code = response.status_code
body = response.content
if status_code >= 400:
detail = response.text[:300]
except requests.RequestException as exc:
raise GrokAuthError("unavailable", str(exc)) from exc
if status_code >= 400:
raise GrokAuthError("token_http_error", f"HTTP {status_code}: {detail}")
data = json.loads(body or b"{}")
if not isinstance(data, dict):
raise GrokAuthError("bad_response", "token endpoint returned non-object")
return data
def _record_from_token_response(
data: dict[str, Any], refresh_fallback: str | None = None
) -> dict[str, Any]:
access = data.get("access_token")
# A refresh response may omit refresh_token when it isn't rotated; keep the old one.
refresh = data.get("refresh_token") or refresh_fallback
expires_in = data.get("expires_in")
if not isinstance(access, str) or not access:
raise GrokAuthError("bad_response", "token response missing access_token")
if not isinstance(refresh, str) or not refresh:
raise GrokAuthError("bad_response", "token response missing refresh_token")
ttl = expires_in if isinstance(expires_in, int | float) else 3600
return {
"type": "oauth",
"provider": PROVIDER,
"access": access,
"refresh": refresh,
"expires_at": time.time() + ttl,
}
def exchange_code(code: str, verifier: str) -> dict[str, Any]:
data = _post_form(
{
"grant_type": "authorization_code",
"client_id": CLIENT_ID,
"code": code,
"code_verifier": verifier,
"redirect_uri": REDIRECT_URI,
}
)
return _record_from_token_response(data)
def refresh_tokens(refresh_token: str) -> dict[str, Any]:
data = _post_form(
{
"grant_type": "refresh_token",
"client_id": CLIENT_ID,
"refresh_token": refresh_token,
}
)
return _record_from_token_response(data, refresh_fallback=refresh_token)
def _access_token(record: dict[str, Any]) -> str:
access = record["access"]
if not isinstance(access, str) or not access:
raise GrokAuthError("bad_response", "stored access token is missing or malformed")
return access
def _near_expiry(record: dict[str, Any]) -> bool:
expires_at = record.get("expires_at")
if not isinstance(expires_at, int | float):
return True
return expires_at - _EXPIRY_SKEW_S <= time.time()
def get_valid_token() -> str:
"""Return a valid access token, refreshing under the cross-process guard if
near expiry."""
record = read_record()
if record is None:
raise GrokAuthError("not_authenticated", "not signed in; run: strix auth login grok")
if not _near_expiry(record):
return _access_token(record)
with _refresh_guard():
record = read_record()
if record is None:
raise GrokAuthError("not_authenticated", "not signed in; run: strix auth login grok")
if not _near_expiry(record):
return _access_token(record)
try:
refreshed = refresh_tokens(record["refresh"])
except GrokAuthError:
# A peer process may have already spent this single-use refresh token.
latest = read_record()
if latest and latest["refresh"] != record["refresh"] and not _near_expiry(latest):
return _access_token(latest)
raise
save_record(refreshed)
return _access_token(refreshed)
def build_openai_client() -> AsyncOpenAI:
"""An ``AsyncOpenAI`` for xAI's API. A per-request hook re-stamps a fresh
bearer token so long scans survive token expiry."""
import asyncio
import httpx
from openai import AsyncOpenAI
get_valid_token() # fail fast at configure time if the sign-in is dead
async def _auth_hook(request: httpx.Request) -> None:
access = await asyncio.to_thread(get_valid_token)
request.headers["Authorization"] = f"Bearer {access}"
http_client = httpx.AsyncClient(
timeout=httpx.Timeout(600.0, connect=30.0),
event_hooks={"request": [_auth_hook]},
)
return AsyncOpenAI(
api_key="strix-grok-oauth", # placeholder; the hook overwrites Authorization
base_url=XAI_BASE_URL,
http_client=http_client,
)
_subscription_client: AsyncOpenAI | None = None
def get_subscription_client() -> AsyncOpenAI:
global _subscription_client # noqa: PLW0603
if _subscription_client is None:
_subscription_client = build_openai_client()
return _subscription_client
SUBSCRIPTION_PREFIX = "grok/"
def subscription_model(model_name: str | None) -> str | None:
"""The model slug behind a ``grok/<model>`` STRIX_LLM, or None."""
name = (model_name or "").strip()
if not name.lower().startswith(SUBSCRIPTION_PREFIX):
return None
return name[len(SUBSCRIPTION_PREFIX) :] or None
def auth_mode(model_name: str | None) -> str:
return "subscription" if subscription_model(model_name) else "api_key"
+16 -7
View File
@@ -20,6 +20,7 @@ from agents.model_settings import ModelSettings
from agents.models.fake_id import FAKE_RESPONSES_ID from agents.models.fake_id import FAKE_RESPONSES_ID
from agents.models.interface import Model from agents.models.interface import Model
from agents.models.multi_provider import MultiProvider from agents.models.multi_provider import MultiProvider
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from agents.models.openai_responses import OpenAIResponsesModel from agents.models.openai_responses import OpenAIResponsesModel
from agents.retry import ( from agents.retry import (
ModelRetryBackoffSettings, ModelRetryBackoffSettings,
@@ -36,7 +37,7 @@ from openai.types.responses import (
from openai.types.responses.response_usage import ResponseUsage from openai.types.responses.response_usage import ResponseUsage
from openai.types.shared import Reasoning from openai.types.shared import Reasoning
from strix.config import codex from strix.config import codex, grok
from strix.config.loader import load_settings from strix.config.loader import load_settings
from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input
from strix.config.tool_call_limits import TurnToolCallLimiter from strix.config.tool_call_limits import TurnToolCallLimiter
@@ -481,6 +482,10 @@ class StrixProvider(MultiProvider):
codex.get_subscription_client(), codex.get_subscription_client(),
reasoning_effort=llm.reasoning_effort, reasoning_effort=llm.reasoning_effort,
) )
elif grok_slug := grok.subscription_model(model_name):
# xAI's API is OpenAI chat-completions compatible; the subscription
# bearer is stamped per-request by the client's auth hook.
model = OpenAIChatCompletionsModel(grok_slug, grok.get_subscription_client())
else: else:
model = super().get_model(model_name) model = super().get_model(model_name)
if llm.disable_streaming: if llm.disable_streaming:
@@ -556,7 +561,7 @@ def configure_sdk_model_defaults(settings: Settings) -> None:
"""Apply Strix config to SDK-native defaults.""" """Apply Strix config to SDK-native defaults."""
llm = settings.llm llm = settings.llm
set_tracing_disabled(True) set_tracing_disabled(True)
if codex.subscription_model(llm.model): if codex.subscription_model(llm.model) or grok.subscription_model(llm.model):
return return
_configure_litellm_compatibility() _configure_litellm_compatibility()
_configure_openrouter_attribution(llm.model) _configure_openrouter_attribution(llm.model)
@@ -652,27 +657,31 @@ def _install_openrouter_stream_cost_capture() -> None:
litellm.OpenrouterConfig = _StrixOpenrouterConfig # type: ignore[misc] litellm.OpenrouterConfig = _StrixOpenrouterConfig # type: ignore[misc]
_OPENROUTER_ATTRIBUTION_HEADERS = { OPENROUTER_ATTRIBUTION_HEADERS = {
"HTTP-Referer": "https://strix.ai", "HTTP-Referer": "https://strix.ai",
"X-Title": "Strix", "X-Title": "Strix",
"X-OpenRouter-Categories": "cli-agent", "X-OpenRouter-Categories": "cli-agent",
} }
def is_openrouter_model(model_name: str | None) -> bool:
return bool(model_name) and "openrouter/" in (model_name or "").strip().lower()
def _configure_openrouter_attribution(model_name: str | None) -> None: def _configure_openrouter_attribution(model_name: str | None) -> None:
import litellm import litellm
current: object = litellm.headers current: object = litellm.headers
existing: dict[str, str] = current if isinstance(current, dict) else {} existing: dict[str, str] = current if isinstance(current, dict) else {}
if not model_name or "openrouter/" not in model_name.strip().lower(): if not is_openrouter_model(model_name):
if any(key in existing for key in _OPENROUTER_ATTRIBUTION_HEADERS): if any(key in existing for key in OPENROUTER_ATTRIBUTION_HEADERS):
remaining = { remaining = {
k: v for k, v in existing.items() if k not in _OPENROUTER_ATTRIBUTION_HEADERS k: v for k, v in existing.items() if k not in OPENROUTER_ATTRIBUTION_HEADERS
} }
litellm.headers = remaining or None # type: ignore[assignment] litellm.headers = remaining or None # type: ignore[assignment]
return return
litellm.headers = {**existing, **_OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment] litellm.headers = {**existing, **OPENROUTER_ATTRIBUTION_HEADERS} # type: ignore[assignment]
def _configure_extra_headers(llm: LlmSettings) -> None: def _configure_extra_headers(llm: LlmSettings) -> None:
+63
View File
@@ -0,0 +1,63 @@
"""Shared helpers across model-subscription providers (ChatGPT/Codex and Grok).
Each provider module (:mod:`strix.config.codex`, :mod:`strix.config.grok`)
exposes the same small surface — ``subscription_model``, ``auth_mode``,
``is_authenticated`` — so callers that only care "is this run on a subscription,
and which provider?" can stay provider-agnostic.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
from strix.config import codex, grok
if TYPE_CHECKING:
from types import ModuleType
_PROVIDERS: tuple[ModuleType, ...] = (codex, grok)
# Human-facing provider names keyed by each module's ``PROVIDER`` constant.
_DISPLAY_NAMES: dict[str, str] = {codex.PROVIDER: "ChatGPT", grok.PROVIDER: "Grok"}
# Prefix LiteLLM keys each provider's model metadata under: ChatGPT models are
# mapped bare ("gpt-5.4"), xAI's only provider-qualified ("xai/grok-4").
_LITELLM_PREFIXES: dict[str, str] = {codex.PROVIDER: "", grok.PROVIDER: "xai/"}
def provider_for_model(model_name: str | None) -> ModuleType | None:
"""Return the subscription provider module that owns ``model_name``'s prefix,
or None when the model isn't a subscription model."""
for provider in _PROVIDERS:
if provider.subscription_model(model_name):
return provider
return None
def auth_mode(model_name: str | None) -> str:
return "subscription" if provider_for_model(model_name) is not None else "api_key"
def provider_label(model_name: str | None) -> str | None:
"""Human-facing name of the subscription provider for ``model_name`` (e.g.
"ChatGPT" or "Grok"), or None when the model isn't a subscription model."""
provider = provider_for_model(model_name)
if provider is None:
return None
return _DISPLAY_NAMES.get(provider.PROVIDER)
def litellm_model_name(model_name: str | None) -> str | None:
"""``model_name`` rewritten to the name LiteLLM maps metadata under.
Subscription prefixes are Strix routing labels LiteLLM never maps, so a
lookup of "grok/grok-4" (or bare "grok-4") finds nothing. Non-subscription
models are returned unchanged.
"""
provider = provider_for_model(model_name)
if provider is None:
return model_name
prefix = _LITELLM_PREFIXES.get(provider.PROVIDER, "")
return f"{prefix}{provider.subscription_model(model_name)}"
+152
View File
@@ -0,0 +1,152 @@
"""Shared on-disk store for subscription OAuth credentials.
Every subscription provider (ChatGPT/Codex, Grok) keeps its record under its own
key in a single ``~/.strix/subscription-auth.json`` file. Reads and writes go
through here so that:
* tokens are written owner-only (mode 0600) from the moment the file is created,
never briefly exposed with umask-derived permissions, and
* concurrent read-modify-write mutations — even across different providers or
processes — are serialized, so one provider's update can't clobber another's.
The lock is reentrant, so a provider may nest a ``save`` inside a longer
``guard`` (e.g. refreshing a token then persisting it) without deadlocking.
"""
from __future__ import annotations
import contextlib
import json
import os
import tempfile
import threading
from pathlib import Path
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from collections.abc import Iterator
from io import TextIOWrapper
class StoreLockError(RuntimeError):
"""The cross-process store lock could not be acquired.
Raised instead of silently proceeding, so a read-modify-write never runs
unlocked (which would let concurrent provider logins/refreshes/logouts race).
"""
def read(path: Path) -> dict[str, Any]:
"""The store's contents, or an empty dict when absent/unreadable."""
try:
data = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return {}
return data if isinstance(data, dict) else {}
def write(path: Path, data: dict[str, Any]) -> None:
"""Atomically replace the store, owner-only from creation.
The temp file is created with a random name via ``mkstemp`` (mode 0600, no
symlink following), so a local attacker can't pre-plant a symlink at a
predictable path to divert the token write.
"""
path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp_name = tempfile.mkstemp(dir=path.parent, prefix=path.name, suffix=".tmp")
tmp = Path(tmp_name)
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
json.dump(data, handle, indent=2)
tmp.replace(path)
except BaseException:
with contextlib.suppress(OSError):
tmp.unlink()
raise
with contextlib.suppress(OSError):
path.chmod(0o600)
class _StoreLock:
"""A reentrant lock serializing store mutations within (thread lock) and
across (flock) Strix processes. Nesting reuses the single held file lock, so
a provider can persist a record inside a longer critical section."""
def __init__(self) -> None:
self._thread_lock = threading.RLock()
self._flock_handle: TextIOWrapper | None = None
self._depth = 0
@contextlib.contextmanager
def hold(self, path: Path) -> Iterator[None]:
with self._thread_lock:
if self._depth == 0:
self._flock_handle = _acquire_flock(path)
self._depth += 1
try:
yield
finally:
self._depth -= 1
if self._depth == 0:
self._release_flock()
def _release_flock(self) -> None:
handle = self._flock_handle
self._flock_handle = None
if handle is None:
return
try:
import fcntl
with contextlib.suppress(OSError):
fcntl.flock(handle.fileno(), fcntl.LOCK_UN)
except ImportError:
pass
finally:
handle.close()
_store_lock = _StoreLock()
def guard(path: Path) -> contextlib.AbstractContextManager[None]:
"""Serialize store mutation across threads and processes (reentrant)."""
return _store_lock.hold(path)
def _acquire_flock(path: Path) -> TextIOWrapper:
"""Hold an exclusive cross-process lock on the store, or raise.
Never returns without the lock held: a missing ``fcntl`` or a failed
``flock`` raises :class:`StoreLockError` so the caller aborts rather than
mutating the store unlocked.
"""
try:
import fcntl
except ImportError as exc: # pragma: no cover - non-POSIX
msg = "cross-process credential locking requires fcntl (a POSIX platform)"
raise StoreLockError(msg) from exc
lock_path = path.with_suffix(".lock")
lock_path.parent.mkdir(parents=True, exist_ok=True)
# O_NOFOLLOW rejects a pre-positioned symlink at the predictable lock path
# (so an attacker can't redirect the open), and no O_TRUNC since the lock
# file is only an flock anchor whose contents we never use.
try:
fd = os.open(str(lock_path), os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW, 0o600)
except OSError as exc:
msg = f"could not open lock file {lock_path}: {exc}"
raise StoreLockError(msg) from exc
handle = os.fdopen(fd, "r+")
try:
while True:
try:
fcntl.flock(handle.fileno(), fcntl.LOCK_EX)
break
except InterruptedError: # EINTR — retry the blocking acquire
continue
except OSError as exc:
handle.close()
msg = f"could not lock {lock_path}: {exc}"
raise StoreLockError(msg) from exc
return handle
+17 -2
View File
@@ -10,10 +10,12 @@ from openai.types.shared import Reasoning
from strix.config.models import ( from strix.config.models import (
DEFAULT_MODEL_RETRY, DEFAULT_MODEL_RETRY,
OPENROUTER_ATTRIBUTION_HEADERS,
bedrock_route_supports_prompt_caching, bedrock_route_supports_prompt_caching,
is_bedrock_route, is_bedrock_route,
is_claude_model, is_claude_model,
is_known_openai_bare_model, is_known_openai_bare_model,
is_openrouter_model,
model_supports_reasoning, model_supports_reasoning,
request_timeout_extra_args, request_timeout_extra_args,
) )
@@ -201,13 +203,15 @@ def make_model_settings(
request_timeout: float | None = None, request_timeout: float | None = None,
prompt_cache: bool = True, prompt_cache: bool = True,
extra_headers: dict[str, str] | None = None, extra_headers: dict[str, str] | None = None,
has_tools: bool = True,
) -> ModelSettings: ) -> ModelSettings:
headers = _request_headers(model_name, extra_headers)
model_settings = ModelSettings( model_settings = ModelSettings(
parallel_tool_calls=False, parallel_tool_calls=False if has_tools else None,
retry=DEFAULT_MODEL_RETRY, retry=DEFAULT_MODEL_RETRY,
include_usage=True, include_usage=True,
extra_args=request_timeout_extra_args(request_timeout), extra_args=request_timeout_extra_args(request_timeout),
extra_headers=dict(extra_headers) if extra_headers else None, extra_headers=headers,
) )
if ( if (
reasoning_effort is not None reasoning_effort is not None
@@ -230,6 +234,17 @@ def make_model_settings(
return model_settings return model_settings
def _request_headers(
model_name: str, extra_headers: dict[str, str] | None
) -> dict[str, str] | None:
headers: dict[str, str] = {}
if is_openrouter_model(model_name):
headers.update(OPENROUTER_ATTRIBUTION_HEADERS)
if extra_headers:
headers.update(extra_headers)
return headers or None
def _reasoning_settings( def _reasoning_settings(
effort: ReasoningEffort, effort: ReasoningEffort,
extra_args: dict[str, Any] | None, extra_args: dict[str, Any] | None,
+166 -67
View File
@@ -1,8 +1,8 @@
"""`strix auth` — ChatGPT subscription sign-in (login / status / logout). """`strix auth` — model-subscription sign-in (login / status / logout).
Signing in only stores OAuth tokens (``~/.strix/subscription-auth.json``); model Signing in only stores OAuth tokens (``~/.strix/subscription-auth.json``); model
selection stays with ``STRIX_LLM``. A ``chatgpt/<model>`` STRIX_LLM runs on the selection stays with ``STRIX_LLM``. A ``chatgpt/<model>`` STRIX_LLM runs on a
subscription. ChatGPT subscription and a ``grok/<model>`` one on a Grok/SuperGrok subscription.
""" """
from __future__ import annotations from __future__ import annotations
@@ -12,6 +12,7 @@ import base64
import logging import logging
import threading import threading
import webbrowser import webbrowser
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, HTTPServer from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
@@ -21,24 +22,76 @@ from rich.console import Console
from rich.panel import Panel from rich.panel import Panel
from rich.text import Text from rich.text import Text
from strix.config import codex, load_settings from strix.config import codex, grok, load_settings, subscription_store
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable from collections.abc import Callable
from types import ModuleType
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_CALLBACK_TIMEOUT_S = 300 _CALLBACK_TIMEOUT_S = 300
# CLI-facing name for the login provider. Internally this is the Codex OAuth
# flow (``codex.PROVIDER``), but users know it as ChatGPT, so that's what the
# command and messaging say. ``codex`` is accepted as an alias.
LOGIN_PROVIDER = "chatgpt"
_ACCEPTED_PROVIDERS = frozenset({LOGIN_PROVIDER, codex.PROVIDER})
_USAGE = "Usage:\n strix auth login chatgpt [--manual]\n strix auth status\n strix auth logout" @dataclass(frozen=True)
class _Provider:
"""A model-subscription provider the ``strix auth`` command can sign into.
``module`` is the provider's OAuth module (:mod:`strix.config.codex` or
:mod:`strix.config.grok`); both expose the same login surface. ``error`` is
that module's auth-error class, caught to report a clean failure.
"""
name: str
module: ModuleType
error: type[Exception]
display: str
example_model: str
blurb: str
_PROVIDERS: dict[str, _Provider] = {
"chatgpt": _Provider(
name="chatgpt",
module=codex,
error=codex.CodexAuthError,
display="ChatGPT",
example_model="chatgpt/gpt-5.4",
blurb="This uses your ChatGPT Plus/Pro plan for inference instead of a metered API key.",
),
"grok": _Provider(
name="grok",
module=grok,
error=grok.GrokAuthError,
display="Grok",
example_model="grok/grok-4",
blurb="This uses your Grok/SuperGrok plan for inference instead of a metered API key.",
),
}
# Internal OAuth provider ids and common vendor names accepted as aliases.
_PROVIDER_ALIASES: dict[str, str] = {
codex.PROVIDER: "chatgpt",
grok.PROVIDER: "grok",
"xai": "grok",
"supergrok": "grok",
}
_DEFAULT_PROVIDER = "chatgpt"
_USAGE = (
"Usage:\n"
" strix auth login [chatgpt|grok] [--manual]\n"
" strix auth status\n"
" strix auth logout [chatgpt|grok]"
)
def _resolve_provider(name: str) -> _Provider | None:
key = _PROVIDER_ALIASES.get(name.lower(), name.lower())
return _PROVIDERS.get(key)
def run_auth(argv: list[str]) -> int: def run_auth(argv: list[str]) -> int:
@@ -49,20 +102,20 @@ def run_auth(argv: list[str]) -> int:
rest = argv[1:] rest = argv[1:]
if subcommand in ("-h", "--help", "help"): if subcommand in ("-h", "--help", "help"):
console.print(_USAGE) console.print(_USAGE, markup=False)
return 0 return 0
handlers: dict[str, Callable[[], int]] = { handlers: dict[str, Callable[[], int]] = {
"login": lambda: _login(console, rest), "login": lambda: _login(console, rest),
"status": lambda: _status(console), "status": lambda: _status(console),
"logout": lambda: _logout(console), "logout": lambda: _logout(console, rest),
} }
handler = handlers.get(subcommand) handler = handlers.get(subcommand)
if handler is not None: if handler is not None:
return handler() return handler()
console.print(f"[red]Unknown auth command:[/] {subcommand}\n") console.print(f"[red]Unknown auth command:[/] {subcommand}\n")
console.print(_USAGE) console.print(_USAGE, markup=False)
return 2 return 2
@@ -71,8 +124,8 @@ def _login(console: Console, argv: list[str]) -> int:
parser.add_argument( parser.add_argument(
"provider", "provider",
nargs="?", nargs="?",
default=LOGIN_PROVIDER, default=_DEFAULT_PROVIDER,
help="Model provider to sign in with (default: chatgpt).", help="Model provider to sign in with (chatgpt or grok; default: chatgpt).",
) )
parser.add_argument( parser.add_argument(
"--manual", "--manual",
@@ -84,39 +137,42 @@ def _login(console: Console, argv: list[str]) -> int:
except SystemExit as exc: # argparse already printed the message except SystemExit as exc: # argparse already printed the message
return int(exc.code or 2) return int(exc.code or 2)
if args.provider.lower() not in _ACCEPTED_PROVIDERS: provider = _resolve_provider(args.provider)
console.print( if provider is None:
f"[red]Unsupported provider:[/] {args.provider}. " supported = ", ".join(f"'{name}'" for name in _PROVIDERS)
f"Only '{LOGIN_PROVIDER}' (ChatGPT subscription) is supported." console.print(f"[red]Unsupported provider:[/] {args.provider}. Supported: {supported}.")
)
return 2 return 2
verifier, challenge = codex.generate_pkce() module = provider.module
state = codex.create_state() verifier, challenge = module.generate_pkce()
authorize_url = codex.build_authorize_url(challenge, state) state = module.create_state()
authorize_url = module.build_authorize_url(challenge, state)
console.print() console.print()
console.print("[bold]Signing in with ChatGPT[/] [dim](provider: chatgpt)[/]")
console.print( console.print(
"[dim]This uses your ChatGPT Plus/Pro plan for inference instead of a metered API key.[/]" f"[bold]Signing in with {provider.display}[/] [dim](provider: {provider.name})[/]"
) )
console.print(f"[dim]{provider.blurb}[/]")
console.print() console.print()
try: try:
record = _run_oauth_flow(console, authorize_url, verifier, state, manual=args.manual) record = _run_oauth_flow(
except codex.CodexAuthError as exc: console, provider, authorize_url, verifier, state, manual=args.manual
)
except provider.error as exc:
return _fail(console, exc) return _fail(console, exc)
except KeyboardInterrupt: except KeyboardInterrupt:
console.print("\n[yellow]Sign-in cancelled.[/]") console.print("\n[yellow]Sign-in cancelled.[/]")
return 130 return 130
codex.save_record(record) module.save_record(record)
_print_success(console) _print_success(console, provider)
return 0 return 0
def _run_oauth_flow( def _run_oauth_flow(
console: Console, console: Console,
provider: _Provider,
authorize_url: str, authorize_url: str,
verifier: str, verifier: str,
state: str, state: str,
@@ -124,7 +180,10 @@ def _run_oauth_flow(
manual: bool, manual: bool,
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Drive the browser (or manual) OAuth flow and return a token record.""" """Drive the browser (or manual) OAuth flow and return a token record."""
server = None if manual else _try_start_callback_server() module = provider.module
server = (
None if manual else _try_start_callback_server(module.CALLBACK_PORT, module.CALLBACK_PATH)
)
console.print("Open this URL in your browser to authorize:") console.print("Open this URL in your browser to authorize:")
console.print(f"[cyan]{authorize_url}[/]") console.print(f"[cyan]{authorize_url}[/]")
@@ -142,8 +201,8 @@ def _run_oauth_flow(
if result is not None: if result is not None:
code, returned_state, error = result code, returned_state, error = result
if error: if error:
raise codex.CodexAuthError("oauth_error", error) raise provider.error("oauth_error", error)
return _finish(code, returned_state, verifier, state, require_state=True) return _finish(provider, code, returned_state, verifier, state, require_state=True)
console.print("[yellow]Timed out waiting for the browser. Falling back to manual paste.[/]") console.print("[yellow]Timed out waiting for the browser. Falling back to manual paste.[/]")
# Manual fallback: the user completes sign-in and pastes the redirect URL # Manual fallback: the user completes sign-in and pastes the redirect URL
@@ -153,12 +212,13 @@ def _run_oauth_flow(
try: try:
pasted = console.input("Paste the full redirect URL (or code#state): ").strip() pasted = console.input("Paste the full redirect URL (or code#state): ").strip()
except EOFError as exc: except EOFError as exc:
raise codex.CodexAuthError("no_input", "no redirect URL provided") from exc raise provider.error("no_input", "no redirect URL provided") from exc
code, returned_state = codex.parse_redirect_input(pasted) code, returned_state = module.parse_redirect_input(pasted)
return _finish(code, returned_state, verifier, state, require_state=False) return _finish(provider, code, returned_state, verifier, state, require_state=False)
def _finish( def _finish(
provider: _Provider,
code: str | None, code: str | None,
returned_state: str | None, returned_state: str | None,
verifier: str, verifier: str,
@@ -167,16 +227,17 @@ def _finish(
require_state: bool, require_state: bool,
) -> dict[str, Any]: ) -> dict[str, Any]:
if not code: if not code:
raise codex.CodexAuthError("no_code", "no authorization code found in the redirect") raise provider.error("no_code", "no authorization code found in the redirect")
# The loopback callback from OpenAI always carries state, so a missing or # The loopback callback from the provider always carries state, so a missing
# mismatched value there is forged (CSRF) and must be rejected. Manual paste # or mismatched value there is forged (CSRF) and must be rejected. Manual
# is user-initiated (the user copies their own redirect), so state is only # paste is user-initiated (the user copies their own redirect), so state is
# validated when the pasted value includes it. # only validated when the pasted value includes it.
if require_state and returned_state is None: if require_state and returned_state is None:
raise codex.CodexAuthError("state_mismatch", "missing state in callback; possible CSRF") raise provider.error("state_mismatch", "missing state in callback; possible CSRF")
if returned_state is not None and returned_state != expected_state: if returned_state is not None and returned_state != expected_state:
raise codex.CodexAuthError("state_mismatch", "state did not match; possible CSRF") raise provider.error("state_mismatch", "state did not match; possible CSRF")
return codex.exchange_code(code, verifier) record: dict[str, Any] = provider.module.exchange_code(code, verifier)
return record
class _CallbackServer: class _CallbackServer:
@@ -203,7 +264,7 @@ class _CallbackServer:
self._httpd.server_close() self._httpd.server_close()
def _try_start_callback_server() -> _CallbackServer | None: def _try_start_callback_server(port: int, path: str) -> _CallbackServer | None:
event = threading.Event() event = threading.Event()
holder: dict[str, Any] = {} holder: dict[str, Any] = {}
@@ -213,7 +274,7 @@ def _try_start_callback_server() -> _CallbackServer | None:
def do_GET(self) -> None: def do_GET(self) -> None:
parsed = urlparse(self.path) parsed = urlparse(self.path)
if parsed.path != codex.CALLBACK_PATH: if parsed.path != path:
self.send_response(404) self.send_response(404)
self.end_headers() self.end_headers()
return return
@@ -230,9 +291,9 @@ def _try_start_callback_server() -> _CallbackServer | None:
event.set() event.set()
try: try:
httpd = HTTPServer(("127.0.0.1", codex.CALLBACK_PORT), Handler) httpd = HTTPServer(("127.0.0.1", port), Handler)
except OSError: except OSError:
logger.debug("could not bind callback port %d", codex.CALLBACK_PORT, exc_info=True) logger.debug("could not bind callback port %d", port, exc_info=True)
return None return None
return _CallbackServer(httpd, event, holder) return _CallbackServer(httpd, event, holder)
@@ -243,30 +304,67 @@ def _first(query: dict[str, list[str]], key: str) -> str | None:
def _status(console: Console) -> int: def _status(console: Console) -> int:
record = codex.read_record()
if record is None:
console.print("[yellow]Not signed in.[/] Run [cyan]strix auth login chatgpt[/] to sign in.")
return 1
settings = load_settings() settings = load_settings()
console.print("[green]Signed in[/] with a ChatGPT subscription.") active_model = settings.llm.model
console.print(f" Account: [bold]{record.get('account_id')}[/]") signed_in_any = False
if codex.subscription_model(settings.llm.model): for provider in _PROVIDERS.values():
console.print(f" Runs use the subscription (STRIX_LLM=[bold]{settings.llm.model}[/]).") record = provider.module.read_record()
else: if record is None:
continue
signed_in_any = True
console.print(f"[green]Signed in[/] with a {provider.display} subscription.")
account_id = record.get("account_id")
if account_id:
console.print(f" Account: [bold]{account_id}[/]")
if provider.module.subscription_model(active_model):
console.print(f" Runs use the subscription (STRIX_LLM=[bold]{active_model}[/]).")
else:
console.print(
f" [yellow]Note:[/] set [cyan]STRIX_LLM[/] to e.g. "
f"[cyan]{provider.example_model}[/] to run on this subscription."
)
if not signed_in_any:
console.print( console.print(
" [yellow]Note:[/] set [cyan]STRIX_LLM[/] to e.g. [cyan]chatgpt/gpt-5.4[/] " "[yellow]Not signed in.[/] Run [cyan]strix auth login chatgpt[/] "
"to run on the subscription." "or [cyan]strix auth login grok[/] to sign in."
) )
return 1
return 0 return 0
def _logout(console: Console) -> int: def _logout(console: Console, argv: list[str]) -> int:
codex.logout() parser = argparse.ArgumentParser(prog="strix auth logout", add_help=True)
console.print("[green]Signed out.[/] Stored subscription credentials removed.") parser.add_argument(
"provider",
nargs="?",
default=None,
help="Provider to sign out of (chatgpt or grok; default: all).",
)
try:
args = parser.parse_args(argv)
except SystemExit as exc:
return int(exc.code or 2)
if args.provider is None:
# Hold the store lock across every provider so a concurrent save/refresh
# can't slip a credential back in between removals (logout-all is atomic).
with subscription_store.guard(codex.AUTH_PATH):
for provider in _PROVIDERS.values():
provider.module.logout()
console.print("[green]Signed out.[/] Stored subscription credentials removed.")
return 0
target = _resolve_provider(args.provider)
if target is None:
supported = ", ".join(f"'{name}'" for name in _PROVIDERS)
console.print(f"[red]Unsupported provider:[/] {args.provider}. Supported: {supported}.")
return 2
target.module.logout()
console.print(f"[green]Signed out of {target.display}.[/] Stored credentials removed.")
return 0 return 0
def _fail(console: Console, exc: codex.CodexAuthError) -> int: def _fail(console: Console, exc: Exception) -> int:
error_text = Text() error_text = Text()
error_text.append("SIGN-IN FAILED", style="bold red") error_text.append("SIGN-IN FAILED", style="bold red")
error_text.append("\n\n", style="white") error_text.append("\n\n", style="white")
@@ -284,17 +382,18 @@ def _fail(console: Console, exc: codex.CodexAuthError) -> int:
return 1 return 1
def _print_success(console: Console) -> None: def _print_success(console: Console, provider: _Provider) -> None:
prefix = provider.module.SUBSCRIPTION_PREFIX
text = Text() text = Text()
text.append("Signed in with your ChatGPT subscription", style="bold #22c55e") text.append(f"Signed in with your {provider.display} subscription", style="bold #22c55e")
text.append("\n\n", style="white") text.append("\n\n", style="white")
text.append("Set ", style="white") text.append("Set ", style="white")
text.append("STRIX_LLM", style="bold white") text.append("STRIX_LLM", style="bold white")
text.append(" to a ", style="white") text.append(" to a ", style="white")
text.append("chatgpt/", style="bold cyan") text.append(prefix, style="bold cyan")
text.append(" model (e.g. ", style="white") text.append(" model (e.g. ", style="white")
text.append("chatgpt/gpt-5.4", style="bold cyan") text.append(provider.example_model, style="bold cyan")
text.append(") — runs are billed to your ChatGPT plan.", style="white") text.append(f") — runs are billed to your {provider.display} plan.", style="white")
text.append("\n\n", style="white") text.append("\n\n", style="white")
text.append("Run a scan as usual, e.g. ", style="white") text.append("Run a scan as usual, e.g. ", style="white")
text.append("strix --target https://example.com", style="bold cyan") text.append("strix --target https://example.com", style="bold cyan")
+11 -1
View File
@@ -8,7 +8,7 @@ from rich.console import Console
from rich.panel import Panel from rich.panel import Panel
from rich.text import Text from rich.text import Text
from strix.config import codex, load_settings from strix.config import codex, grok, load_settings
from strix.interface.utils import ( from strix.interface.utils import (
check_docker_connection, check_docker_connection,
image_exists, image_exists,
@@ -37,6 +37,16 @@ def validate_environment() -> None:
logger.info("Environment OK (ChatGPT subscription)") logger.info("Environment OK (ChatGPT subscription)")
return return
if grok.subscription_model(settings.llm.model):
if not grok.is_authenticated():
console.print(
f"[red]STRIX_LLM={settings.llm.model} uses your Grok subscription, "
"but you're not signed in.[/] Run [cyan]strix auth login grok[/] first."
)
sys.exit(1)
logger.info("Environment OK (Grok subscription)")
return
if not settings.llm.model: if not settings.llm.model:
missing_required_vars.append("STRIX_LLM") missing_required_vars.append("STRIX_LLM")
+11 -1
View File
@@ -224,6 +224,7 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
request_timeout=llm.timeout, request_timeout=llm.timeout,
prompt_cache=False, prompt_cache=False,
extra_headers=settings.dedupe.extra_headers, extra_headers=settings.dedupe.extra_headers,
has_tools=False,
) )
if deduper_extra: if deduper_extra:
merged = {**(deduper_settings.extra_args or {}), **deduper_extra} merged = {**(deduper_settings.extra_args or {}), **deduper_extra}
@@ -435,7 +436,16 @@ def main() -> None:
start_background_check() start_background_check()
if not args.non_interactive and prompt_update_if_available(Console()): if not args.non_interactive and prompt_update_if_available(Console()):
if is_binary_install() and sys.platform != "win32": if is_binary_install() and sys.platform != "win32":
os.execv(sys.executable, sys.argv) # noqa: S606 # nosec B606 # The PyInstaller onefile bootloader passes its state to the child
# process via environment variables; if they leak into the re-exec,
# the new binary reuses the old extracted application instead of
# unpacking itself, so the pre-update version runs again.
env = {
key: value
for key, value in os.environ.items()
if not key.startswith("_PYI_") and key != "_MEIPASS2"
}
os.execve(sys.executable, sys.argv, env) # noqa: S606 # nosec B606
sys.exit(0) sys.exit(0)
check_docker_installed() check_docker_installed()
+6 -3
View File
@@ -14,7 +14,7 @@ import logging
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from strix.config import Settings, codex, load_settings from strix.config import Settings, load_settings, subscription
from strix.core.paths import run_dir_for from strix.core.paths import run_dir_for
from strix.interface.utils import ( from strix.interface.utils import (
assign_workspace_subdirs, assign_workspace_subdirs,
@@ -78,6 +78,7 @@ async def preflight_model_connection(
request_timeout=resolved_settings.llm.timeout, request_timeout=resolved_settings.llm.timeout,
prompt_cache=False, prompt_cache=False,
extra_headers=resolved_settings.llm.extra_headers, extra_headers=resolved_settings.llm.extra_headers,
has_tools=False,
) )
await asyncio.wait_for( await asyncio.wait_for(
model.get_response( model.get_response(
@@ -225,7 +226,7 @@ def telemetry_start(args: argparse.Namespace) -> None:
model = load_settings().llm.model model = load_settings().llm.model
kwargs = { kwargs = {
"model": model, "model": model,
"auth_mode": codex.auth_mode(model), "auth_mode": subscription.auth_mode(model),
"scan_mode": args.scan_mode, "scan_mode": args.scan_mode,
"is_whitebox": is_whitebox_scan(args.targets_info), "is_whitebox": is_whitebox_scan(args.targets_info),
"interactive": not args.non_interactive, "interactive": not args.non_interactive,
@@ -240,13 +241,15 @@ def _persist_run_record(args: argparse.Namespace) -> None:
run_dir = run_dir_for(args.run_name) run_dir = run_dir_for(args.run_name)
run_dir.mkdir(parents=True, exist_ok=True) run_dir.mkdir(parents=True, exist_ok=True)
model = load_settings().llm.model
run_record = { run_record = {
"run_id": args.run_name, "run_id": args.run_name,
"run_name": args.run_name, "run_name": args.run_name,
"status": "running", "status": "running",
"start_time": datetime.now(UTC).isoformat(), "start_time": datetime.now(UTC).isoformat(),
"end_time": None, "end_time": None,
"auth_mode": codex.auth_mode(load_settings().llm.model), "auth_mode": subscription.auth_mode(model),
"subscription_provider": subscription.provider_label(model),
"targets_info": args.targets_info, "targets_info": args.targets_info,
"scan_mode": args.scan_mode, "scan_mode": args.scan_mode,
"instruction": args.instruction, "instruction": args.instruction,
+5 -1
View File
@@ -24,7 +24,7 @@ from strix.interface.tui.backend.projection import (
sanitize_terminal_text, sanitize_terminal_text,
terminal_projection, terminal_projection,
) )
from strix.interface.utils import is_subscription_run from strix.interface.utils import is_subscription_run, subscription_label
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -162,8 +162,11 @@ class TuiController:
if self.report_state is not None: if self.report_state is not None:
usage = dict(self.report_state.get_total_llm_usage()) usage = dict(self.report_state.get_total_llm_usage())
subscription = False subscription = False
subscription_name = ""
with contextlib.suppress(Exception): with contextlib.suppress(Exception):
subscription = is_subscription_run(self.report_state) subscription = is_subscription_run(self.report_state)
if subscription:
subscription_name = subscription_label(self.report_state)
model_warning = "" model_warning = ""
if model and not is_recommended_or_frontier_model(model): if model and not is_recommended_or_frontier_model(model):
model_warning = ( model_warning = (
@@ -200,6 +203,7 @@ class TuiController:
], ],
"usage": terminal_projection(usage, max_string=256, max_items=20), "usage": terminal_projection(usage, max_string=256, max_items=20),
"subscription": subscription, "subscription": subscription,
"subscription_label": terminal_projection(subscription_name, max_string=64),
"viewer_status": self.viewer_status, "viewer_status": self.viewer_status,
"viewer_url": terminal_projection(self.viewer_url, max_string=1024), "viewer_url": terminal_projection(self.viewer_url, max_string=1024),
"error": terminal_projection(self.error, max_string=2 * 1024), "error": terminal_projection(self.error, max_string=2 * 1024),
@@ -175,6 +175,7 @@ def bounded_state_projection(state: dict[str, Any]) -> dict[str, Any]:
"messages": [], "messages": [],
"usage": {}, "usage": {},
"subscription": state["subscription"], "subscription": state["subscription"],
"subscription_label": state["subscription_label"],
"viewer_status": state["viewer_status"], "viewer_status": state["viewer_status"],
"viewer_url": None, "viewer_url": None,
"error": terminal_projection(state["error"], max_string=256), "error": terminal_projection(state["error"], max_string=256),
+13 -2
View File
@@ -1067,11 +1067,12 @@ func TestBudgetPauseShowsOneWarningToastUntilResumed(t *testing.T) {
func TestStatsViewShowsSubscription(t *testing.T) { func TestStatsViewShowsSubscription(t *testing.T) {
model := New(nil) model := New(nil)
model.snapshot.Model = "gpt-5" model.snapshot.Model = "grok/grok-4"
model.snapshot.Subscription = true model.snapshot.Subscription = true
model.snapshot.SubscriptionLabel = "Grok subscription"
model.snapshot.Usage = map[string]any{"total_tokens": float64(1200), "cost": 3.5} model.snapshot.Usage = map[string]any{"total_tokens": float64(1200), "cost": 3.5}
stats := ansi.Strip(model.statsView()) stats := ansi.Strip(model.statsView())
if !strings.Contains(stats, "ChatGPT subscription") { if !strings.Contains(stats, "Grok subscription") {
t.Fatalf("stats missing subscription line: %q", stats) t.Fatalf("stats missing subscription line: %q", stats)
} }
if strings.Contains(stats, "$") { if strings.Contains(stats, "$") {
@@ -1079,6 +1080,16 @@ func TestStatsViewShowsSubscription(t *testing.T) {
} }
} }
func TestStatsViewSubscriptionFallsBackWithoutLabel(t *testing.T) {
model := New(nil)
model.snapshot.Model = "gpt-5"
model.snapshot.Subscription = true
stats := ansi.Strip(model.statsView())
if !strings.Contains(stats, "Subscription") {
t.Fatalf("stats missing generic subscription line: %q", stats)
}
}
func TestVulnerabilityMarkdownReport(t *testing.T) { func TestVulnerabilityMarkdownReport(t *testing.T) {
report := vulnerabilityMarkdownReport(map[string]any{ report := vulnerabilityMarkdownReport(map[string]any{
"title": "SQLi in login", "title": "SQLi in login",
+5 -1
View File
@@ -596,7 +596,11 @@ func (m Model) statsView() string {
if b.Len() > 0 { if b.Len() > 0 {
b.WriteString("\n") b.WriteString("\n")
} }
b.WriteString(lipgloss.NewStyle().Foreground(green).Render("ChatGPT subscription")) label := m.snapshot.SubscriptionLabel
if label == "" {
label = "Subscription"
}
b.WriteString(lipgloss.NewStyle().Foreground(green).Render(label))
} }
total := numberValue(m.snapshot.Usage["total_tokens"]) total := numberValue(m.snapshot.Usage["total_tokens"])
if total > 0 { if total > 0 {
@@ -68,6 +68,7 @@ type Snapshot struct {
Vulnerabilities []map[string]any `json:"-"` Vulnerabilities []map[string]any `json:"-"`
Usage map[string]any `json:"usage"` Usage map[string]any `json:"usage"`
Subscription bool `json:"subscription"` Subscription bool `json:"subscription"`
SubscriptionLabel string `json:"subscription_label"`
ViewerStatus string `json:"viewer_status"` ViewerStatus string `json:"viewer_status"`
ViewerURL *string `json:"viewer_url"` ViewerURL *string `json:"viewer_url"`
Error *string `json:"error"` Error *string `json:"error"`
+22 -4
View File
@@ -262,9 +262,27 @@ def is_subscription_run(report_state: Any) -> bool:
record = getattr(report_state, "run_record", None) record = getattr(report_state, "run_record", None)
if isinstance(record, dict) and record.get("auth_mode"): if isinstance(record, dict) and record.get("auth_mode"):
return record.get("auth_mode") == "subscription" return record.get("auth_mode") == "subscription"
from strix.config import codex from strix.config import subscription
return codex.auth_mode(load_settings().llm.model) == "subscription" return subscription.auth_mode(load_settings().llm.model) == "subscription"
def subscription_label(report_state: Any) -> str:
"""Human label for the active model subscription (e.g. "Grok subscription").
Prefers the persisted run record so a resumed run keeps its original provider
even if STRIX_LLM later points at a different one; falls back to current
settings.
"""
record = getattr(report_state, "run_record", None)
if isinstance(record, dict):
provider = record.get("subscription_provider")
if isinstance(provider, str) and provider:
return f"{provider} subscription"
from strix.config import subscription
label = subscription.provider_label(load_settings().llm.model)
return f"{label} subscription" if label else "Subscription"
def _int_stat(usage: dict[str, Any], key: str) -> int: def _int_stat(usage: dict[str, Any], key: str) -> int:
@@ -368,7 +386,7 @@ def build_live_stats_text(report_state: Any) -> Text:
stats_text.append(str(model), style="white") stats_text.append(str(model), style="white")
if is_subscription_run(report_state): if is_subscription_run(report_state):
stats_text.append(" · ", style="dim white") stats_text.append(" · ", style="dim white")
stats_text.append("ChatGPT subscription", style="#22c55e") stats_text.append(subscription_label(report_state), style="#22c55e")
stats_text.append("\n") stats_text.append("\n")
vuln_count = len(report_state.vulnerability_reports) vuln_count = len(report_state.vulnerability_reports)
@@ -414,7 +432,7 @@ def build_tui_stats_text(report_state: Any) -> Text:
subscription = is_subscription_run(report_state) subscription = is_subscription_run(report_state)
if subscription: if subscription:
stats_text.append("\n") stats_text.append("\n")
stats_text.append("ChatGPT subscription", style="#22c55e") stats_text.append(subscription_label(report_state), style="#22c55e")
usage = _llm_usage(report_state) usage = _llm_usage(report_state)
if usage and _int_stat(usage, "total_tokens") > 0: if usage and _int_stat(usage, "total_tokens") > 0:
@@ -101,6 +101,7 @@ export function RunDetails({
const totalTokens = num(usage.total_tokens); const totalTokens = num(usage.total_tokens);
const cost = num(usage.cost); const cost = num(usage.cost);
const subscription = str(raw.auth_mode) === "subscription"; const subscription = str(raw.auth_mode) === "subscription";
const subscriptionProvider = str(raw.subscription_provider);
const sub = (n: number, word: string) => ( const sub = (n: number, word: string) => (
<span className="text-[#666]"> ({formatNumber(n)} {word})</span> <span className="text-[#666]"> ({formatNumber(n)} {word})</span>
@@ -180,7 +181,7 @@ export function RunDetails({
<Field label="Provider"> <Field label="Provider">
<span className="inline-flex items-center gap-1.5"> <span className="inline-flex items-center gap-1.5">
<span className="rounded-full border border-[#22c55e]/40 bg-[#22c55e]/10 px-2 py-0.5 text-[11px] text-[#22c55e]"> <span className="rounded-full border border-[#22c55e]/40 bg-[#22c55e]/10 px-2 py-0.5 text-[11px] text-[#22c55e]">
ChatGPT subscription {subscriptionProvider ? `${subscriptionProvider} subscription` : "Subscription"}
</span> </span>
</span> </span>
</Field> </Field>
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -6,7 +6,7 @@
<meta name="viewport" content="width=device-width, initial-scale=1.0" /> <meta name="viewport" content="width=device-width, initial-scale=1.0" />
<meta name="color-scheme" content="dark" /> <meta name="color-scheme" content="dark" />
<title>Strix Results</title> <title>Strix Results</title>
<script type="module" crossorigin src="./assets/index-DBJ-RJqo.js"></script> <script type="module" crossorigin src="./assets/index-XDX3roAH.js"></script>
<link rel="stylesheet" crossorigin href="./assets/index-DKbLYAbP.css"> <link rel="stylesheet" crossorigin href="./assets/index-DKbLYAbP.css">
</head> </head>
<body> <body>
+34 -1
View File
@@ -6,6 +6,7 @@ import json
import logging import logging
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from strix.config import subscription
from strix.core.paths import run_record_path from strix.core.paths import run_record_path
from strix.interface.tui.live_view import TuiLiveView from strix.interface.tui.live_view import TuiLiveView
@@ -57,7 +58,39 @@ def read_run_summary(run_dir: Path) -> dict[str, Any]:
record = {} record = {}
status = record.get("status") status = record.get("status")
finished = status in _TERMINAL_STATUSES and bool(record.get("end_time")) finished = status in _TERMINAL_STATUSES and bool(record.get("end_time"))
return {**record, "finished": finished} summary = {**record, "finished": finished}
_backfill_subscription_provider(summary)
return summary
def _first_recorded_model(record: dict[str, Any]) -> str | None:
"""The first non-empty per-agent model slug in a run record, or None."""
usage = record.get("llm_usage")
if not isinstance(usage, dict):
return None
agents = usage.get("agents")
if not isinstance(agents, list):
return None
for agent in agents:
if isinstance(agent, dict):
model = agent.get("model")
if isinstance(model, str) and model:
return model
return None
def _backfill_subscription_provider(record: dict[str, Any]) -> None:
"""Name the subscription provider for runs recorded before that field
existed, deriving it from the recorded ``provider/model`` slug so the viewer
labels them correctly without a rescan. Newer runs already carry the field.
"""
if record.get("subscription_provider"):
return
if record.get("auth_mode") != "subscription":
return
label = subscription.provider_label(_first_recorded_model(record))
if label:
record["subscription_provider"] = label
def primary_target(record: dict[str, Any]) -> str | None: def primary_target(record: dict[str, Any]) -> str | None:
+1
View File
@@ -294,6 +294,7 @@ async def _summarize(model: str, prompt: str, max_tokens: int) -> str | None:
request_timeout=llm.timeout, request_timeout=llm.timeout,
prompt_cache=False, prompt_cache=False,
extra_headers=llm.extra_headers, extra_headers=llm.extra_headers,
has_tools=False,
).resolve(ModelSettings(max_tokens=max_tokens)) ).resolve(ModelSettings(max_tokens=max_tokens))
try: try:
response = ( response = (
+9 -5
View File
@@ -10,7 +10,7 @@ from typing import Any
import litellm import litellm
from strix.config import load_settings from strix.config import load_settings, subscription
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -19,7 +19,6 @@ logger = logging.getLogger(__name__)
# ``litellm/``, ``ollama/`` ...). Strip a leading provider segment on lookup. # ``litellm/``, ``ollama/`` ...). Strip a leading provider segment on lookup.
_STRIPPABLE_PREFIXES = ( _STRIPPABLE_PREFIXES = (
"openai/", "openai/",
"chatgpt/",
"litellm/", "litellm/",
"any-llm/", "any-llm/",
"ollama/", "ollama/",
@@ -30,6 +29,8 @@ _DEFAULT_OUTPUT_TOKENS = 8_192
def _lookup_key(model: str) -> str: def _lookup_key(model: str) -> str:
if subscription.provider_for_model(model) is not None:
return subscription.litellm_model_name(model) or model
for prefix in _STRIPPABLE_PREFIXES: for prefix in _STRIPPABLE_PREFIXES:
if model.startswith(prefix): if model.startswith(prefix):
return model[len(prefix) :] return model[len(prefix) :]
@@ -46,9 +47,12 @@ def _safe_get_model_info(model: str) -> dict[str, Any] | None:
@lru_cache(maxsize=128) @lru_cache(maxsize=128)
def _model_info(model: str) -> dict[str, int]: def _model_info(model: str) -> dict[str, int]:
lookup_key = _lookup_key(model) lookup_key = _lookup_key(model)
# Provider-qualified ChatGPT lookups may start a synchronous device-login # Subscription prefixes are never LiteLLM keys, and a provider-qualified
# poll. LiteLLM keys the metadata by the underlying model slug. # ChatGPT lookup may start a synchronous device-login poll: only ask about
candidates = (lookup_key,) if model.startswith("chatgpt/") else (model, lookup_key) # the resolved name.
candidates = (
(lookup_key,) if subscription.provider_for_model(model) is not None else (model, lookup_key)
)
for candidate in candidates: for candidate in candidates:
info = _safe_get_model_info(candidate) info = _safe_get_model_info(candidate)
if info is not None: if info is not None:
+1
View File
@@ -62,6 +62,7 @@ def _dedupe_model_settings(
# must never receive the main endpoint's credentials. A dedicated model # must never receive the main endpoint's credentials. A dedicated model
# gets its own DEDUPE_LLM_EXTRA_HEADERS instead. # gets its own DEDUPE_LLM_EXTRA_HEADERS instead.
extra_headers=dedupe.extra_headers if dedupe.model else llm.extra_headers, extra_headers=dedupe.extra_headers if dedupe.model else llm.extra_headers,
has_tools=False,
) )
extra = _dedupe_extra_args(dedupe) extra = _dedupe_extra_args(dedupe)
if extra: if extra:
+54
View File
@@ -0,0 +1,54 @@
"""LiteLLM model-name resolution for local cost estimates."""
from __future__ import annotations
from functools import lru_cache
from typing import Any, cast
@lru_cache(maxsize=512)
def resolve_litellm_model(model: str) -> str | None:
"""Return a provider-qualified model name that LiteLLM can price."""
try:
import litellm
normalized = model.strip()
for prefix in ("litellm/", "any-llm/", "openai/"):
if normalized.startswith(prefix):
normalized = normalized.removeprefix(prefix)
break
if not normalized:
return None
model_cost = cast(
"dict[str, dict[str, Any]]",
getattr(litellm, "model_cost"), # noqa: B009
)
bare_entry = model_cost.get(normalized)
if "/" not in normalized and isinstance(bare_entry, dict):
provider = bare_entry.get("litellm_provider")
if isinstance(provider, str) and provider:
return f"{provider}/{normalized}"
if "/" in normalized and isinstance(bare_entry, dict):
return normalized
names = [normalized]
if "/" in normalized:
names.append(normalized.rsplit("/", 1)[-1])
for name in names:
matches = sorted(key for key in model_cost if key.endswith(f"/{name}"))
if not matches:
continue
prices = {
(
model_cost[key].get("input_cost_per_token"),
model_cost[key].get("output_cost_per_token"),
)
for key in matches
if isinstance(model_cost.get(key), dict)
}
if len(matches) == 1 or len(prices) == 1:
return matches[0]
return None # noqa: TRY300
except Exception: # noqa: BLE001
return None
+10 -4
View File
@@ -11,9 +11,10 @@ from uuid import uuid4
from agents.usage import Usage from agents.usage import Usage
from strix.config import codex from strix.config import subscription
from strix.config.loader import load_settings from strix.config.loader import load_settings
from strix.core.paths import run_dir_for from strix.core.paths import run_dir_for
from strix.report.pricing import resolve_litellm_model
from strix.report.sarif import write_sarif from strix.report.sarif import write_sarif
from strix.report.usage import LLMUsageLedger from strix.report.usage import LLMUsageLedger
from strix.report.writer import ( from strix.report.writer import (
@@ -122,7 +123,8 @@ class ReportState:
self.scan_results: dict[str, Any] | None = None self.scan_results: dict[str, Any] | None = None
self.scan_config: dict[str, Any] | None = None self.scan_config: dict[str, Any] | None = None
self._llm_usage = LLMUsageLedger() self._llm_usage = LLMUsageLedger()
auth_mode = codex.auth_mode(load_settings().llm.model) model = load_settings().llm.model
auth_mode = subscription.auth_mode(model)
self._llm_usage.zero_cost = auth_mode == "subscription" self._llm_usage.zero_cost = auth_mode == "subscription"
self.run_record: dict[str, Any] = { self.run_record: dict[str, Any] = {
"run_id": self.run_id, "run_id": self.run_id,
@@ -131,6 +133,7 @@ class ReportState:
"end_time": None, "end_time": None,
"status": "running", "status": "running",
"auth_mode": auth_mode, "auth_mode": auth_mode,
"subscription_provider": subscription.provider_label(model),
"targets_info": [], "targets_info": [],
"llm_usage": self._build_llm_usage_record(), "llm_usage": self._build_llm_usage_record(),
} }
@@ -696,10 +699,13 @@ def _estimate_response_cost(kwargs: Any, completion_response: Any) -> float | No
candidates.append(model.rsplit("/", 1)[-1]) candidates.append(model.rsplit("/", 1)[-1])
for candidate in candidates: for candidate in candidates:
resolved = resolve_litellm_model(candidate)
if not resolved:
continue
try: try:
value = completion_cost( value = completion_cost(
completion_response={"model": candidate, "usage": usage_payload}, completion_response={"model": resolved, "usage": usage_payload},
model=candidate, model=resolved,
) )
except Exception: # nosec B112 # noqa: BLE001, S112 except Exception: # nosec B112 # noqa: BLE001, S112
continue continue
+30 -29
View File
@@ -7,6 +7,8 @@ from typing import Any
from agents.usage import Usage, deserialize_usage, serialize_usage from agents.usage import Usage, deserialize_usage, serialize_usage
from strix.report.pricing import resolve_litellm_model
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -18,7 +20,9 @@ class LLMUsageLedger:
self._total_usage = Usage() self._total_usage = Usage()
self._agent_usage: dict[str, Usage] = {} self._agent_usage: dict[str, Usage] = {}
self._agent_metadata: dict[str, dict[str, str]] = {} self._agent_metadata: dict[str, dict[str, str]] = {}
self._total_cost = 0.0 self._observed_cost = 0.0
self._estimated_cost = 0.0
self._has_observed_cost = False
# When True, tokens are still tracked but cost stays $0 — the run is on a # When True, tokens are still tracked but cost stays $0 — the run is on a
# model subscription, so there is no metered per-token charge to report. # model subscription, so there is no metered per-token charge to report.
self.zero_cost = False self.zero_cost = False
@@ -44,10 +48,10 @@ class LLMUsageLedger:
if model: if model:
metadata["model"] = model metadata["model"] = model
if not self.zero_cost and not _is_litellm_routed(model): if not self.zero_cost:
estimated = _estimate_litellm_cost(usage, model) estimated = _estimate_litellm_cost(usage, model)
if estimated: if estimated:
self._total_cost += estimated self._estimated_cost += estimated
return True return True
@@ -55,15 +59,18 @@ class LLMUsageLedger:
if self.zero_cost: if self.zero_cost:
return return
if isinstance(cost, int | float) and cost > 0: if isinstance(cost, int | float) and cost > 0:
self._total_cost += float(cost) self._observed_cost += float(cost)
self._has_observed_cost = True
@property @property
def total_cost(self) -> float: def total_cost(self) -> float:
return _round_cost(self._total_cost) if self.zero_cost:
return 0.0
return _round_cost(self._observed_cost if self._has_observed_cost else self._estimated_cost)
def to_record(self) -> dict[str, Any]: def to_record(self) -> dict[str, Any]:
record = serialize_usage(self._total_usage) record = serialize_usage(self._total_usage)
record["cost"] = _round_cost(self._total_cost) record["cost"] = self.total_cost
record["agents"] = [] record["agents"] = []
agent_tokens = {aid: _resolve_total_tokens(u) for aid, u in self._agent_usage.items()} agent_tokens = {aid: _resolve_total_tokens(u) for aid, u in self._agent_usage.items()}
@@ -72,7 +79,7 @@ class LLMUsageLedger:
usage = self._agent_usage[agent_id] usage = self._agent_usage[agent_id]
metadata = self._agent_metadata.get(agent_id, {}) metadata = self._agent_metadata.get(agent_id, {})
agent_cost = ( agent_cost = (
self._total_cost * (agent_tokens[agent_id] / total_tokens) if total_tokens else 0.0 self.total_cost * (agent_tokens[agent_id] / total_tokens) if total_tokens else 0.0
) )
agent_record = serialize_usage(usage) agent_record = serialize_usage(usage)
@@ -92,7 +99,9 @@ class LLMUsageLedger:
self._total_usage = Usage() self._total_usage = Usage()
self._agent_usage.clear() self._agent_usage.clear()
self._agent_metadata.clear() self._agent_metadata.clear()
self._total_cost = 0.0 self._observed_cost = 0.0
self._estimated_cost = 0.0
self._has_observed_cost = False
if not isinstance(raw_usage, dict): if not isinstance(raw_usage, dict):
return return
@@ -103,7 +112,9 @@ class LLMUsageLedger:
logger.exception("Failed to hydrate aggregate llm_usage from run.json") logger.exception("Failed to hydrate aggregate llm_usage from run.json")
self._total_usage = Usage() self._total_usage = Usage()
self._total_cost = _float_or_zero(raw_usage.get("cost")) persisted_cost = _float_or_zero(raw_usage.get("cost"))
self._observed_cost = persisted_cost
self._estimated_cost = persisted_cost
for raw_agent in raw_usage.get("agents") or []: for raw_agent in raw_usage.get("agents") or []:
if not isinstance(raw_agent, dict): if not isinstance(raw_agent, dict):
@@ -136,15 +147,6 @@ def _resolve_total_tokens(usage: Usage) -> int:
return prompt + completion return prompt + completion
def _is_litellm_routed(model: str | None) -> bool:
if not model:
return False
name = model.strip().lower()
if "/" not in name:
return False
return not name.startswith("openai/")
def _usage_has_activity(usage: Usage) -> bool: def _usage_has_activity(usage: Usage) -> bool:
return bool( return bool(
usage.requests usage.requests
@@ -201,24 +203,23 @@ def _estimate_litellm_entry_cost(entry: Any, model: str) -> float | None:
candidates = [model] candidates = [model]
if "/" in model: if "/" in model:
candidates.append(model.split("/", 1)[-1]) candidates.append(model.rsplit("/", 1)[-1])
cost: Any = None
for candidate in candidates: for candidate in candidates:
resolved = resolve_litellm_model(candidate)
if not resolved:
continue
try: try:
cost = completion_cost( cost = completion_cost(
completion_response={"model": candidate, "usage": usage_payload}, completion_response={"model": resolved, "usage": usage_payload},
model=model, model=resolved,
) )
break
except Exception: # nosec B112 # noqa: BLE001, S112 except Exception: # nosec B112 # noqa: BLE001, S112
continue continue
if cost > 0:
if cost is None: return float(cost)
logger.debug("LiteLLM cost estimate unavailable for model %s", model) logger.debug("LiteLLM cost estimate unavailable for model %s", model)
return None return None
return cost if isinstance(cost, int | float) and cost >= 0 else None
def _litellm_model_name(model: str | None) -> str | None: def _litellm_model_name(model: str | None) -> str | None:
+35 -2
View File
@@ -58,6 +58,26 @@ agent-browser screenshot
The browser stays running across commands so these feel like a single The browser stays running across commands so these feel like a single
session. Use `agent-browser close` (or `close --all`) when you're done. session. Use `agent-browser close` (or `close --all`) when you're done.
The default session is **shared with every other agent in the sandbox** — if
another agent navigates it, your page and your refs are gone from under you. So
claim your own by passing `--session <your-agent-name>` on **every** command:
```bash
agent-browser --session recon-3 open https://example.com
agent-browser --session recon-3 snapshot -i
agent-browser --session recon-3 close # when done with the target
```
The examples in the rest of this skill omit `--session` to keep them readable;
keep passing yours. Each session is a separate Chromium (~340 MB) on a shared
box, so hold one rather than several, and close it when you're finished.
A browser left idle for 3 minutes is reclaimed automatically to free memory for
the other agents; the next command relaunches it, but the page, tabs, refs and
cookies are gone. If you're authenticated and about to go do something else for a
while, save the state first (see
[Persist session across runs](#persist-session-across-runs)).
## Reading a page ## Reading a page
```bash ```bash
@@ -307,6 +327,16 @@ agent-browser --session b fill @e1 "bob@test.com"
`AGENT_BROWSER_SESSION=myapp` sets the default session for the current `AGENT_BROWSER_SESSION=myapp` sets the default session for the current
shell. shell.
Use a session named after yourself for your own work — that's what keeps a
concurrent agent from navigating the page out from under you. Every session is a
separate Chromium though, so hold one at a time rather than a collection, and
close each one when its flow is finished:
```bash
agent-browser --session a close
agent-browser --session b close
```
### Mock network requests ### Mock network requests
```bash ```bash
@@ -368,8 +398,11 @@ agent-browser dialog dismiss # cancel
## Readiness & recovery ## Readiness & recovery
The first `agent-browser open` in a session launches the headless-Chrome The first `agent-browser open` in a session launches the headless-Chrome
daemon; later commands reuse it. Distinguish the two failure modes and react daemon; later commands reuse it. A daemon left idle for 3 minutes shuts itself
differently — do **not** blindly re-run the same failing command in a loop: down to free memory for the other agents, so an `open` after a long gap is a
fresh browser rather than a resumed one — expect to re-navigate, and re-`state
load` if you were logged in. Distinguish the failure modes and react differently
— do **not** blindly re-run the same failing command in a loop:
- **Daemon / connection failure** (`Failed to connect`, `connection refused`, - **Daemon / connection failure** (`Failed to connect`, `connection refused`,
socket missing, `browser not running`): the daemon isn't up or has died. Run socket missing, `browser not running`): the daemon isn't up or has died. Run
+52 -16
View File
@@ -6,23 +6,34 @@ from typing import TYPE_CHECKING, Any
import pytest import pytest
from strix.config import codex from strix.config import codex, grok
from strix.interface import auth_cli from strix.interface import auth_cli
if TYPE_CHECKING: if TYPE_CHECKING:
from pathlib import Path from pathlib import Path
_CHATGPT = auth_cli._PROVIDERS["chatgpt"]
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(codex, "AUTH_PATH", tmp_path / "home" / ".strix" / "subscription-auth.json") store = tmp_path / "home" / ".strix" / "subscription-auth.json"
monkeypatch.setattr(codex, "AUTH_PATH", store)
monkeypatch.setattr(grok, "AUTH_PATH", store)
def test_login_provider_is_chatgpt() -> None: def test_default_provider_is_chatgpt() -> None:
assert auth_cli.LOGIN_PROVIDER == "chatgpt" assert auth_cli._DEFAULT_PROVIDER == "chatgpt"
assert codex.PROVIDER in auth_cli._ACCEPTED_PROVIDERS assert set(auth_cli._PROVIDERS) == {"chatgpt", "grok"}
assert "chatgpt" in auth_cli._ACCEPTED_PROVIDERS
def test_provider_aliases_resolve() -> None:
assert auth_cli._resolve_provider(codex.PROVIDER) is _CHATGPT
assert auth_cli._resolve_provider("ChatGPT") is _CHATGPT
assert auth_cli._resolve_provider("grok") is auth_cli._PROVIDERS["grok"]
assert auth_cli._resolve_provider("xai") is auth_cli._PROVIDERS["grok"]
assert auth_cli._resolve_provider("gemini") is None
def test_unknown_subcommand_returns_usage_error() -> None: def test_unknown_subcommand_returns_usage_error() -> None:
@@ -51,32 +62,32 @@ def test_finish_requires_state_on_loopback(monkeypatch: pytest.MonkeyPatch) -> N
# Loopback (require_state=True): missing or mismatched state is rejected. # Loopback (require_state=True): missing or mismatched state is rejected.
with pytest.raises(codex.CodexAuthError) as missing: with pytest.raises(codex.CodexAuthError) as missing:
auth_cli._finish("code", None, "verifier", "expected", require_state=True) auth_cli._finish(_CHATGPT, "code", None, "verifier", "expected", require_state=True)
assert missing.value.code == "state_mismatch" assert missing.value.code == "state_mismatch"
with pytest.raises(codex.CodexAuthError) as mismatch: with pytest.raises(codex.CodexAuthError) as mismatch:
auth_cli._finish("code", "wrong", "verifier", "expected", require_state=True) auth_cli._finish(_CHATGPT, "code", "wrong", "verifier", "expected", require_state=True)
assert mismatch.value.code == "state_mismatch" assert mismatch.value.code == "state_mismatch"
# Matching state proceeds to the exchange. # Matching state proceeds to the exchange.
assert auth_cli._finish("code", "expected", "verifier", "expected", require_state=True) == { assert auth_cli._finish(
"ok": True _CHATGPT, "code", "expected", "verifier", "expected", require_state=True
} ) == {"ok": True}
def test_finish_manual_paste_allows_absent_state(monkeypatch: pytest.MonkeyPatch) -> None: def test_finish_manual_paste_allows_absent_state(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(codex, "exchange_code", lambda *_: {"ok": True}) monkeypatch.setattr(codex, "exchange_code", lambda *_: {"ok": True})
# Manual paste (require_state=False): a bare code with no state is accepted, # Manual paste (require_state=False): a bare code with no state is accepted,
# but a present-and-wrong state is still rejected. # but a present-and-wrong state is still rejected.
assert auth_cli._finish("code", None, "verifier", "expected", require_state=False) == { assert auth_cli._finish(
"ok": True _CHATGPT, "code", None, "verifier", "expected", require_state=False
} ) == {"ok": True}
with pytest.raises(codex.CodexAuthError): with pytest.raises(codex.CodexAuthError):
auth_cli._finish("code", "wrong", "verifier", "expected", require_state=False) auth_cli._finish(_CHATGPT, "code", "wrong", "verifier", "expected", require_state=False)
def test_finish_rejects_missing_code() -> None: def test_finish_rejects_missing_code() -> None:
with pytest.raises(codex.CodexAuthError) as exc: with pytest.raises(codex.CodexAuthError) as exc:
auth_cli._finish(None, "expected", "verifier", "expected", require_state=True) auth_cli._finish(_CHATGPT, None, "expected", "verifier", "expected", require_state=True)
assert exc.value.code == "no_code" assert exc.value.code == "no_code"
@@ -84,6 +95,31 @@ def test_model_subcommand_removed() -> None:
assert auth_cli.run_auth(["model", "gpt-5.5"]) == 2 assert auth_cli.run_auth(["model", "gpt-5.5"]) == 2
def _sign_in_both() -> None:
codex.save_record({"type": "oauth", "access": "c", "refresh": "r", "account_id": "a"})
grok.save_record({"type": "oauth", "access": "g", "refresh": "r"})
def test_logout_all_removes_every_provider() -> None:
_sign_in_both()
assert codex.is_authenticated()
assert grok.is_authenticated()
assert auth_cli.run_auth(["logout"]) == 0
assert not codex.is_authenticated()
assert not grok.is_authenticated()
def test_logout_single_provider_leaves_the_other() -> None:
_sign_in_both()
assert auth_cli.run_auth(["logout", "grok"]) == 0
assert codex.is_authenticated()
assert not grok.is_authenticated()
@pytest.mark.parametrize("provider", ["chatgpt", "codex", "ChatGPT"]) @pytest.mark.parametrize("provider", ["chatgpt", "codex", "ChatGPT"])
def test_login_accepts_provider_aliases(provider: str, monkeypatch: pytest.MonkeyPatch) -> None: def test_login_accepts_provider_aliases(provider: str, monkeypatch: pytest.MonkeyPatch) -> None:
reached = {"flow": False} reached = {"flow": False}
+11
View File
@@ -39,6 +39,17 @@ def test_context_window_chatgpt_prefix_skips_provider_auth(
context_budget._model_info.cache_clear() context_budget._model_info.cache_clear()
def test_context_window_grok_prefix_resolves_to_xai() -> None:
# LiteLLM maps xAI models only provider-qualified: neither "grok/grok-4" nor
# bare "grok-4" resolves, so the subscription prefix becomes "xai/".
context_budget._model_info.cache_clear()
try:
assert context_budget.context_window("grok/grok-4") == 256_000
assert context_budget.output_limit("grok/grok-4") == 256_000
finally:
context_budget._model_info.cache_clear()
def test_context_window_unmapped_uses_fallback(monkeypatch: pytest.MonkeyPatch) -> None: def test_context_window_unmapped_uses_fallback(monkeypatch: pytest.MonkeyPatch) -> None:
context_budget._model_info.cache_clear() context_budget._model_info.cache_clear()
+1 -1
View File
@@ -143,7 +143,7 @@ def test_cost_callback_estimates_cost_with_bare_model_fallback() -> None:
} }
def fake_completion_cost(**kwargs: object) -> float: def fake_completion_cost(**kwargs: object) -> float:
if kwargs["model"] == "gpt-4o-mini": if kwargs["model"] == "openai/gpt-4o-mini":
return 0.025 return 0.025
raise ValueError(kwargs["model"]) raise ValueError(kwargs["model"])
+267
View File
@@ -0,0 +1,267 @@
"""Tests for Grok (xAI) subscription auth: PKCE, token handling, store."""
from __future__ import annotations
import base64
import hashlib
import json
import time
from typing import TYPE_CHECKING, Any
from unittest import mock
import pytest
import requests
from strix.config import grok
if TYPE_CHECKING:
from pathlib import Path
@pytest.fixture(autouse=True)
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path:
path = tmp_path / "home" / ".strix" / "subscription-auth.json"
monkeypatch.setattr(grok, "AUTH_PATH", path)
return path
def test_pkce_challenge_matches_verifier_and_is_unpadded() -> None:
verifier, challenge = grok.generate_pkce()
expected = (
base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
)
assert challenge == expected
assert "=" not in verifier
assert "=" not in challenge
def test_authorize_url_carries_pkce_client_and_grok_scope() -> None:
url = grok.build_authorize_url("chal", "st8")
assert grok.AUTHORIZE_URL in url
assert "code_challenge=chal" in url
assert "code_challenge_method=S256" in url
assert f"client_id={grok.CLIENT_ID}" in url
assert "state=st8" in url
# The Grok-CLI scope is what unlocks subscription inference.
assert "grok-cli%3Aaccess" in url
assert "offline_access" in url
def test_redirect_uri_is_loopback() -> None:
assert grok.REDIRECT_URI == "http://127.0.0.1:56121/callback"
def test_post_form_returns_parsed_body() -> None:
resp = mock.MagicMock()
resp.status_code = 200
resp.content = b'{"access_token": "tok"}'
resp.__enter__.return_value = resp
with mock.patch.object(requests, "post", return_value=resp) as post:
data = grok._post_form({"grant_type": "refresh_token"})
assert data == {"access_token": "tok"}
assert post.call_args.kwargs["timeout"] == grok._TOKEN_TIMEOUT
def test_post_form_raises_on_http_error() -> None:
resp = mock.MagicMock()
resp.status_code = 400
resp.text = "invalid_grant"
resp.__enter__.return_value = resp
with (
mock.patch.object(requests, "post", return_value=resp),
pytest.raises(grok.GrokAuthError) as exc,
):
grok._post_form({"grant_type": "refresh_token"})
assert exc.value.code == "token_http_error"
@pytest.mark.parametrize(
("value", "expected"),
[
("http://127.0.0.1:56121/callback?code=AAA&state=BBB", ("AAA", "BBB")),
("AAA#BBB", ("AAA", "BBB")),
("code=AAA&state=BBB", ("AAA", "BBB")),
("AAA", ("AAA", None)),
("", (None, None)),
],
)
def test_parse_redirect_input(value: str, expected: tuple[str | None, str | None]) -> None:
assert grok.parse_redirect_input(value) == expected
@pytest.mark.parametrize(
("model", "expected"),
[
("grok/grok-4", "grok-4"),
("Grok/Grok-4", "Grok-4"),
(" grok/grok-4 ", "grok-4"),
("xai/grok-4", None), # metered API path
("chatgpt/gpt-5.4", None),
("grok-4", None),
("grok/", None),
("", None),
(None, None),
],
)
def test_subscription_model(model: str | None, expected: str | None) -> None:
assert grok.subscription_model(model) == expected
def test_auth_mode() -> None:
assert grok.auth_mode("grok/grok-4") == "subscription"
assert grok.auth_mode("xai/grok-4") == "api_key"
assert grok.auth_mode("chatgpt/gpt-5.4") == "api_key"
assert grok.auth_mode(None) == "api_key"
def _record(access: str, refresh: str, expires_at: float) -> dict[str, Any]:
return {
"type": "oauth",
"provider": "grok",
"access": access,
"refresh": refresh,
"expires_at": expires_at,
}
def test_store_roundtrip_and_logout() -> None:
assert grok.read_record() is None
assert grok.is_authenticated() is False
grok.save_record(_record("a1", "r1", time.time() + 3600))
record = grok.read_record()
assert record is not None
assert record["access"] == "a1"
assert grok.is_authenticated() is True
grok.logout()
assert grok.read_record() is None
grok.logout() # no-op when already gone
def test_store_file_permissions_are_owner_only(_tmp_store: Path) -> None:
grok.save_record(_record("a1", "r1", time.time() + 3600))
assert (_tmp_store.stat().st_mode & 0o777) == 0o600
def test_store_shares_file_with_other_providers(_tmp_store: Path) -> None:
# Grok must not clobber a co-resident ChatGPT record in the shared store.
_tmp_store.parent.mkdir(parents=True, exist_ok=True)
_tmp_store.write_text(json.dumps({"codex": {"type": "oauth", "access": "x"}}))
grok.save_record(_record("a1", "r1", time.time() + 3600))
on_disk = json.loads(_tmp_store.read_text())
assert on_disk["codex"] == {"type": "oauth", "access": "x"}
assert on_disk["grok"]["access"] == "a1"
grok.logout()
# Removing grok leaves the other provider's record and the file intact.
assert json.loads(_tmp_store.read_text()) == {"codex": {"type": "oauth", "access": "x"}}
def test_read_record_rejects_incomplete_records() -> None:
grok.save_record({"type": "oauth", "access": "a"}) # missing refresh
assert grok.read_record() is None
assert grok.is_authenticated() is False
def test_get_valid_token_returns_stored_when_fresh(monkeypatch: pytest.MonkeyPatch) -> None:
def _boom(_payload: dict[str, str]) -> dict[str, Any]:
msg = "should not refresh a fresh token"
raise AssertionError(msg)
monkeypatch.setattr(grok, "_post_form", _boom)
grok.save_record(_record("access-fresh", "r1", time.time() + 3600))
assert grok.get_valid_token() == "access-fresh"
def test_get_valid_token_refreshes_and_persists_rotation(monkeypatch: pytest.MonkeyPatch) -> None:
calls = {"n": 0}
def _fake_post(payload: dict[str, str]) -> dict[str, Any]:
calls["n"] += 1
assert payload["grant_type"] == "refresh_token"
assert payload["refresh_token"] == "r1"
return {"access_token": "access-new", "refresh_token": "r2", "expires_in": 3600}
monkeypatch.setattr(grok, "_post_form", _fake_post)
grok.save_record(_record("stale", "r1", time.time() - 10)) # already expired
assert grok.get_valid_token() == "access-new"
assert calls["n"] == 1
record = grok.read_record()
assert record is not None
assert record["refresh"] == "r2" # rotated refresh written back
def test_refresh_keeps_old_refresh_when_response_omits_it(monkeypatch: pytest.MonkeyPatch) -> None:
def _fake_post(_payload: dict[str, str]) -> dict[str, Any]:
return {"access_token": "access-new", "expires_in": 3600} # no refresh_token
monkeypatch.setattr(grok, "_post_form", _fake_post)
grok.save_record(_record("stale", "r1", time.time() - 10))
assert grok.get_valid_token() == "access-new"
record = grok.read_record()
assert record is not None
assert record["refresh"] == "r1" # fell back to the prior refresh token
def test_get_valid_token_uses_token_rotated_by_another_process(
monkeypatch: pytest.MonkeyPatch,
) -> None:
records = [
_record("stale", "r1", time.time() - 10),
_record("fresh-from-other-process", "r2", time.time() + 3600),
]
calls = {"n": 0}
def _fake_read() -> dict[str, Any]:
record = records[min(calls["n"], len(records) - 1)]
calls["n"] += 1
return record
def _boom(_payload: dict[str, str]) -> dict[str, Any]:
msg = "must not refresh a token another process already rotated"
raise AssertionError(msg)
monkeypatch.setattr(grok, "read_record", _fake_read)
monkeypatch.setattr(grok, "_post_form", _boom)
assert grok.get_valid_token() == "fresh-from-other-process"
def test_get_valid_token_recovers_when_refresh_loses_race(
monkeypatch: pytest.MonkeyPatch,
) -> None:
grok.save_record(_record("stale", "r1", time.time() - 10))
def _fake_post(_payload: dict[str, str]) -> dict[str, Any]:
grok.save_record(_record("fresh-from-peer", "r2", time.time() + 3600))
raise grok.GrokAuthError("token_http_error", "HTTP 400: invalid_grant")
monkeypatch.setattr(grok, "_post_form", _fake_post)
assert grok.get_valid_token() == "fresh-from-peer"
def test_get_valid_token_reraises_refresh_error_without_rotation(
monkeypatch: pytest.MonkeyPatch,
) -> None:
grok.save_record(_record("stale", "r1", time.time() - 10))
def _fake_post(_payload: dict[str, str]) -> dict[str, Any]:
raise grok.GrokAuthError("token_http_error", "HTTP 400: invalid_grant")
monkeypatch.setattr(grok, "_post_form", _fake_post)
with pytest.raises(grok.GrokAuthError):
grok.get_valid_token()
def test_get_valid_token_raises_when_not_signed_in() -> None:
with pytest.raises(grok.GrokAuthError) as exc:
grok.get_valid_token()
assert exc.value.code == "not_authenticated"
+113
View File
@@ -0,0 +1,113 @@
"""Grok subscription routing through StrixProvider.get_model."""
from __future__ import annotations
import argparse
from unittest import mock
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
from strix.config import grok, subscription
from strix.config.models import StrixProvider, _TurnGuardModel
from strix.interface import scan_setup, utils
from strix.report import state as state_mod
def test_grok_prefix_routes_to_chat_completions(monkeypatch) -> None: # type: ignore[no-untyped-def]
client = mock.MagicMock()
monkeypatch.setattr(grok, "get_subscription_client", lambda: client)
model = StrixProvider().get_model("grok/grok-4")
assert isinstance(model, _TurnGuardModel)
assert isinstance(model._inner, OpenAIChatCompletionsModel)
# The provider strips the grok/ prefix and passes xAI's bare model slug.
assert model._inner.model == "grok-4"
def test_non_subscription_model_is_not_hijacked_by_grok(monkeypatch) -> None: # type: ignore[no-untyped-def]
def _boom() -> object:
msg = "grok client must not be built for a non-grok model"
raise AssertionError(msg)
monkeypatch.setattr(grok, "get_subscription_client", _boom)
# A metered xai/* key model must fall through to the normal provider path,
# not the subscription route.
model = StrixProvider().get_model("xai/grok-4")
assert isinstance(model, _TurnGuardModel)
assert not isinstance(model._inner, OpenAIChatCompletionsModel)
def test_provider_label_names_the_subscription() -> None:
assert subscription.provider_label("grok/grok-4") == "Grok"
assert subscription.provider_label("chatgpt/gpt-5.4") == "ChatGPT"
# Metered API-key models are not subscriptions.
assert subscription.provider_label("xai/grok-4") is None
assert subscription.provider_label("openai/gpt-5.4") is None
def test_litellm_model_name_maps_subscription_prefixes() -> None:
# Model metadata (context window, output cap) is keyed "xai/…" for Grok and
# bare for ChatGPT; the routing prefixes themselves are never LiteLLM keys.
assert subscription.litellm_model_name("grok/grok-4") == "xai/grok-4"
assert subscription.litellm_model_name("chatgpt/gpt-5.4") == "gpt-5.4"
# Non-subscription models pass through untouched.
assert subscription.litellm_model_name("xai/grok-4") == "xai/grok-4"
assert subscription.litellm_model_name("openai/gpt-5.4") == "openai/gpt-5.4"
assert subscription.litellm_model_name(None) is None
def test_run_record_reports_grok_provider(monkeypatch) -> None: # type: ignore[no-untyped-def]
settings = mock.MagicMock()
settings.llm.model = "grok/grok-4"
monkeypatch.setattr(state_mod, "load_settings", lambda: settings)
record = state_mod.ReportState(run_name="run-test").run_record
assert record["auth_mode"] == "subscription"
assert record["subscription_provider"] == "Grok"
def test_subscription_label_prefers_persisted_provider(monkeypatch) -> None: # type: ignore[no-untyped-def]
settings = mock.MagicMock()
settings.llm.model = "chatgpt/gpt-5.4" # current settings point at ChatGPT
monkeypatch.setattr(utils, "load_settings", lambda: settings)
# A resumed Grok run keeps its persisted provider even though settings changed.
resumed = mock.MagicMock(
run_record={"auth_mode": "subscription", "subscription_provider": "Grok"}
)
assert utils.subscription_label(resumed) == "Grok subscription"
# With no persisted provider, it derives the label from settings (not a
# hardcoded default).
fresh = mock.MagicMock(run_record={})
assert utils.subscription_label(fresh) == "ChatGPT subscription"
def test_persisted_run_record_carries_provider(tmp_path, monkeypatch) -> None: # type: ignore[no-untyped-def]
settings = mock.MagicMock()
settings.llm.model = "grok/grok-4"
monkeypatch.setattr(scan_setup, "load_settings", lambda: settings)
monkeypatch.setattr(scan_setup, "run_dir_for", lambda _name: tmp_path)
captured: dict[str, object] = {}
monkeypatch.setattr(
"strix.report.writer.write_run_record", lambda _dir, rec: captured.update(rec)
)
args = argparse.Namespace(
run_name="run-test",
targets_info=[],
scan_mode="scan",
instruction=None,
non_interactive=True,
local_sources=[],
diff_scope={"active": False},
scope_mode="mode",
diff_base=None,
)
scan_setup._persist_run_record(args)
# The resume/viewer record must carry the provider so resumed runs stay labeled.
assert captured["auth_mode"] == "subscription"
assert captured["subscription_provider"] == "Grok"
+39
View File
@@ -299,6 +299,16 @@ def test_make_model_settings_forces_required_for_anyllm_routed_openai_model() ->
assert settings.tool_choice == "required" assert settings.tool_choice == "required"
def test_make_model_settings_disables_parallel_tool_calls_by_default() -> None:
assert make_model_settings("none", model_name="gpt-4o").parallel_tool_calls is False
def test_make_model_settings_omits_parallel_tool_calls_without_tools() -> None:
settings = make_model_settings("none", model_name="gpt-4o", has_tools=False)
assert settings.parallel_tool_calls is None
def test_make_model_settings_sets_request_timeout() -> None: def test_make_model_settings_sets_request_timeout() -> None:
settings = make_model_settings( settings = make_model_settings(
"none", "none",
@@ -351,3 +361,32 @@ def test_make_model_settings_timeout_survives_reasoning_resolve() -> None:
assert settings.extra_args is not None assert settings.extra_args is not None
assert settings.extra_args["timeout"] == 120.0 assert settings.extra_args["timeout"] == 120.0
def test_openrouter_attribution_rides_on_the_request_headers() -> None:
# litellm.headers is ignored once a request carries any header of its own,
# so the attribution must be part of the per-request headers.
headers = make_model_settings(
None, model_name="openrouter/anthropic/claude-sonnet-4-5"
).extra_headers
assert headers == {
"HTTP-Referer": "https://strix.ai",
"X-Title": "Strix",
"X-OpenRouter-Categories": "cli-agent",
}
def test_openrouter_attribution_absent_for_other_providers() -> None:
assert make_model_settings(None, model_name="anthropic/claude-sonnet-4-5").extra_headers is None
def test_user_headers_override_openrouter_attribution() -> None:
headers = make_model_settings(
None,
model_name="openrouter/anthropic/claude-sonnet-4-5",
extra_headers={"X-Title": "Custom", "X-Tenant": "acme"},
).extra_headers
assert headers is not None
assert headers["X-Title"] == "Custom"
assert headers["X-Tenant"] == "acme"
assert headers["HTTP-Referer"] == "https://strix.ai"
+120
View File
@@ -0,0 +1,120 @@
from __future__ import annotations
from unittest.mock import patch
import litellm
from agents.usage import Usage
from strix.report.pricing import resolve_litellm_model
from strix.report.usage import LLMUsageLedger
def test_resolves_common_bare_model_names() -> None:
resolve_litellm_model.cache_clear()
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"
def test_resolver_returns_none_for_unresolvable_model() -> None:
resolve_litellm_model.cache_clear()
assert resolve_litellm_model("provider/not-a-real-model") is None
def test_ledger_uses_estimate_when_routed_provider_reports_no_cost() -> None:
usage = Usage()
usage.requests = 1
usage.input_tokens = 1000
usage.output_tokens = 200
usage.total_tokens = 1200
ledger = LLMUsageLedger()
with patch("litellm.completion_cost", return_value=0.42):
ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash")
assert ledger.total_cost == 0.42
def test_ledger_prefers_observed_cost_over_estimate() -> None:
usage = Usage()
usage.requests = 1
usage.input_tokens = 1000
usage.output_tokens = 200
usage.total_tokens = 1200
ledger = LLMUsageLedger()
with patch("litellm.completion_cost", return_value=0.42):
ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash")
ledger.record_observed_cost(0.17)
assert ledger.total_cost == 0.17
def test_hydrated_estimate_continues_accumulating_new_estimates() -> None:
usage = Usage()
usage.requests = 1
usage.input_tokens = 1000
usage.output_tokens = 200
usage.total_tokens = 1200
ledger = LLMUsageLedger()
ledger.hydrate({"cost": 0.42})
with patch("litellm.completion_cost", return_value=0.17):
ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash")
assert ledger.total_cost == 0.59
def test_zero_cost_disables_both_observed_and_estimated_costs() -> None:
usage = Usage()
usage.requests = 1
usage.input_tokens = 1000
usage.output_tokens = 200
usage.total_tokens = 1200
ledger = LLMUsageLedger()
ledger.zero_cost = True
with patch("litellm.completion_cost", return_value=0.42) as estimate:
ledger.record(agent_id="a", usage=usage, model="deepseek-v4-flash")
ledger.record_observed_cost(1.0)
estimate.assert_not_called()
assert ledger.total_cost == 0.0
def test_resolver_uses_provider_when_bare_entry_has_one() -> None:
original = litellm.model_cost
litellm.model_cost = {
"example": {
"litellm_provider": "example-provider",
"input_cost_per_token": 1.0,
"output_cost_per_token": 2.0,
}
}
try:
resolve_litellm_model.cache_clear()
assert resolve_litellm_model("example") == "example-provider/example"
finally:
litellm.model_cost = original
resolve_litellm_model.cache_clear()
def test_resolver_does_not_guess_between_differently_priced_providers() -> None:
original = litellm.model_cost
litellm.model_cost = {
"provider-a/example": {
"input_cost_per_token": 1.0,
"output_cost_per_token": 2.0,
},
"provider-b/example": {
"input_cost_per_token": 3.0,
"output_cost_per_token": 4.0,
},
}
try:
resolve_litellm_model.cache_clear()
assert resolve_litellm_model("example") is None
finally:
litellm.model_cost = original
resolve_litellm_model.cache_clear()
+107
View File
@@ -0,0 +1,107 @@
"""Shared subscription credential store: secure writes and cross-provider locking."""
from __future__ import annotations
import fcntl
import stat
from typing import TYPE_CHECKING
import pytest
from strix.config import codex, grok, subscription_store
if TYPE_CHECKING:
from pathlib import Path
def test_write_creates_owner_only_file(tmp_path: Path) -> None:
path = tmp_path / ".strix" / "subscription-auth.json"
subscription_store.write(path, {"grok": {"type": "oauth", "access": "a", "refresh": "r"}})
assert stat.S_IMODE(path.stat().st_mode) == 0o600
# No stray temp file is left behind.
assert not path.with_suffix(".json.tmp").exists()
def test_write_does_not_follow_a_symlink_at_target(tmp_path: Path) -> None:
store_dir = tmp_path / ".strix"
store_dir.mkdir()
outside = tmp_path / "attacker-target.json"
path = store_dir / "subscription-auth.json"
path.symlink_to(outside) # attacker pre-plants a symlink at the store path
subscription_store.write(path, {"grok": {"type": "oauth", "access": "a", "refresh": "r"}})
# The atomic rename replaced the symlink with a real file; nothing was
# written through it to the attacker-chosen location.
assert not path.is_symlink()
assert not outside.exists()
assert subscription_store.read(path)["grok"]["access"] == "a"
def test_providers_share_store_without_clobbering(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
store = tmp_path / ".strix" / "subscription-auth.json"
monkeypatch.setattr(codex, "AUTH_PATH", store)
monkeypatch.setattr(grok, "AUTH_PATH", store)
codex.save_record({"type": "oauth", "access": "c", "refresh": "r", "account_id": "acct"})
grok.save_record({"type": "oauth", "access": "g", "refresh": "r"})
data = subscription_store.read(store)
assert data["codex"]["access"] == "c"
assert data["grok"]["access"] == "g"
# Logging one provider out leaves the other's credential intact.
grok.logout()
remaining = subscription_store.read(store)
assert "grok" not in remaining
assert remaining["codex"]["access"] == "c"
def test_guard_is_reentrant(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
store = tmp_path / ".strix" / "subscription-auth.json"
monkeypatch.setattr(grok, "AUTH_PATH", store)
# Persisting while already holding the guard must not deadlock — this mirrors
# a token refresh saving its new record inside the refresh critical section.
with subscription_store.guard(store):
grok.save_record({"type": "oauth", "access": "g", "refresh": "r"})
record = grok.read_record()
assert record is not None
assert record["access"] == "g"
def test_mutation_aborts_when_lock_cannot_be_acquired(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
store = tmp_path / ".strix" / "subscription-auth.json"
monkeypatch.setattr(grok, "AUTH_PATH", store)
def _no_lock(*_args: object, **_kwargs: object) -> None:
raise OSError("no locks available")
monkeypatch.setattr(fcntl, "flock", _no_lock)
# Rather than silently doing an unlocked read-modify-write, the store raises
# and writes nothing.
with pytest.raises(subscription_store.StoreLockError):
grok.save_record({"type": "oauth", "access": "g", "refresh": "r"})
assert not store.exists()
def test_lock_file_rejects_a_pre_positioned_symlink(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
store_dir = tmp_path / ".strix"
store_dir.mkdir()
store = store_dir / "subscription-auth.json"
monkeypatch.setattr(grok, "AUTH_PATH", store)
# Attacker pre-plants a symlink where the lock file would be created.
outside = tmp_path / "attacker-target"
store.with_suffix(".lock").symlink_to(outside)
with pytest.raises(subscription_store.StoreLockError):
grok.save_record({"type": "oauth", "access": "g", "refresh": "r"})
# The symlink target was never created/truncated through the lock open.
assert not outside.exists()
+22
View File
@@ -110,6 +110,28 @@ def test_state_populates_model_warning_for_non_frontier_model() -> None:
assert "not a recommended frontier model" in warning assert "not a recommended frontier model" in warning
def test_snapshot_carries_the_subscription_provider_label() -> None:
os.environ["STRIX_LLM"] = "grok/grok-4"
loader._cached = None
snapshot = TuiController(args()).snapshot()
# The TUI renders this label, so it must name the actual provider rather
# than assuming ChatGPT.
assert snapshot["subscription"] is True
assert snapshot["subscription_label"] == "Grok subscription"
def test_snapshot_has_no_subscription_label_for_api_key_runs() -> None:
os.environ["STRIX_LLM"] = "openai/gpt-5.4"
loader._cached = None
snapshot = TuiController(args()).snapshot()
assert snapshot["subscription"] is False
assert snapshot["subscription_label"] == ""
def test_setup_restores_prepared_cli_targets() -> None: def test_setup_restores_prepared_cli_targets() -> None:
setup_args = args() setup_args = args()
setup_args.targets_info = [ setup_args.targets_info = [
+47
View File
@@ -70,6 +70,53 @@ def test_read_run_summary_finished_flag(tmp_path: Path) -> None:
assert read_run_summary(partial)["finished"] is False assert read_run_summary(partial)["finished"] is False
def _write_record(base: Path, name: str, record: dict[str, object]) -> Path:
run_dir = base / "strix_runs" / name
run_dir.mkdir(parents=True)
(run_dir / "run.json").write_text(json.dumps(record), encoding="utf-8")
return run_dir
def test_read_run_summary_backfills_subscription_provider(tmp_path: Path) -> None:
# An older subscription run recorded no provider name; it is derived from
# the recorded provider/model slug so the viewer can label it.
run_dir = _write_record(
tmp_path,
"grok-run",
{
"auth_mode": "subscription",
"llm_usage": {"agents": [{"agent_id": "root", "model": "grok/grok-4"}]},
},
)
assert read_run_summary(run_dir)["subscription_provider"] == "Grok"
def test_read_run_summary_keeps_explicit_provider(tmp_path: Path) -> None:
run_dir = _write_record(
tmp_path,
"chatgpt-run",
{
"auth_mode": "subscription",
"subscription_provider": "ChatGPT",
"llm_usage": {"agents": [{"agent_id": "root", "model": "grok/grok-4"}]},
},
)
# An explicit field is authoritative and never overwritten by the slug.
assert read_run_summary(run_dir)["subscription_provider"] == "ChatGPT"
def test_read_run_summary_ignores_api_key_runs(tmp_path: Path) -> None:
run_dir = _write_record(
tmp_path,
"api-key-run",
{
"auth_mode": "api_key",
"llm_usage": {"agents": [{"agent_id": "root", "model": "openai/gpt-5.4"}]},
},
)
assert "subscription_provider" not in read_run_summary(run_dir)
def test_read_missing_artifacts_return_defaults(tmp_path: Path) -> None: def test_read_missing_artifacts_return_defaults(tmp_path: Path) -> None:
run_dir = _make_run(tmp_path, "empty", status="running", end_time=None) run_dir = _make_run(tmp_path, "empty", status="running", end_time=None)
assert read_vulnerabilities(run_dir) == [] assert read_vulnerabilities(run_dir) == []
Generated
+1 -1
View File
@@ -2378,7 +2378,7 @@ wheels = [
[[package]] [[package]]
name = "strix-agent" name = "strix-agent"
version = "1.5.1" version = "1.5.3"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "caido-sdk-client" }, { name = "caido-sdk-client" },