Sign in with a ChatGPT subscription for inference (#854)

Co-authored-by: Jonathan Singer <jonathansinger@Jonathans-MacBook-Pro.local>
Co-authored-by: Jonathan Singer <jonathansinger@Mac-3004.lan>
Co-authored-by: Ahmed Allam <ahmed39652003@gmail.com>
This commit is contained in:
yoni-at-strix
2026-07-24 15:41:19 -07:00
committed by GitHub
co-authored by Jonathan Singer Jonathan Singer Ahmed Allam
parent 93af2b94a2
commit cd8270c98b
30 changed files with 1899 additions and 93 deletions
+14
View File
@@ -267,6 +267,20 @@ export STRIX_REASONING_EFFORT="high" # control thinking effort (default: high,
> [!NOTE]
> Strix automatically saves your configuration to `~/.strix/cli-config.json`, so you don't have to re-enter it on every run.
#### Sign in with a ChatGPT subscription
Instead of a metered API key, you can run Strix on your ChatGPT Plus/Pro subscription:
```bash
strix auth login chatgpt # sign in with your ChatGPT account
export STRIX_LLM="chatgpt/gpt-5.4" # chatgpt/<model> runs on the subscription
strix --target ./app-directory
strix auth status # show the active sign-in
strix auth logout # forget the sign-in
```
**Recommended models for best results:**
- [OpenAI GPT-5.4](https://openai.com/api/) - `openai/gpt-5.4`
+9 -1
View File
@@ -215,6 +215,10 @@ ignore = [
# Test doubles use fixture tokens/passwords and match a callee signature whose
# args they intentionally ignore.
"tests/test_viewer_auth.py" = ["S105", "S106", "ARG001"]
"tests/test_codex_auth.py" = ["S105", "S106", "SLF001"]
# Stdlib HTTP handler overrides (do_GET/do_POST).
"strix/interface/auth_cli.py" = ["N802"]
"tests/test_codex_streaming.py" = ["N802"]
"tests/test_report_pdf.py" = ["S105", "S106"]
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
# circular dependency with strix.telemetry / strix.viewer.report_pdf.
@@ -253,9 +257,13 @@ ignore = [
"strix/core/runner.py" = ["TC003", "PLR0912", "PLR0915", "PLC0415"]
# ReportState carries scan artifact/report fields and
# a runtime ``Callable`` annotation on ``vulnerability_found_callback``.
"strix/report/state.py" = ["TC003", "PLR0912", "PLR0915", "E501", "PERF401"]
"strix/report/state.py" = ["TC003", "PLR0912", "PLR0915", "E501", "PERF401", "PLC0415"]
"strix/report/usage.py" = ["PLC0415"]
"strix/telemetry/logging.py" = ["PLC0415"]
"strix/config/models.py" = ["PLC0415"]
# Heavy inference deps (httpx, openai) imported lazily so auth-status checks
# don't pull them in.
"strix/config/codex.py" = ["PLC0415"]
# Interface utility branches per scope-mode / target-type combination;
# splitting would obscure the decision tree without simplifying it.
"strix/interface/utils.py" = ["PLR0912", "BLE001", "PLC0415"]
+411
View File
@@ -0,0 +1,411 @@
"""ChatGPT (Codex) subscription auth: OAuth login, token refresh, and the OpenAI
client that routes inference through the ChatGPT backend.
Mirrors OpenAI's Codex CLI: OAuth 2.0 + PKCE against ``auth.openai.com``, with the
access token sent as a ``Bearer`` token to ``chatgpt.com/backend-api/codex``. Using
a ChatGPT subscription outside OpenAI's own products is not officially supported by
OpenAI; the user chooses this path knowingly. The OAuth constants are OpenAI's own
Codex 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 threading
import time
import urllib.error
import urllib.parse
import urllib.request
from pathlib import Path
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from collections.abc import Iterator
from openai import AsyncOpenAI
logger = logging.getLogger(__name__)
PROVIDER = "codex"
CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann"
AUTHORIZE_URL = "https://auth.openai.com/oauth/authorize"
TOKEN_URL = "https://auth.openai.com/oauth/token" # noqa: S105 # nosec B105 - URL, not a secret
CALLBACK_HOST = "localhost"
CALLBACK_PORT = 1455
CALLBACK_PATH = "/auth/callback"
REDIRECT_URI = f"http://{CALLBACK_HOST}:{CALLBACK_PORT}{CALLBACK_PATH}"
SCOPE = "openid profile email offline_access"
CODEX_BASE_URL = "https://chatgpt.com/backend-api/codex"
ORIGINATOR = "codex_cli_rs"
_ACCOUNT_CLAIM = "https://api.openai.com/auth"
_TOKEN_TIMEOUT = 30
_EXPIRY_SKEW_S = 300
_refresh_lock = threading.Lock()
# 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_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:
AUTH_PATH.parent.mkdir(parents=True, exist_ok=True)
tmp = AUTH_PATH.with_suffix(".json.tmp")
tmp.write_text(json.dumps(data, indent=2), encoding="utf-8")
with contextlib.suppress(OSError):
tmp.chmod(0o600)
tmp.replace(AUTH_PATH)
with contextlib.suppress(OSError):
AUTH_PATH.chmod(0o600)
def read_record() -> dict[str, Any] | None:
record = _read_store().get(PROVIDER)
if not isinstance(record, dict) or record.get("type") != "oauth":
return None
if not (record.get("access") and record.get("refresh") and record.get("account_id")):
return None
return record
def is_authenticated() -> bool:
return read_record() is not None
def save_record(record: dict[str, Any]) -> None:
data = _read_store()
data[PROVIDER] = record
_write_store(data)
def logout() -> None:
data = _read_store()
if PROVIDER not in data:
return
del data[PROVIDER]
if data:
_write_store(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 _refresh_lock:
try:
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):
def __init__(self, code: str, message: str | None = None) -> None:
self.code = code
super().__init__(message or code)
class CodexContentGuardrailError(Exception):
"""The ChatGPT backend refused a request via its content guardrail.
Terminal — retrying identical content never clears the block."""
def __init__(self, model: str, original: BaseException | None = None) -> None:
self.model = model
self.original = original
super().__init__(
f"'{model}' was blocked by ChatGPT's content guardrails "
f"(flagged as a possible cybersecurity risk). "
f"Set STRIX_LLM to a model that isn't blocked and re-run."
)
_GUARDRAIL_MARKERS = (
"flagged for possible cybersecurity risk",
"trusted access for cyber",
)
def is_content_guardrail_error(exc: BaseException) -> bool:
if isinstance(exc, CodexContentGuardrailError):
return True
text = str(exc).lower()
return any(marker in text for marker in _GUARDRAIL_MARKERS)
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,
"id_token_add_organizations": "true",
"codex_cli_simplified_flow": "true",
"originator": ORIGINATOR,
}
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]:
body = urllib.parse.urlencode(payload).encode("ascii")
request = urllib.request.Request( # noqa: S310 - fixed https OAuth endpoint
TOKEN_URL,
data=body,
headers={
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
},
method="POST",
)
try:
with urllib.request.urlopen( # noqa: S310 # nosec B310 - fixed https endpoint
request, timeout=_TOKEN_TIMEOUT
) as response:
data = json.loads(response.read() or b"{}")
except urllib.error.HTTPError as exc:
detail = exc.read().decode("utf-8", "replace")[:300]
raise CodexAuthError("token_http_error", f"HTTP {exc.code}: {detail}") from exc
except (urllib.error.URLError, TimeoutError, OSError) as exc:
raise CodexAuthError("unavailable", str(exc)) from exc
if not isinstance(data, dict):
raise CodexAuthError("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 CodexAuthError("bad_response", "token response missing access_token")
if not isinstance(refresh, str) or not refresh:
raise CodexAuthError("bad_response", "token response missing refresh_token")
account_id = _account_id_from_jwt(access) or _account_id_from_jwt(
data.get("id_token") if isinstance(data.get("id_token"), str) else ""
)
if not account_id:
raise CodexAuthError("no_account_id", "could not read chatgpt_account_id from token")
ttl = expires_in if isinstance(expires_in, int | float) else 3600
return {
"type": "oauth",
"provider": PROVIDER,
"access": access,
"refresh": refresh,
"account_id": account_id,
"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 _account_id_from_jwt(token: str | None) -> str | None:
"""Read the account id claim without verifying the JWT (the server enforces
authenticity on use); it feeds the ``chatgpt-account-id`` header."""
if not token or token.count(".") != 2:
return None
payload_b64 = token.split(".")[1]
padding = "=" * (-len(payload_b64) % 4)
try:
payload = json.loads(base64.urlsafe_b64decode(payload_b64 + padding))
except (ValueError, json.JSONDecodeError):
return None
if not isinstance(payload, dict):
return None
auth = payload.get(_ACCOUNT_CLAIM)
if isinstance(auth, dict):
account_id = auth.get("chatgpt_account_id")
if isinstance(account_id, str) and account_id:
return account_id
organizations = payload.get("organizations")
if isinstance(organizations, list) and organizations and isinstance(organizations[0], dict):
org_id = organizations[0].get("id")
if isinstance(org_id, str) and org_id:
return org_id
return None
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() -> tuple[str, str]:
"""Return ``(access_token, account_id)``, refreshing under the cross-process
guard if near expiry."""
record = read_record()
if record is None:
raise CodexAuthError("not_authenticated", "not signed in; run: strix auth login")
if not _near_expiry(record):
return record["access"], record["account_id"]
with _refresh_guard():
record = read_record()
if record is None:
raise CodexAuthError("not_authenticated", "not signed in; run: strix auth login")
if not _near_expiry(record):
return record["access"], record["account_id"]
try:
refreshed = refresh_tokens(record["refresh"])
except CodexAuthError:
# 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 latest["access"], latest["account_id"]
raise
save_record(refreshed)
return refreshed["access"], refreshed["account_id"]
def build_openai_client() -> AsyncOpenAI:
"""An ``AsyncOpenAI`` for the ChatGPT backend. 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, account_id = await asyncio.to_thread(get_valid_token)
request.headers["Authorization"] = f"Bearer {access}"
request.headers["chatgpt-account-id"] = account_id
http_client = httpx.AsyncClient(
timeout=httpx.Timeout(600.0, connect=30.0),
event_hooks={"request": [_auth_hook]},
)
return AsyncOpenAI(
api_key="strix-codex-oauth", # placeholder; the hook overwrites Authorization
base_url=CODEX_BASE_URL,
http_client=http_client,
default_headers={
"OpenAI-Beta": "responses=experimental",
"originator": ORIGINATOR,
},
)
_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 = "chatgpt/"
def subscription_model(model_name: str | None) -> str | None:
"""The model slug behind a ``chatgpt/<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"
+128 -12
View File
@@ -2,23 +2,38 @@
from __future__ import annotations
import contextlib
import inspect
import os
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
from agents import set_default_openai_api, set_default_openai_key, set_tracing_disabled
from agents import (
set_default_openai_api,
set_default_openai_key,
set_tracing_disabled,
)
from agents.model_settings import ModelSettings
from agents.models.multi_provider import MultiProvider
from agents.models.openai_responses import OpenAIResponsesModel
from agents.retry import (
ModelRetryBackoffSettings,
ModelRetrySettings,
RetryPolicyContext,
retry_policies,
)
from openai.types.shared import Reasoning
from strix.config import codex
from strix.config.loader import load_settings
if TYPE_CHECKING:
from agents.models.interface import ModelProvider
from collections.abc import AsyncIterator
from strix.config.settings import Settings
from agents.models.interface import Model, ModelProvider
from openai import AsyncOpenAI
from strix.config.settings import ReasoningEffort, Settings
def request_timeout_extra_args(timeout_s: float | None) -> dict[str, float] | None:
@@ -33,9 +48,93 @@ def _retry_statusless_provider_errors(context: RetryPolicyContext) -> bool:
normalized = context.normalized
if normalized.is_abort:
return False
if codex.is_content_guardrail_error(context.error):
return False
return normalized.status_code is None
class _CodexResponsesModel(OpenAIResponsesModel):
"""Responses model for the ChatGPT subscription backend (always streamed, stateless)."""
def __init__(
self,
model: str,
openai_client: AsyncOpenAI,
*,
reasoning_effort: ReasoningEffort | None = None,
) -> None:
super().__init__(model, openai_client)
self._reasoning_effort = reasoning_effort
def _codex_settings(self, model_settings: ModelSettings) -> ModelSettings:
overrides = ModelSettings(store=False, response_include=["reasoning.encrypted_content"])
effort = self._reasoning_effort
if effort and effort != "none":
# Clamp to efforts the backend accepts.
if effort == "minimal":
effort = "low"
elif effort == "xhigh":
effort = "high"
overrides = overrides.resolve(ModelSettings(reasoning=Reasoning(effort=effort)))
return model_settings.resolve(overrides)
async def _fetch_response(self, *args: Any, stream: bool = False, **kwargs: Any) -> Any:
if len(args) >= 3: # model_settings is positional arg 2
args = (*args[:2], self._codex_settings(args[2]), *args[3:])
try:
events = await super()._fetch_response(*args, stream=True, **kwargs) # type: ignore[call-overload]
except Exception as exc:
guardrail = self._as_guardrail(exc)
if guardrail is not None:
raise guardrail from exc
raise
guarded = self._guarded(events)
if stream:
return guarded
final_response = None
async for event in guarded:
if getattr(event, "type", None) == "response.completed":
final_response = event.response
if final_response is None:
msg = "ChatGPT backend stream ended without a completed response"
raise RuntimeError(msg)
return final_response
def _as_guardrail(self, exc: BaseException) -> codex.CodexContentGuardrailError | None:
if isinstance(exc, codex.CodexContentGuardrailError):
return exc
if codex.is_content_guardrail_error(exc):
return codex.CodexContentGuardrailError(self.model, exc)
return None
async def _guarded(self, events: Any) -> AsyncIterator[Any]:
"""Convert mid-stream guardrail rejections and close the stream on exit."""
try:
async for event in events:
yield event
except Exception as exc:
guardrail = self._as_guardrail(exc)
if guardrail is not None:
raise guardrail from exc
raise
finally:
await self._aclose(events)
@staticmethod
async def _aclose(events: Any) -> None:
aclose = getattr(events, "aclose", None)
if callable(aclose):
with contextlib.suppress(Exception):
await aclose()
return
close = getattr(events, "close", None)
if callable(close):
with contextlib.suppress(Exception):
result = close()
if inspect.isawaitable(result):
await result
class StrixProvider(MultiProvider):
"""Route any non-OpenAI prefix through LiteLLM with the prefix preserved,
so users type ``deepseek/deepseek-chat`` rather than
@@ -59,6 +158,16 @@ class StrixProvider(MultiProvider):
return self._get_fallback_provider("litellm"), f"ollama_chat/{stripped_model_name}"
return self._get_fallback_provider("litellm"), original_model_name
def get_model(self, model_name: str | None) -> Model:
slug = codex.subscription_model(model_name)
if slug:
return _CodexResponsesModel(
slug,
codex.get_subscription_client(),
reasoning_effort=load_settings().llm.reasoning_effort,
)
return super().get_model(model_name)
DEFAULT_MODEL_RETRY = ModelRetrySettings(
max_retries=5,
@@ -77,39 +186,42 @@ DEFAULT_MODEL_RETRY = ModelRetrySettings(
)
RECOMMENDED_MODEL_NAMES = (
"openai/gpt-5.6",
"openai/gpt-5.6-sol",
"openai/gpt-5.6-terra",
"openai/gpt-5.5",
"openai/gpt-5.6-luna",
"openai/gpt-5.6",
"openai/gpt-5.5-pro",
"openai/gpt-5.5",
"openai/gpt-5.4",
"openai/gpt-5.3-codex",
"anthropic/claude-fable-5",
"anthropic/claude-opus-5",
"anthropic/claude-opus-4-8",
"anthropic/claude-opus-4-7",
"anthropic/claude-sonnet-5",
"anthropic/claude-sonnet-4-6",
"vertex_ai/gemini-3.1-pro-preview",
"gemini/gemini-3.1-pro-preview",
"gemini/gemini-3.6-flash",
"deepseek/deepseek-v4-pro",
"deepseek/deepseek-v4-flash",
"dashscope/qwen3.8-max",
"dashscope/qwen3.7-max-2026-06-08",
"moonshot/kimi-k3",
"moonshot/kimi-k2.7-code",
"moonshot/kimi-k2.6",
)
_RECOMMENDED_MODEL_NAME_SET = frozenset(name.lower() for name in RECOMMENDED_MODEL_NAMES)
FRONTIER_MODEL_FAMILIES = (
(("azure", "azure_ai", "bedrock_mantle", "openai"), ("gpt-5",)),
(("azure", "azure_ai", "bedrock_mantle", "chatgpt", "openai"), ("gpt-5",)),
(
("anthropic", "azure_ai", "bedrock", "claude", "databricks", "snowflake", "vertex_ai"),
("claude-fable-5", "claude-opus-4", "claude-sonnet-5", "claude-sonnet-4"),
("claude-fable-5", "claude-opus-5", "claude-opus-4", "claude-sonnet-5", "claude-sonnet-4"),
),
(("google", "gemini", "vertex_ai"), ("gemini-3",)),
(("deepseek",), ("deepseek-v4", "deepseek-r1", "deepseek-reasoner")),
(("alibaba", "dashscope", "qwen"), ("qwen3.7", "qwen3.5", "qwen3-max")),
(("moonshot", "moonshotai", "kimi"), ("kimi-k2.7", "kimi-k2.6", "kimi-k2.5")),
(("alibaba", "dashscope", "qwen"), ("qwen3.8", "qwen3.7", "qwen3-max")),
(("moonshot", "moonshotai", "kimi"), ("kimi-k3", "kimi-k2.7", "kimi-k2.6")),
)
@@ -117,6 +229,8 @@ def configure_sdk_model_defaults(settings: Settings) -> None:
"""Apply Strix config to SDK-native defaults."""
llm = settings.llm
set_tracing_disabled(True)
if codex.subscription_model(llm.model):
return
_configure_litellm_compatibility()
_configure_openrouter_attribution(llm.model)
if llm.api_key:
@@ -211,6 +325,8 @@ def _configure_litellm_default(name: str, value: str) -> None:
def uses_chat_completions_tool_schema(model_name: str, settings: Settings) -> bool:
"""Return whether the resolved SDK route can only receive JSON function tools."""
if codex.subscription_model(model_name):
return False
model = model_name.strip().lower()
if "/" in model and not model.startswith("openai/"):
return True
+18 -3
View File
@@ -41,6 +41,7 @@ class AgentCoordinator:
self.names: dict[str, str] = {}
self.metadata: dict[str, dict[str, Any]] = {}
self.pending_counts: dict[str, int] = {}
self.errors: dict[str, str] = {}
self.runtimes: dict[str, AgentRuntime] = {}
self._lock = asyncio.Lock()
self._snapshot_path: Path | None = None
@@ -107,16 +108,23 @@ class AgentCoordinator:
async with self._lock:
if agent_id in self.statuses:
self.statuses[agent_id] = "running"
self.errors.pop(agent_id, None)
await self._maybe_snapshot()
async def park_waiting(self, agent_id: str) -> None:
await self.set_status(agent_id, "waiting")
async def set_status(self, agent_id: str, status: Status | str) -> None:
async def set_status(
self, agent_id: str, status: Status | str, *, error: str | None = None
) -> None:
async with self._lock:
if agent_id not in self.statuses:
return
self.statuses[agent_id] = status # type: ignore[assignment]
if error is not None:
self.errors[agent_id] = error
elif status == "running":
self.errors.pop(agent_id, None)
runtime = self.runtimes.setdefault(agent_id, AgentRuntime())
runtime.wake.set()
logger.info("agent.status %s=%s", agent_id, status)
@@ -246,9 +254,14 @@ class AgentCoordinator:
async def graph_snapshot(
self,
) -> tuple[dict[str, str | None], dict[str, Status], dict[str, str]]:
) -> tuple[dict[str, str | None], dict[str, Status], dict[str, str], dict[str, str]]:
async with self._lock:
return dict(self.parent_of), dict(self.statuses), dict(self.names)
return (
dict(self.parent_of),
dict(self.statuses),
dict(self.names),
dict(self.errors),
)
def _message_to_session_item(self, message: dict[str, Any]) -> TResponseInputItem:
sender = str(message.get("from", "unknown"))
@@ -286,6 +299,7 @@ class AgentCoordinator:
"names": dict(self.names),
"metadata": {aid: dict(md) for aid, md in self.metadata.items()},
"pending_counts": dict(self.pending_counts),
"errors": dict(self.errors),
}
async def restore(self, snap: dict[str, Any]) -> None:
@@ -295,6 +309,7 @@ class AgentCoordinator:
self.names = dict(snap.get("names", {}))
self.metadata = {aid: dict(md) for aid, md in snap.get("metadata", {}).items()}
self.pending_counts = dict(snap.get("pending_counts", {}))
self.errors = dict(snap.get("errors", {}))
for aid in self.statuses:
self.runtimes.setdefault(aid, AgentRuntime())
+1 -3
View File
@@ -437,10 +437,8 @@ async def _run_cycle( # noqa: PLR0912, PLR0915
else:
status = "crashed"
logger.exception("agent run failed for %s; parking as %s", agent_id, status)
await coordinator.set_status(agent_id, status)
await coordinator.set_status(agent_id, status, error=str(exc) or type(exc).__name__)
await _notify_parent_on_crash(coordinator, agent_id, status)
if context.get("parent_id") is None and status in {"failed", "crashed"}:
raise
return None
else:
await _settle_run_result(coordinator, agent_id, interactive)
+419
View File
@@ -0,0 +1,419 @@
"""`strix auth` — ChatGPT subscription sign-in (login / status / logout).
Signing in only stores OAuth tokens (``~/.strix/subscription-auth.json``); model
selection stays with ``STRIX_LLM``. A ``chatgpt/<model>`` STRIX_LLM runs on the
subscription.
"""
from __future__ import annotations
import argparse
import base64
import logging
import threading
import webbrowser
from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path
from typing import TYPE_CHECKING, Any
from urllib.parse import parse_qs, urlparse
from rich.console import Console
from rich.panel import Panel
from rich.text import Text
from strix.config import codex, load_settings
if TYPE_CHECKING:
from collections.abc import Callable
logger = logging.getLogger(__name__)
_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"
def run_auth(argv: list[str]) -> int:
"""Entry point for ``strix auth …``. Returns a process exit code."""
console = Console()
# Bare `strix auth` (no subcommand) defaults to login.
subcommand = argv[0] if argv else "login"
rest = argv[1:]
if subcommand in ("-h", "--help", "help"):
console.print(_USAGE)
return 0
handlers: dict[str, Callable[[], int]] = {
"login": lambda: _login(console, rest),
"status": lambda: _status(console),
"logout": lambda: _logout(console),
}
handler = handlers.get(subcommand)
if handler is not None:
return handler()
console.print(f"[red]Unknown auth command:[/] {subcommand}\n")
console.print(_USAGE)
return 2
def _login(console: Console, argv: list[str]) -> int:
parser = argparse.ArgumentParser(prog="strix auth login", add_help=True)
parser.add_argument(
"provider",
nargs="?",
default=LOGIN_PROVIDER,
help="Model provider to sign in with (default: chatgpt).",
)
parser.add_argument(
"--manual",
action="store_true",
help="Skip the local callback server and paste the redirect URL by hand.",
)
try:
args = parser.parse_args(argv)
except SystemExit as exc: # argparse already printed the message
return int(exc.code or 2)
if args.provider.lower() not in _ACCEPTED_PROVIDERS:
console.print(
f"[red]Unsupported provider:[/] {args.provider}. "
f"Only '{LOGIN_PROVIDER}' (ChatGPT subscription) is supported."
)
return 2
verifier, challenge = codex.generate_pkce()
state = codex.create_state()
authorize_url = codex.build_authorize_url(challenge, state)
console.print()
console.print("[bold]Signing in with ChatGPT[/] [dim](provider: chatgpt)[/]")
console.print(
"[dim]This uses your ChatGPT Plus/Pro plan for inference instead of a metered API key.[/]"
)
console.print()
try:
record = _run_oauth_flow(console, authorize_url, verifier, state, manual=args.manual)
except codex.CodexAuthError as exc:
return _fail(console, exc)
except KeyboardInterrupt:
console.print("\n[yellow]Sign-in cancelled.[/]")
return 130
codex.save_record(record)
_print_success(console)
return 0
def _run_oauth_flow(
console: Console,
authorize_url: str,
verifier: str,
state: str,
*,
manual: bool,
) -> dict[str, Any]:
"""Drive the browser (or manual) OAuth flow and return a token record."""
server = None if manual else _try_start_callback_server()
console.print("Open this URL in your browser to authorize:")
console.print(f"[cyan]{authorize_url}[/]")
console.print()
if not manual:
try:
webbrowser.open(authorize_url)
except Exception: # noqa: BLE001 - opening a browser is best-effort
logger.debug("could not open browser", exc_info=True)
if server is not None:
console.print("[dim]Waiting for you to finish signing in…[/]")
result = server.wait(_CALLBACK_TIMEOUT_S)
server.shutdown()
if result is not None:
code, returned_state, error = result
if error:
raise codex.CodexAuthError("oauth_error", error)
return _finish(code, returned_state, verifier, state, require_state=True)
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
# (the browser lands on a localhost page that won't load if no server is up;
# the address bar still holds the code+state).
console.print()
try:
pasted = console.input("Paste the full redirect URL (or code#state): ").strip()
except EOFError as exc:
raise codex.CodexAuthError("no_input", "no redirect URL provided") from exc
code, returned_state = codex.parse_redirect_input(pasted)
return _finish(code, returned_state, verifier, state, require_state=False)
def _finish(
code: str | None,
returned_state: str | None,
verifier: str,
expected_state: str,
*,
require_state: bool,
) -> dict[str, Any]:
if not code:
raise codex.CodexAuthError("no_code", "no authorization code found in the redirect")
# The loopback callback from OpenAI always carries state, so a missing or
# mismatched value there is forged (CSRF) and must be rejected. Manual paste
# is user-initiated (the user copies their own redirect), so state is only
# validated when the pasted value includes it.
if require_state and returned_state is None:
raise codex.CodexAuthError("state_mismatch", "missing state in callback; possible CSRF")
if returned_state is not None and returned_state != expected_state:
raise codex.CodexAuthError("state_mismatch", "state did not match; possible CSRF")
return codex.exchange_code(code, verifier)
class _CallbackServer:
"""A one-shot local HTTP server that catches the OAuth redirect."""
def __init__(self, httpd: HTTPServer, event: threading.Event, holder: dict[str, Any]) -> None:
self._httpd = httpd
self._event = event
self._holder = holder
self._thread = threading.Thread(target=httpd.serve_forever, daemon=True)
self._thread.start()
def wait(self, timeout: float) -> tuple[str | None, str | None, str | None] | None:
if not self._event.wait(timeout):
return None
return (
self._holder.get("code"),
self._holder.get("state"),
self._holder.get("error"),
)
def shutdown(self) -> None:
self._httpd.shutdown()
self._httpd.server_close()
def _try_start_callback_server() -> _CallbackServer | None:
event = threading.Event()
holder: dict[str, Any] = {}
class Handler(BaseHTTPRequestHandler):
def log_message(self, *args: Any) -> None: # silence default stderr logging
pass
def do_GET(self) -> None:
parsed = urlparse(self.path)
if parsed.path != codex.CALLBACK_PATH:
self.send_response(404)
self.end_headers()
return
query = parse_qs(parsed.query)
holder["code"] = _first(query, "code")
holder["state"] = _first(query, "state")
holder["error"] = _first(query, "error_description") or _first(query, "error")
body = _render_callback_html().encode("utf-8")
self.send_response(200)
self.send_header("Content-Type", "text/html; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
event.set()
try:
httpd = HTTPServer(("127.0.0.1", codex.CALLBACK_PORT), Handler)
except OSError:
logger.debug("could not bind callback port %d", codex.CALLBACK_PORT, exc_info=True)
return None
return _CallbackServer(httpd, event, holder)
def _first(query: dict[str, list[str]], key: str) -> str | None:
values = query.get(key)
return values[0] if values else None
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()
console.print("[green]Signed in[/] with a ChatGPT subscription.")
console.print(f" Account: [bold]{record.get('account_id')}[/]")
if codex.subscription_model(settings.llm.model):
console.print(f" Runs use the subscription (STRIX_LLM=[bold]{settings.llm.model}[/]).")
else:
console.print(
" [yellow]Note:[/] set [cyan]STRIX_LLM[/] to e.g. [cyan]chatgpt/gpt-5.4[/] "
"to run on the subscription."
)
return 0
def _logout(console: Console) -> int:
codex.logout()
console.print("[green]Signed out.[/] Stored subscription credentials removed.")
return 0
def _fail(console: Console, exc: codex.CodexAuthError) -> int:
error_text = Text()
error_text.append("SIGN-IN FAILED", style="bold red")
error_text.append("\n\n", style="white")
error_text.append(f"{exc}", style="white")
console.print()
console.print(
Panel(
error_text,
title="[bold white]STRIX",
title_align="left",
border_style="red",
padding=(1, 2),
)
)
return 1
def _print_success(console: Console) -> None:
text = Text()
text.append("Signed in with your ChatGPT subscription", style="bold #22c55e")
text.append("\n\n", style="white")
text.append("Set ", style="white")
text.append("STRIX_LLM", style="bold white")
text.append(" to a ", style="white")
text.append("chatgpt/", style="bold cyan")
text.append(" model (e.g. ", style="white")
text.append("chatgpt/gpt-5.4", style="bold cyan")
text.append(") — runs are billed to your ChatGPT plan.", style="white")
text.append("\n\n", style="white")
text.append("Run a scan as usual, e.g. ", style="white")
text.append("strix --target https://example.com", style="bold cyan")
console.print()
console.print(
Panel(
text,
title="[bold white]STRIX",
title_align="left",
border_style="#22c55e",
padding=(1, 2),
)
)
console.print()
_LOGO_PATH = Path(__file__).resolve().parent.parent / "viewer" / "static" / "logo.png"
def _logo_img_tag() -> str:
"""Return an ``<img>`` for the Strix logo as an inline data URI, or "".
The callback page is served offline by the local OAuth server, so the logo
is embedded rather than linked. Missing/unreadable file degrades to just the
"Strix" wordmark.
"""
try:
data = _LOGO_PATH.read_bytes()
except OSError:
return ""
encoded = base64.b64encode(data).decode("ascii")
return f'<img class="logo" src="data:image/png;base64,{encoded}" alt="" />'
def _render_callback_html() -> str:
return _CALLBACK_HTML.replace("<!--LOGO-->", _logo_img_tag())
_CALLBACK_HTML = """<!doctype html>
<html lang="en"><head><meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>Strix — signed in</title>
<style>
:root { color-scheme: dark; }
* { box-sizing: border-box; }
body {
margin: 0; min-height: 100vh; padding: 24px;
font-family: 'Geist', 'Geist Sans', ui-sans-serif, system-ui, -apple-system,
"Segoe UI", Roboto, Helvetica, Arial, sans-serif;
-webkit-font-smoothing: antialiased; -moz-osx-font-smoothing: grayscale;
background: #000; color: #ededed;
display: flex; flex-direction: column; align-items: center; justify-content: center;
}
.topbar {
position: absolute; top: 20px; left: 22px;
display: flex; align-items: center; gap: 6px; text-decoration: none;
}
.topbar .logo { width: 40px; height: 40px; display: block; }
.topbar span {
font-size: 1.1rem; font-weight: 600; letter-spacing: -.01em; color: #fff;
transition: color .15s ease;
}
.topbar:hover span { color: #c9c9c9; }
.brand {
font-size: 2.1rem; font-weight: 700; letter-spacing: -.02em; color: #fff;
text-align: center; margin: 0 0 10px;
}
h1 {
font-size: 1.35rem; font-weight: 600; letter-spacing: -.01em; color: #f5f5f5;
text-align: center; margin: 0 0 28px;
}
.card {
width: 100%; max-width: 430px; text-align: center;
background: #171717; border: 1px solid rgba(255, 255, 255, .06);
border-radius: 24px; padding: 40px 40px 34px;
}
.badge {
margin: 0 auto 22px; width: 52px; height: 52px; border-radius: 50%;
display: flex; align-items: center; justify-content: center; font-size: 23px; color: #fff;
background: rgba(255, 255, 255, .05); border: 1px solid rgba(255, 255, 255, .14);
}
.msg { margin: 0 auto; max-width: 34ch; color: #b5b5b5; line-height: 1.6; font-size: .98rem; }
.rule { height: 1px; background: rgba(255, 255, 255, .07); margin: 26px 0 0; }
.tagline { margin: 22px 0 0; color: #7c7c7c; font-size: .9rem; line-height: 1.55; }
.tagline b { color: #ededed; font-weight: 500; }
.links {
margin-top: 18px; display: flex; gap: 8px; justify-content: center;
align-items: center; flex-wrap: wrap; font-size: .84rem;
}
.links a { color: #a3a3a3; text-decoration: none; transition: color .15s ease; }
.links a:hover { color: #fff; }
.links .dot { color: #3a3a3a; }
.close { margin: 24px 0 0; color: #5a5a5a; font-size: .78rem; text-align: center; }
</style></head>
<body>
<a class="topbar" href="https://strix.ai" target="_blank" rel="noopener"
aria-label="Strix — strix.ai">
<!--LOGO-->
<span>Strix</span>
</a>
<div class="brand">Strix</div>
<h1>You're signed in</h1>
<main class="card">
<div class="badge">✓</div>
<p class="msg">Strix is connected to your ChatGPT subscription. Head back to your
terminal — your security test runs there.</p>
<div class="rule"></div>
<p class="tagline">Autonomous AI hackers that <b>find and fix</b> your app's
vulnerabilities.</p>
<nav class="links">
<a href="https://strix.ai" target="_blank" rel="noopener">strix.ai</a>
<span class="dot">·</span>
<a href="https://docs.strix.ai" target="_blank" rel="noopener">docs</a>
<span class="dot">·</span>
<a href="https://discord.gg/strix-ai" target="_blank" rel="noopener">community</a>
</nav>
</main>
<p class="close">You can close this tab.</p>
</body></html>"""
__all__ = ["run_auth"]
+66 -10
View File
@@ -20,6 +20,7 @@ from rich.text import Text
from strix.config import (
apply_config_override,
codex,
load_settings,
persist_current,
)
@@ -92,6 +93,16 @@ def validate_environment() -> None:
settings = load_settings()
if codex.subscription_model(settings.llm.model):
if not codex.is_authenticated():
console.print(
f"[red]STRIX_LLM={settings.llm.model} uses your ChatGPT subscription, "
"but you're not signed in.[/] Run [cyan]strix auth login chatgpt[/] first."
)
sys.exit(1)
logger.info("Environment OK (ChatGPT subscription)")
return
if not settings.llm.model:
missing_required_vars.append("STRIX_LLM")
@@ -274,6 +285,29 @@ def _provider_import_hint(exc: BaseException, model: str) -> str | None:
return None
def _subscription_error_hint(exc: BaseException) -> str | None:
"""Return an actionable hint for a known ChatGPT-subscription error, or None."""
if not codex.subscription_model(load_settings().llm.model):
return None
joined = " ".join(_exception_messages(exc)).lower()
if "not supported when using codex with a chatgpt account" in joined:
return (
"This model isn't available on your ChatGPT subscription. "
"Set STRIX_LLM to a model your plan includes (e.g. chatgpt/gpt-5.4)."
)
if (
"error code: 401" in joined
or "http 401" in joined
or "unauthorized" in joined
or "invalid_grant" in joined
):
return (
"Your ChatGPT sign-in has expired or was revoked. Sign in again:\n"
" strix auth login chatgpt"
)
return None
async def warm_up_llm(show_model_warning: bool = True) -> None:
console = Console()
logger.info("Warming up LLM connection")
@@ -283,8 +317,8 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
settings = load_settings()
configure_sdk_model_defaults(settings)
llm = settings.llm
raw_model = (llm.model or "").strip()
if (
raw_model
and "/" not in raw_model
@@ -363,20 +397,33 @@ async def warm_up_llm(show_model_warning: bool = True) -> None:
except Exception as e:
logger.exception("LLM warm-up failed")
error_text = Text()
error_text.append("LLM CONNECTION FAILED", style="bold red")
error_text.append("\n\n", style="white")
error_text.append("Could not establish connection to the language model.\n", style="white")
error_text.append("Please check your configuration and try again.\n", style="white")
hint = _provider_import_hint(e, raw_model)
if hint is not None:
error_text.append(f"\n{hint}\n", style="bold yellow")
error_text.append(f"\nError: {e}", style="dim white")
sub_hint = _subscription_error_hint(e)
if sub_hint is not None:
# The model/backend answered with a clear, actionable rejection —
# show that instead of a generic "connection failed".
border_style = "yellow"
error_text.append("MODEL NOT AVAILABLE ON SUBSCRIPTION", style="bold yellow")
error_text.append("\n\n", style="white")
error_text.append(f"{sub_hint}\n", style="white")
error_text.append(f"\nDetails: {e}", style="dim white")
else:
border_style = "red"
error_text.append("LLM CONNECTION FAILED", style="bold red")
error_text.append("\n\n", style="white")
error_text.append(
"Could not establish connection to the language model.\n", style="white"
)
error_text.append("Please check your configuration and try again.\n", style="white")
hint = _provider_import_hint(e, raw_model)
if hint is not None:
error_text.append(f"\n{hint}\n", style="bold yellow")
error_text.append(f"\nError: {e}", style="dim white")
panel = Panel(
error_text,
title="[bold white]STRIX",
title_align="left",
border_style="red",
border_style=border_style,
padding=(1, 2),
)
@@ -682,6 +729,7 @@ def _persist_run_record(args: argparse.Namespace) -> None:
"status": "running",
"start_time": datetime.now(UTC).isoformat(),
"end_time": None,
"auth_mode": codex.auth_mode(load_settings().llm.model),
"targets_info": args.targets_info,
"scan_mode": args.scan_mode,
"instruction": args.instruction,
@@ -882,6 +930,13 @@ def main() -> None:
run_view(sys.argv[2:])
return
# `strix auth …` manages model-subscription sign-in and exits; it needs no
# target, Docker, or scan setup.
if len(sys.argv) > 1 and sys.argv[1] == "auth":
from strix.interface.auth_cli import run_auth
sys.exit(run_auth(sys.argv[2:]))
args = parse_arguments()
if args.config:
@@ -949,6 +1004,7 @@ def main() -> None:
_telemetry_start_kwargs = {
"model": load_settings().llm.model,
"auth_mode": codex.auth_mode(load_settings().llm.model),
"scan_mode": args.scan_mode,
"is_whitebox": is_whitebox_scan(args.targets_info),
"interactive": not args.non_interactive,
+20 -7
View File
@@ -802,6 +802,7 @@ class StrixTUIApp(App): # type: ignore[misc]
self._scan_stop_event = threading.Event()
self._scan_completed = threading.Event()
self._scan_error: BaseException | None = None
self._error_noted_agents: set[str] = set()
self._spinner_frame_index: int = 0
self._sweep_num_squares: int = 6
@@ -1015,22 +1016,32 @@ class StrixTUIApp(App): # type: ignore[misc]
else:
self._agent_graph_sync_future = None
try:
parent_of, statuses, names = future.result()
parent_of, statuses, names, errors = future.result()
except Exception:
logger.exception("TUI agent graph sync failed")
else:
for agent_id, status in statuses.items():
error = errors.get(agent_id)
self.live_view.upsert_agent(
agent_id,
name=names.get(agent_id, agent_id),
parent_id=parent_of.get(agent_id),
status=status,
error_message=error,
)
if status in {"failed", "crashed"} and error:
if agent_id not in self._error_noted_agents:
self._error_noted_agents.add(agent_id)
self.live_view.record_agent_error(agent_id, error)
else:
self._error_noted_agents.discard(agent_id)
if self._scan_loop is None or self._scan_loop.is_closed():
return
async def collect() -> tuple[dict[str, str | None], dict[str, Any], dict[str, str]]:
async def collect() -> tuple[
dict[str, str | None], dict[str, Any], dict[str, str], dict[str, str]
]:
return await self.coordinator.graph_snapshot()
self._agent_graph_sync_future = asyncio.run_coroutine_threadsafe(collect(), self._scan_loop)
@@ -1049,6 +1060,7 @@ class StrixTUIApp(App): # type: ignore[misc]
"waiting": "",
"completed": "🟢",
"failed": "🔴",
"crashed": "🔴",
"stopped": "",
}
@@ -1234,13 +1246,12 @@ class StrixTUIApp(App): # type: ignore[misc]
text.append(msg)
return (text, Text(), False)
if status == "failed":
if status in {"failed", "crashed"}:
error_msg = agent_data.get("error_message", "")
text = Text()
if error_msg:
text.append(error_msg, style="red")
else:
text.append("Scan failed", style="red")
text.append(error_msg or "Agent failed", style="red")
text.append(" · ", style="dim")
text.append("Send message to resume", style="dim")
self._stop_dot_animation()
return (text, Text(), False)
@@ -1539,6 +1550,7 @@ class StrixTUIApp(App): # type: ignore[misc]
"waiting": "",
"completed": "🟢",
"failed": "🔴",
"crashed": "🔴",
"stopped": "",
}
@@ -1584,6 +1596,7 @@ class StrixTUIApp(App): # type: ignore[misc]
"waiting": "",
"completed": "🟢",
"failed": "🔴",
"crashed": "🔴",
"stopped": "",
}
+11
View File
@@ -86,6 +86,17 @@ class TuiLiveView:
current["error_message"] = error_message
current["updated_at"] = now
def record_agent_error(self, agent_id: str, error: str) -> None:
self._append_event(
agent_id,
"chat",
{
"role": "assistant",
"content": (f"An error occurred: {error}\nI'm now waiting for new instructions."),
"metadata": {"source": "agent_error"},
},
)
def record_user_message(self, agent_id: str, content: str) -> None:
self._append_event(
agent_id,
+38 -6
View File
@@ -253,6 +253,20 @@ def _llm_usage(report_state: Any) -> dict[str, Any]:
return usage if isinstance(usage, dict) else {}
def _is_subscription(report_state: Any) -> bool:
"""Whether this run uses a model subscription (no metered cost).
Prefers the run record so it's correct for hydrated/resumed runs; falls back
to current settings.
"""
record = getattr(report_state, "run_record", None)
if isinstance(record, dict) and record.get("auth_mode"):
return record.get("auth_mode") == "subscription"
from strix.config import codex
return codex.auth_mode(load_settings().llm.model) == "subscription"
def _int_stat(usage: dict[str, Any], key: str) -> int:
try:
return max(0, int(usage.get(key) or 0))
@@ -283,11 +297,16 @@ def _build_llm_usage_stats(
*,
live: bool = False,
) -> None:
subscription = _is_subscription(report_state)
usage = _llm_usage(report_state)
if not usage or _int_stat(usage, "requests") <= 0:
stats_text.append("\n")
stats_text.append("Cost ", style="dim")
stats_text.append("$0.0000 ", style="#fbbf24")
if subscription:
stats_text.append("$0.00 ", style="#22c55e")
stats_text.append("(subscription) ", style="dim")
else:
stats_text.append("$0.0000 ", style="#fbbf24")
stats_text.append("· ", style="dim white")
stats_text.append("Tokens ", style="dim")
stats_text.append("0", style="white")
@@ -312,7 +331,12 @@ def _build_llm_usage_stats(
stats_text.append("Output Tokens ", style="dim")
stats_text.append(format_token_count(output_tokens), style="white")
if live or cost > 0:
if subscription:
stats_text.append(" · ", style="dim white")
stats_text.append("Cost ", style="dim")
stats_text.append("$0.00", style="#22c55e")
stats_text.append(" (subscription)", style="dim")
elif live or cost > 0:
stats_text.append(" · ", style="dim white")
stats_text.append("Cost ", style="dim")
stats_text.append(f"${cost:.4f}", style="#fbbf24")
@@ -337,6 +361,9 @@ def build_live_stats_text(report_state: Any) -> Text:
model = load_settings().llm.model or "unknown"
stats_text.append("Model ", style="dim")
stats_text.append(str(model), style="white")
if _is_subscription(report_state):
stats_text.append(" · ", style="dim white")
stats_text.append("ChatGPT subscription", style="#22c55e")
stats_text.append("\n")
vuln_count = len(report_state.vulnerability_reports)
@@ -379,6 +406,10 @@ def build_tui_stats_text(report_state: Any) -> Text:
model = load_settings().llm.model or "unknown"
stats_text.append(str(model), style="white")
subscription = _is_subscription(report_state)
if subscription:
stats_text.append("\n")
stats_text.append("ChatGPT subscription", style="#22c55e")
usage = _llm_usage(report_state)
if usage and _int_stat(usage, "total_tokens") > 0:
@@ -388,7 +419,10 @@ def build_tui_stats_text(report_state: Any) -> Text:
style="white",
)
cost = _float_stat(usage, "cost")
if cost > 0:
if subscription:
stats_text.append(" · ", style="white")
stats_text.append("$0.00", style="white")
elif cost > 0:
stats_text.append(" · ", style="white")
stats_text.append(f"${cost:.2f}", style="white")
@@ -1147,9 +1181,7 @@ def read_target_list_file(path_str: str) -> list[str]:
if (target := line.strip()) and not target.startswith("#")
]
except UnicodeDecodeError as e:
raise ValueError(
f"Target list file '{path_str}' must be valid UTF-8 text: {e!s}"
) from e
raise ValueError(f"Target list file '{path_str}' must be valid UTF-8 text: {e!s}") from e
except OSError as e:
raise ValueError(f"Failed to read target list file '{path_str}': {e!s}") from e
+5
View File
@@ -10,6 +10,8 @@ from uuid import uuid4
from agents.usage import Usage
from strix.config import codex
from strix.config.loader import load_settings
from strix.core.paths import run_dir_for
from strix.report.sarif import write_sarif
from strix.report.usage import LLMUsageLedger
@@ -117,12 +119,15 @@ class ReportState:
self.scan_results: dict[str, Any] | None = None
self.scan_config: dict[str, Any] | None = None
self._llm_usage = LLMUsageLedger()
auth_mode = codex.auth_mode(load_settings().llm.model)
self._llm_usage.zero_cost = auth_mode == "subscription"
self.run_record: dict[str, Any] = {
"run_id": self.run_id,
"run_name": self.run_name,
"start_time": self.start_time,
"end_time": None,
"status": "running",
"auth_mode": auth_mode,
"targets_info": [],
"llm_usage": self._build_llm_usage_record(),
}
+6 -1
View File
@@ -19,6 +19,9 @@ class LLMUsageLedger:
self._agent_usage: dict[str, Usage] = {}
self._agent_metadata: dict[str, dict[str, str]] = {}
self._total_cost = 0.0
# 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.
self.zero_cost = False
def record(
self,
@@ -41,7 +44,7 @@ class LLMUsageLedger:
if model:
metadata["model"] = model
if not _is_litellm_routed(model):
if not self.zero_cost and not _is_litellm_routed(model):
estimated = _estimate_litellm_cost(usage, model)
if estimated:
self._total_cost += estimated
@@ -49,6 +52,8 @@ class LLMUsageLedger:
return True
def record_observed_cost(self, cost: float) -> None:
if self.zero_cost:
return
if isinstance(cost, int | float) and cost > 0:
self._total_cost += float(cost)
+1 -1
View File
@@ -15,7 +15,7 @@ We collect only very **basic** usage data including:
**Session Errors:** Duration and error types (not messages or stack traces)\
**System Context:** OS type, architecture, Strix version\
**Scan Context:** Scan mode (quick/standard/deep), scan type (whitebox/blackbox)\
**Model Usage:** Which LLM model is being used (not prompts or responses)\
**Model Usage:** Which LLM model is being used and whether it runs via an API key or a model subscription (not prompts or responses)\
**Feature Usage:** Which built-in skills are loaded\
**Aggregate Metrics:** Vulnerability counts by severity and weakness category (CWE)
+3
View File
@@ -58,12 +58,14 @@ def start(
is_whitebox: bool,
interactive: bool,
has_instructions: bool,
auth_mode: str | None = None,
) -> None:
_send(
"scan_started",
{
**base_props(),
"model": model or "unknown",
"auth_mode": auth_mode or "api_key",
"scan_mode": scan_mode or "unknown",
"scan_type": "whitebox" if is_whitebox else "blackbox",
"interactive": interactive,
@@ -133,6 +135,7 @@ def end(report_state: "ReportState", exit_reason: str = "completed") -> None:
"scan_ended",
{
**base_props(),
"auth_mode": report_state.run_record.get("auth_mode") or "api_key",
"exit_reason": report_state.scan_ended_exit_reason,
"duration_seconds": round(duration),
"vulnerabilities_total": len(report_state.vulnerability_reports),
+3
View File
@@ -59,6 +59,7 @@ def start(
is_whitebox: bool,
interactive: bool,
has_instructions: bool,
auth_mode: str | None = None,
) -> None:
_send(
"scan_started",
@@ -66,6 +67,7 @@ def start(
**base_props(),
"session": SESSION_ID,
"model": model or "unknown",
"auth_mode": auth_mode or "api_key",
"scan_mode": scan_mode or "unknown",
"scan_type": "whitebox" if is_whitebox else "blackbox",
"interactive": interactive,
@@ -140,6 +142,7 @@ def end(report_state: ReportState, exit_reason: str = "completed") -> None:
{
**base_props(),
"session": SESSION_ID,
"auth_mode": report_state.run_record.get("auth_mode") or "api_key",
"exit_reason": report_state.scan_ended_exit_reason,
"duration_seconds": round(duration),
"vulnerabilities_total": len(report_state.vulnerability_reports),
+2 -2
View File
@@ -87,7 +87,7 @@ async def view_agent_graph(ctx: RunContextWrapper) -> str:
default=str,
)
parent_of, statuses, names = await coordinator.graph_snapshot()
parent_of, statuses, names, _ = await coordinator.graph_snapshot()
lines: list[str] = []
@@ -635,7 +635,7 @@ async def stop_agent(
ensure_ascii=False,
default=str,
)
_, statuses, _ = await coordinator.graph_snapshot()
_, statuses, _, _ = await coordinator.graph_snapshot()
if target_agent_id not in statuses:
return json.dumps(
{"success": False, "error": f"Unknown agent_id: {target_agent_id}"},
-9
View File
@@ -63,7 +63,6 @@
"integrity": "sha512-RgHBCvtjbOK2gXSNBNIkNoEc9qoVEtau3hj8gEqKQuL3HZAibKarWFEI3Lfm6EYKkLalOh8eSrj9b+ch9H/VBA==",
"dev": true,
"license": "MIT",
"peer": true,
"dependencies": {
"@babel/code-frame": "^7.29.7",
"@babel/generator": "^7.29.7",
@@ -1605,7 +1604,6 @@
"resolved": "https://registry.npmjs.org/@types/react/-/react-19.2.17.tgz",
"integrity": "sha512-MXfmqaVPEVgkBT/aY0aGCkRWWtByiYQXo3xdQ8r5RzuFrPiRn8Gar2tQdXSUQ2GKV3bkXckek89V8wQBY2Q/Aw==",
"license": "MIT",
"peer": true,
"dependencies": {
"csstype": "^3.2.2"
}
@@ -1616,7 +1614,6 @@
"integrity": "sha512-jp2L/eY6fn+KgVVQAOqYItbF0VY/YApe5Mz2F0aykSO8gx31bYCZyvSeYxCHKvzHG5eZjc+zyaS5BrBWya2+kQ==",
"devOptional": true,
"license": "MIT",
"peer": true,
"peerDependencies": {
"@types/react": "^19.2.0"
}
@@ -1739,7 +1736,6 @@
}
],
"license": "MIT",
"peer": true,
"dependencies": {
"baseline-browser-mapping": "^2.10.42",
"caniuse-lite": "^1.0.30001803",
@@ -1920,7 +1916,6 @@
"resolved": "https://registry.npmjs.org/d3-selection/-/d3-selection-3.0.0.tgz",
"integrity": "sha512-fmTRWbNMmsmWq6xJV8D19U/gw/bwrHfNXxrIN+HfZgnzqTHp9jOmKMhsTUjXOJnZOdZY9Q28y4yebKzqDKlxlQ==",
"license": "ISC",
"peer": true,
"engines": {
"node": ">=12"
}
@@ -3571,7 +3566,6 @@
"integrity": "sha512-RvwwcruNjI1ncT5xRakeyS9Lf8lcItv34KD+aif+VH9kduAyfYBipGh12274xtenIPZ119/R9BdTBa8gAwSh0A==",
"dev": true,
"license": "MIT",
"peer": true,
"engines": {
"node": ">=12"
},
@@ -3623,7 +3617,6 @@
"resolved": "https://registry.npmjs.org/react/-/react-19.2.7.tgz",
"integrity": "sha512-HNe9WslTbXmFK8o8cmwgAeJFSBvt1bPdHCVKtaaV+WlAN36mpT4hcRpwbf3fY56ar2oIXzsBpOAiIRHAdY0OlQ==",
"license": "MIT",
"peer": true,
"engines": {
"node": ">=0.10.0"
}
@@ -3633,7 +3626,6 @@
"resolved": "https://registry.npmjs.org/react-dom/-/react-dom-19.2.7.tgz",
"integrity": "sha512-t0BRVXvbiE/o20Hfw669rLbMCDWtYZLvmJigy2f0MxsXF+71pxhR3xOkspmsO8h3ZlNzyibAmtCa3l4lYKk6gQ==",
"license": "MIT",
"peer": true,
"dependencies": {
"scheduler": "^0.27.0"
},
@@ -4109,7 +4101,6 @@
"integrity": "sha512-NTKlcQjlAK7MlQoyb6LgaqHc8sso/pVyUJYWMws3jg21uTJw/LddqIFPcPqP6PzpgbIcZyKI85sFE4HBrQDA8A==",
"dev": true,
"license": "MIT",
"peer": true,
"dependencies": {
"esbuild": "^0.25.0",
"fdir": "^6.4.4",
@@ -100,6 +100,7 @@ export function RunDetails({
const reasoning = num(rec(arr(usage.output_tokens_details)[0]).reasoning_tokens);
const totalTokens = num(usage.total_tokens);
const cost = num(usage.cost);
const subscription = str(raw.auth_mode) === "subscription";
const sub = (n: number, word: string) => (
<span className="text-[#666]"> ({formatNumber(n)} {word})</span>
@@ -175,6 +176,15 @@ export function RunDetails({
{hasUsage ? (
<dl className="space-y-2.5 tabular-nums">
<Field label="Model">{models.length ? models.join(", ") : "n/a"}</Field>
{subscription && (
<Field label="Provider">
<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]">
ChatGPT subscription
</span>
</span>
</Field>
)}
<Field label="Run time">{fmtDuration(durationSeconds)}</Field>
{requests != null && <Field label="Requests">{formatNumber(requests)}</Field>}
{inputTokens != null && (
@@ -190,7 +200,14 @@ export function RunDetails({
</Field>
)}
{totalTokens != null && <Field label="Total tokens">{formatNumber(totalTokens)}</Field>}
{cost != null && <Field label="Cost">${cost.toFixed(2)}</Field>}
{subscription ? (
<Field label="Cost">
<span className="text-[#22c55e]">$0.00</span>
<span className="text-[#666]"> (subscription)</span>
</Field>
) : (
cost != null && <Field label="Cost">${cost.toFixed(2)}</Field>
)}
{agents.length > 0 && <Field label="Agents">{formatNumber(agents.length)}</Field>}
</dl>
) : (
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+2 -2
View File
@@ -6,8 +6,8 @@
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<meta name="color-scheme" content="dark" />
<title>Strix Results</title>
<script type="module" crossorigin src="./assets/index-BNKUksp9.js"></script>
<link rel="stylesheet" crossorigin href="./assets/index-BdiSGmzb.css">
<script type="module" crossorigin src="./assets/index-e2r6VuTm.js"></script>
<link rel="stylesheet" crossorigin href="./assets/index-vV8wxCG6.css">
</head>
<body>
<div id="root"></div>
+106
View File
@@ -0,0 +1,106 @@
"""Tests for the `strix auth` CLI: subcommand routing and provider naming."""
from __future__ import annotations
from typing import TYPE_CHECKING, Any
import pytest
from strix.config import codex
from strix.interface import auth_cli
if TYPE_CHECKING:
from pathlib import Path
@pytest.fixture(autouse=True)
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(codex, "AUTH_PATH", tmp_path / "home" / ".strix" / "subscription-auth.json")
def test_login_provider_is_chatgpt() -> None:
assert auth_cli.LOGIN_PROVIDER == "chatgpt"
assert codex.PROVIDER in auth_cli._ACCEPTED_PROVIDERS
assert "chatgpt" in auth_cli._ACCEPTED_PROVIDERS
def test_unknown_subcommand_returns_usage_error() -> None:
assert auth_cli.run_auth(["bogus"]) == 2
def test_help_returns_zero() -> None:
assert auth_cli.run_auth(["--help"]) == 0
def test_status_not_signed_in() -> None:
assert auth_cli.run_auth(["status"]) == 1
def test_login_rejects_unsupported_provider(monkeypatch: pytest.MonkeyPatch) -> None:
def _should_not_run(*_args: Any, **_kwargs: Any) -> dict[str, Any]:
msg = "OAuth flow must not start for an unsupported provider"
raise AssertionError(msg)
monkeypatch.setattr(auth_cli, "_run_oauth_flow", _should_not_run)
assert auth_cli.run_auth(["login", "gemini"]) == 2
def test_finish_requires_state_on_loopback(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(codex, "exchange_code", lambda *_: {"ok": True})
# Loopback (require_state=True): missing or mismatched state is rejected.
with pytest.raises(codex.CodexAuthError) as missing:
auth_cli._finish("code", None, "verifier", "expected", require_state=True)
assert missing.value.code == "state_mismatch"
with pytest.raises(codex.CodexAuthError) as mismatch:
auth_cli._finish("code", "wrong", "verifier", "expected", require_state=True)
assert mismatch.value.code == "state_mismatch"
# Matching state proceeds to the exchange.
assert auth_cli._finish("code", "expected", "verifier", "expected", require_state=True) == {
"ok": True
}
def test_finish_manual_paste_allows_absent_state(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(codex, "exchange_code", lambda *_: {"ok": True})
# Manual paste (require_state=False): a bare code with no state is accepted,
# but a present-and-wrong state is still rejected.
assert auth_cli._finish("code", None, "verifier", "expected", require_state=False) == {
"ok": True
}
with pytest.raises(codex.CodexAuthError):
auth_cli._finish("code", "wrong", "verifier", "expected", require_state=False)
def test_finish_rejects_missing_code() -> None:
with pytest.raises(codex.CodexAuthError) as exc:
auth_cli._finish(None, "expected", "verifier", "expected", require_state=True)
assert exc.value.code == "no_code"
def test_model_subcommand_removed() -> None:
assert auth_cli.run_auth(["model", "gpt-5.5"]) == 2
@pytest.mark.parametrize("provider", ["chatgpt", "codex", "ChatGPT"])
def test_login_accepts_provider_aliases(provider: str, monkeypatch: pytest.MonkeyPatch) -> None:
reached = {"flow": False}
def _fake_flow(*_args: Any, **_kwargs: Any) -> dict[str, Any]:
reached["flow"] = True
return {
"type": "oauth",
"provider": "codex",
"access": "a",
"refresh": "r",
"account_id": "acct",
"expires_at": 0,
}
monkeypatch.setattr(auth_cli, "_run_oauth_flow", _fake_flow)
monkeypatch.setattr(codex, "save_record", lambda _record: None)
assert auth_cli.run_auth(["login", provider]) == 0
assert reached["flow"] is True
+294
View File
@@ -0,0 +1,294 @@
"""Tests for ChatGPT (Codex) subscription auth: PKCE, token handling, store."""
from __future__ import annotations
import base64
import hashlib
import json
import time
from typing import TYPE_CHECKING, Any
import pytest
from strix.config import codex
if TYPE_CHECKING:
from pathlib import Path
def _fake_jwt(account_id: str) -> str:
def seg(obj: dict[str, Any]) -> str:
return base64.urlsafe_b64encode(json.dumps(obj).encode()).rstrip(b"=").decode()
header = seg({"alg": "none"})
payload = seg({"https://api.openai.com/auth": {"chatgpt_account_id": account_id}})
return f"{header}.{payload}.sig"
@pytest.fixture(autouse=True)
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path:
path = tmp_path / "home" / ".strix" / "subscription-auth.json"
monkeypatch.setattr(codex, "AUTH_PATH", path)
return path
def test_pkce_challenge_matches_verifier_and_is_unpadded() -> None:
verifier, challenge = codex.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_and_client() -> None:
url = codex.build_authorize_url("chal", "st8")
assert codex.AUTHORIZE_URL in url
assert "code_challenge=chal" in url
assert "code_challenge_method=S256" in url
assert f"client_id={codex.CLIENT_ID}" in url
assert "state=st8" in url
@pytest.mark.parametrize(
("value", "expected"),
[
("http://localhost:1455/auth/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 codex.parse_redirect_input(value) == expected
@pytest.mark.parametrize(
("model", "expected"),
[
("chatgpt/gpt-5.4", "gpt-5.4"),
("ChatGPT/GPT-5.5", "GPT-5.5"),
(" chatgpt/gpt-5.4 ", "gpt-5.4"),
("openai/gpt-5.4", None), # metered API path
("anthropic/claude-opus-4-8", None),
("gpt-5.4", None),
("chatgpt/", None),
("", None),
(None, None),
],
)
def test_subscription_model(model: str | None, expected: str | None) -> None:
assert codex.subscription_model(model) == expected
def test_auth_mode() -> None:
assert codex.auth_mode("chatgpt/gpt-5.4") == "subscription"
assert codex.auth_mode("openai/gpt-5.4") == "api_key"
assert codex.auth_mode("anthropic/claude-opus-4-8") == "api_key"
assert codex.auth_mode(None) == "api_key"
def test_is_content_guardrail_error() -> None:
# The backend's real wording (from a live gpt-5.6-sol block).
raw = RuntimeError(
"This content was flagged for possible cybersecurity risk. If this seems "
"wrong, try rephrasing. To get authorized, join the Trusted Access for Cyber program."
)
assert codex.is_content_guardrail_error(raw) is True
# The already-typed error is recognized regardless of its message wording.
assert codex.is_content_guardrail_error(codex.CodexContentGuardrailError("gpt-5.6-sol")) is True
# Unrelated errors are not misclassified.
assert codex.is_content_guardrail_error(RuntimeError("rate limit exceeded")) is False
def test_content_guardrail_error_message() -> None:
err = codex.CodexContentGuardrailError("gpt-5.6-sol")
assert err.model == "gpt-5.6-sol"
assert "gpt-5.6-sol" in str(err)
assert "STRIX_LLM" in str(err)
def test_account_id_from_jwt() -> None:
assert codex._account_id_from_jwt(_fake_jwt("acct-42")) == "acct-42"
assert codex._account_id_from_jwt("not-a-jwt") is None
assert codex._account_id_from_jwt("") is None
def test_store_roundtrip_and_logout() -> None:
assert codex.read_record() is None
assert codex.is_authenticated() is False
codex.save_record(
{
"type": "oauth",
"provider": "codex",
"access": _fake_jwt("acct-42"),
"refresh": "r1",
"account_id": "acct-42",
"expires_at": time.time() + 3600,
}
)
record = codex.read_record()
assert record is not None
assert record["account_id"] == "acct-42"
assert codex.is_authenticated() is True
codex.logout()
assert codex.read_record() is None
codex.logout() # no-op when already gone
def test_read_record_rejects_incomplete_records() -> None:
codex.save_record({"type": "oauth", "access": "a"}) # missing refresh/account
assert codex.read_record() is None
assert codex.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(codex, "_post_form", _boom)
codex.save_record(
{
"type": "oauth",
"provider": "codex",
"access": "access-fresh",
"refresh": "r1",
"account_id": "acct-42",
"expires_at": time.time() + 3600,
}
)
assert codex.get_valid_token() == ("access-fresh", "acct-42")
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": _fake_jwt("acct-42"), "refresh_token": "r2", "expires_in": 3600}
monkeypatch.setattr(codex, "_post_form", _fake_post)
codex.save_record(
{
"type": "oauth",
"provider": "codex",
"access": "stale",
"refresh": "r1",
"account_id": "acct-42",
"expires_at": time.time() - 10, # already expired
}
)
_access, account_id = codex.get_valid_token()
assert calls["n"] == 1
assert account_id == "acct-42"
# Rotated refresh token was written back to the store.
record = codex.read_record()
assert record is not None
assert record["refresh"] == "r2"
def test_get_valid_token_uses_token_rotated_by_another_process(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Simulate a parallel Strix process rotating the token while we wait for the
# refresh guard: the pre-guard read sees the stale token, the in-guard read
# sees the winner's fresh one, so we must NOT exchange the now-dead refresh.
records = [
{
"type": "oauth",
"provider": "codex",
"access": "stale",
"refresh": "r1",
"account_id": "acct",
"expires_at": time.time() - 10,
},
{
"type": "oauth",
"provider": "codex",
"access": "fresh-from-other-process",
"refresh": "r2",
"account_id": "acct",
"expires_at": 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(codex, "read_record", _fake_read)
monkeypatch.setattr(codex, "_post_form", _boom)
access, account_id = codex.get_valid_token()
assert access == "fresh-from-other-process"
assert account_id == "acct"
def _expired_record(refresh: str, access: str) -> dict[str, Any]:
return {
"type": "oauth",
"provider": "codex",
"access": access,
"refresh": refresh,
"account_id": "acct-42",
"expires_at": time.time() - 10,
}
def test_get_valid_token_recovers_when_refresh_loses_race(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Lock failed open: our in-guard read still saw the stale token, so we tried to
# refresh and lost the race (invalid_grant). By then a peer has saved a fresh
# token — recover from it instead of failing the scan on the dead one.
codex.save_record(_expired_record("r1", "stale"))
def _fake_post(_payload: dict[str, str]) -> dict[str, Any]:
codex.save_record(
{
"type": "oauth",
"provider": "codex",
"access": "fresh-from-peer",
"refresh": "r2",
"account_id": "acct-42",
"expires_at": time.time() + 3600,
}
)
raise codex.CodexAuthError("token_http_error", "HTTP 400: invalid_grant")
monkeypatch.setattr(codex, "_post_form", _fake_post)
assert codex.get_valid_token() == ("fresh-from-peer", "acct-42")
def test_get_valid_token_reraises_refresh_error_without_rotation(
monkeypatch: pytest.MonkeyPatch,
) -> None:
# Refresh fails and no peer rotated the token: surface the error, don't mask it.
codex.save_record(_expired_record("r1", "stale"))
def _fake_post(_payload: dict[str, str]) -> dict[str, Any]:
raise codex.CodexAuthError("token_http_error", "HTTP 400: invalid_grant")
monkeypatch.setattr(codex, "_post_form", _fake_post)
with pytest.raises(codex.CodexAuthError):
codex.get_valid_token()
def test_get_valid_token_raises_when_not_signed_in() -> None:
with pytest.raises(codex.CodexAuthError) as exc:
codex.get_valid_token()
assert exc.value.code == "not_authenticated"
+225
View File
@@ -0,0 +1,225 @@
"""Regression test for the ChatGPT Codex backend's streaming requirement.
The backend rejects non-streamed requests with ``{"detail": "Stream must be set
to true"}``. ``_CodexResponsesModel`` must therefore issue a streamed request
even from the non-streaming ``get_response`` path and aggregate the events into
a single response. A local server that mimics that behaviour proves the wrapper
works where the stock responses model would fail.
"""
from __future__ import annotations
import json
import threading
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import TYPE_CHECKING, Any
import pytest
from agents.model_settings import ModelSettings
from agents.models.interface import ModelTracing
from agents.models.openai_responses import OpenAIResponsesModel
from openai import AsyncOpenAI, BadRequestError
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from strix.config import codex
from strix.config.models import _CodexResponsesModel
if TYPE_CHECKING:
from collections.abc import AsyncIterator, Iterator
def _response_payload() -> dict[str, Any]:
return {
"id": "resp_1",
"object": "response",
"created_at": 0,
"status": "completed",
"model": "gpt-5.5",
"output": [
{
"type": "message",
"id": "m1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "OK", "annotations": []}],
}
],
"usage": {
"input_tokens": 1,
"output_tokens": 1,
"total_tokens": 2,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
},
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
"metadata": {},
"temperature": 1.0,
"top_p": 1.0,
"error": None,
"incomplete_details": None,
"instructions": None,
"max_output_tokens": None,
}
_CAPTURED: dict[str, Any] = {}
class _Handler(BaseHTTPRequestHandler):
def log_message(self, *args: Any) -> None:
pass
def do_POST(self) -> None:
length = int(self.headers.get("Content-Length", 0))
body = json.loads(self.rfile.read(length) or b"{}")
_CAPTURED.clear()
_CAPTURED.update(body)
if not body.get("stream"):
payload = json.dumps({"detail": "Stream must be set to true"}).encode()
self.send_response(400)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.end_headers()
self.wfile.write(payload)
return
event = {
"type": "response.completed",
"sequence_number": 0,
"response": _response_payload(),
}
frame = f"event: response.completed\ndata: {json.dumps(event)}\n\n".encode()
self.send_response(200)
self.send_header("Content-Type", "text/event-stream")
self.end_headers()
self.wfile.write(frame)
@pytest.fixture
def backend_url() -> Iterator[str]:
server = HTTPServer(("127.0.0.1", 0), _Handler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
try:
yield f"http://127.0.0.1:{server.server_address[1]}/backend-api/codex"
finally:
server.shutdown()
server.server_close()
def _client(base_url: str) -> AsyncOpenAI:
return AsyncOpenAI(api_key="tok", base_url=base_url)
def _call_kwargs() -> dict[str, Any]:
return {
"system_instructions": "s",
"input": "hi",
"model_settings": ModelSettings(
store=False, response_include=["reasoning.encrypted_content"]
),
"tools": [],
"output_schema": None,
"handoffs": [],
"tracing": ModelTracing.DISABLED,
"previous_response_id": None,
"conversation_id": None,
"prompt": None,
}
@pytest.mark.asyncio
async def test_stock_model_fails_on_non_streamed_backend(backend_url: str) -> None:
model = OpenAIResponsesModel(model="gpt-5.5", openai_client=_client(backend_url))
with pytest.raises(BadRequestError, match="Stream must be set to true"):
await model.get_response(**_call_kwargs())
@pytest.mark.asyncio
async def test_codex_model_streams_and_aggregates(backend_url: str) -> None:
model = _CodexResponsesModel(model="gpt-5.5", openai_client=_client(backend_url))
response = await model.get_response(**_call_kwargs())
message = response.output[0]
assert isinstance(message, ResponseOutputMessage)
text = message.content[0]
assert isinstance(text, ResponseOutputText)
assert text.text == "OK"
assert response.usage.total_tokens == 2
class _TrackingStream:
"""An async iterator that yields, then raises, and records if it was closed."""
def __init__(self, events: list[Any], error: Exception | None) -> None:
self._events = iter(events)
self._error = error
self.closed = False
def __aiter__(self) -> _TrackingStream:
return self
async def __anext__(self) -> Any:
try:
return next(self._events)
except StopIteration:
if self._error is not None:
raise self._error from None
raise StopAsyncIteration from None
async def aclose(self) -> None:
self.closed = True
async def _drain(gen: AsyncIterator[Any]) -> list[Any]:
return [event async for event in gen]
@pytest.mark.asyncio
async def test_guarded_converts_guardrail_error() -> None:
# A mid-stream backend rejection becomes a typed, model-tagged error.
model = _CodexResponsesModel(model="gpt-5.6-sol", openai_client=_client("http://x/backend-api"))
guardrail = RuntimeError("This content was flagged for possible cybersecurity risk.")
stream = _TrackingStream(["a", "b"], guardrail)
with pytest.raises(codex.CodexContentGuardrailError) as exc_info:
await _drain(model._guarded(stream))
assert exc_info.value.model == "gpt-5.6-sol"
assert stream.closed is True # underlying stream is released
@pytest.mark.asyncio
async def test_guarded_passes_through_other_errors() -> None:
# A non-guardrail error propagates unchanged (still not swallowed).
model = _CodexResponsesModel(model="gpt-5.5", openai_client=_client("http://x/backend-api"))
boom = RuntimeError("some unrelated failure")
stream = _TrackingStream(["a"], boom)
with pytest.raises(RuntimeError, match="some unrelated failure"):
await _drain(model._guarded(stream))
assert stream.closed is True
@pytest.mark.asyncio
async def test_guarded_yields_all_events_when_clean() -> None:
model = _CodexResponsesModel(model="gpt-5.4", openai_client=_client("http://x/backend-api"))
stream = _TrackingStream(["a", "b", "c"], None)
assert await _drain(model._guarded(stream)) == ["a", "b", "c"]
assert stream.closed is True
@pytest.mark.asyncio
async def test_codex_model_self_enforces_backend_requirements(backend_url: str) -> None:
# The caller passes ordinary settings; the model must impose the backend's
# requirements (stream, store=false, encrypted reasoning) and the configured
# reasoning effort itself.
model = _CodexResponsesModel(
model="gpt-5.4", openai_client=_client(backend_url), reasoning_effort="high"
)
kwargs = _call_kwargs()
kwargs["model_settings"] = ModelSettings() # nothing special from the caller
await model.get_response(**kwargs)
assert _CAPTURED["stream"] is True
assert _CAPTURED["store"] is False
assert _CAPTURED["include"] == ["reasoning.encrypted_content"]
assert _CAPTURED["reasoning"] == {"effort": "high"}
+17 -4
View File
@@ -13,12 +13,15 @@ import asyncio
from agents.retry import ModelRetryNormalizedError, RetryPolicyContext
from strix.config import codex
from strix.config.models import DEFAULT_MODEL_RETRY, _retry_statusless_provider_errors
def _context(normalized: ModelRetryNormalizedError) -> RetryPolicyContext:
def _context(
normalized: ModelRetryNormalizedError, error: Exception | None = None
) -> RetryPolicyContext:
return RetryPolicyContext(
error=RuntimeError("boom"),
error=error or RuntimeError("boom"),
attempt=1,
max_retries=5,
stream=True,
@@ -27,11 +30,11 @@ def _context(normalized: ModelRetryNormalizedError) -> RetryPolicyContext:
)
def _retries(normalized: ModelRetryNormalizedError) -> bool:
def _retries(normalized: ModelRetryNormalizedError, error: Exception | None = None) -> bool:
"""Evaluate the composed DEFAULT_MODEL_RETRY policy for a normalized error."""
policy = DEFAULT_MODEL_RETRY.policy
assert policy is not None
decision = asyncio.run(policy(_context(normalized)))
decision = asyncio.run(policy(_context(normalized, error)))
return bool(getattr(decision, "retry", decision))
@@ -63,6 +66,16 @@ def test_timeout_error_is_retried() -> None:
assert _retries(ModelRetryNormalizedError(is_network_error=True)) is True
def test_content_guardrail_error_is_not_retried() -> None:
# A guardrail block is status-less, so it would match the statusless policy;
# the guard must keep it from being retried (retrying never clears it).
guardrail = codex.CodexContentGuardrailError("gpt-5.6-sol")
assert _retries(ModelRetryNormalizedError(status_code=None), guardrail) is False
# A raw provider error carrying the backend's wording is excluded too.
raw = RuntimeError("This content was flagged for possible cybersecurity risk.")
assert _retries(ModelRetryNormalizedError(status_code=None), raw) is False
def test_policy_helper_matches_statusless_only() -> None:
assert _retry_statusless_provider_errors(_context(ModelRetryNormalizedError())) is True
assert (
+4
View File
@@ -42,9 +42,11 @@ def test_recommended_models_are_matched_case_insensitively() -> None:
"model_name",
[
"gpt-5.5",
"chatgpt/gpt-5.4",
"litellm/openai/gpt-5.4-pro",
"azure_ai/gpt-5.5-pro",
"bedrock_mantle/openai.gpt-5.5",
"anthropic/claude-opus-5",
"anthropic/claude-opus-4-8",
"anthropic.claude-opus-4-8",
"anthropic/claude-opus-4-7",
@@ -60,8 +62,10 @@ def test_recommended_models_are_matched_case_insensitively() -> None:
"deepseek/deepseek-reasoner",
"dashscope/qwen3-max-2026-01-23",
"qwen3.7-max",
"dashscope/qwen3.8-max",
"moonshot/kimi-k2.6",
"kimi-k2.7-code",
"moonshot/kimi-k3",
],
)
def test_frontier_model_families_are_accepted(model_name: str) -> None:
+2 -1
View File
@@ -46,7 +46,8 @@ def _patch_engine_scaffold(
reasoning_effort="high",
force_required_tool_choice=False,
timeout=300,
)
),
runtime=types.SimpleNamespace(max_context_images=3),
)
monkeypatch.setattr(runner, "load_settings", lambda: settings)
monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _settings: None)
+46
View File
@@ -0,0 +1,46 @@
"""Subscription runs track tokens but report zero cost."""
from __future__ import annotations
from agents.usage import Usage
from strix.report.usage import LLMUsageLedger
def _usage() -> Usage:
usage = Usage()
usage.requests = 1
usage.input_tokens = 1000
usage.output_tokens = 200
usage.total_tokens = 1200
return usage
def test_zero_cost_ledger_keeps_tokens_but_reports_no_cost() -> None:
ledger = LLMUsageLedger()
ledger.zero_cost = True
ledger.record(agent_id="a", usage=_usage(), agent_name="strix", model="gpt-5.5")
record = ledger.to_record()
assert record["cost"] == 0.0
assert record["total_tokens"] == 1200
assert record["input_tokens"] == 1000
assert record["output_tokens"] == 200
assert ledger.total_cost == 0.0
def test_zero_cost_ledger_ignores_observed_cost() -> None:
ledger = LLMUsageLedger()
ledger.zero_cost = True
ledger.record_observed_cost(4.20)
assert ledger.total_cost == 0.0
def test_normal_ledger_still_estimates_cost() -> None:
# Sanity check the flag is opt-in: without it, an OpenAI-native model still
# accrues an estimated cost (proves zeroing is what suppresses it).
ledger = LLMUsageLedger()
ledger.record(agent_id="a", usage=_usage(), agent_name="strix", model="gpt-5.5")
assert ledger.to_record()["total_tokens"] == 1200
# Cost estimation depends on litellm's cost map; it should be >= 0 and not error.
assert ledger.total_cost >= 0.0