mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 09:26:39 +02:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
399c15627b | ||
|
|
dce70c643a |
@@ -315,6 +315,18 @@ strix auth status # show the active sign-in
|
||||
strix auth logout # forget the sign-in
|
||||
```
|
||||
|
||||
#### Sign in with an OpenCode subscription
|
||||
|
||||
You can also run Strix on [OpenCode Zen](https://opencode.ai/docs/zen/) credits or an [OpenCode Go](https://opencode.ai/docs/go/) subscription:
|
||||
|
||||
```bash
|
||||
strix auth login opencode # paste your API key from opencode.ai/auth
|
||||
|
||||
export STRIX_LLM="opencode/claude-sonnet-5" # opencode/<model> runs on Zen credits
|
||||
export STRIX_LLM="opencode-go/kimi-k3" # opencode-go/<model> runs on the Go subscription
|
||||
strix --target ./app-directory
|
||||
```
|
||||
|
||||
**Recommended models for best results:**
|
||||
|
||||
- [OpenAI GPT-5.4](https://openai.com/api/) - `openai/gpt-5.4`
|
||||
|
||||
+1
-7
@@ -230,7 +230,6 @@ ignore = [
|
||||
# args they intentionally ignore.
|
||||
"tests/test_viewer_auth.py" = ["S105", "S106", "ARG001"]
|
||||
"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.
|
||||
"scripts/tui_sidecar_hook.py" = ["INP001"]
|
||||
# Stdlib HTTP handler overrides (do_GET/do_POST).
|
||||
@@ -245,9 +244,6 @@ ignore = [
|
||||
# Stdlib HTTP handler overrides (do_GET/do_POST) and lazy imports that avoid a
|
||||
# circular dependency with strix.telemetry / strix.interface.viewer.report_pdf.
|
||||
"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.
|
||||
"strix/interface/viewer/cli.py" = ["PLC0415"]
|
||||
# Lazy imports inside functions to avoid circular dependency with
|
||||
@@ -284,6 +280,7 @@ ignore = [
|
||||
# a runtime ``Callable`` annotation on ``vulnerability_found_callback``.
|
||||
"strix/report/state.py" = ["TC003", "PLR0912", "PLR0915", "E501", "PERF401", "PLC0415"]
|
||||
"strix/report/usage.py" = ["PLC0415"]
|
||||
"strix/report/pricing.py" = ["PLC0415"]
|
||||
# Lazy import of strix.config.models avoids a circular dependency between the
|
||||
# report pipeline and the config layer.
|
||||
"strix/report/dedupe.py" = ["PLC0415"]
|
||||
@@ -292,9 +289,6 @@ ignore = [
|
||||
# Heavy inference deps (httpx, openai) imported lazily so auth-status checks
|
||||
# don't pull them in.
|
||||
"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;
|
||||
# splitting would obscure the decision tree without simplifying it.
|
||||
"strix/interface/utils.py" = ["PLR0912", "BLE001", "PLC0415"]
|
||||
|
||||
+61
-18
@@ -16,6 +16,7 @@ import hashlib
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
import threading
|
||||
import time
|
||||
import urllib.parse
|
||||
from pathlib import Path
|
||||
@@ -23,7 +24,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
import requests
|
||||
|
||||
from strix.config import subscription_store
|
||||
from strix.utils.secret_files import write_secret_text
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -53,12 +54,50 @@ _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:
|
||||
write_secret_text(AUTH_PATH, json.dumps(data, indent=2))
|
||||
|
||||
|
||||
def read_provider_record(provider: str) -> dict[str, Any] | None:
|
||||
"""Raw record for *provider* from the shared subscription-auth store."""
|
||||
record = _read_store().get(provider)
|
||||
return record if isinstance(record, dict) else None
|
||||
|
||||
|
||||
def save_provider_record(provider: str, record: dict[str, Any]) -> None:
|
||||
data = _read_store()
|
||||
data[provider] = record
|
||||
_write_store(data)
|
||||
|
||||
|
||||
def remove_provider_record(provider: str) -> 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()
|
||||
|
||||
|
||||
def read_record() -> dict[str, Any] | None:
|
||||
record = subscription_store.read(AUTH_PATH).get(PROVIDER)
|
||||
record = read_provider_record(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")):
|
||||
@@ -71,31 +110,35 @@ def is_authenticated() -> bool:
|
||||
|
||||
|
||||
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)
|
||||
save_provider_record(PROVIDER, record)
|
||||
|
||||
|
||||
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()
|
||||
remove_provider_record(PROVIDER)
|
||||
|
||||
|
||||
@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
|
||||
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):
|
||||
|
||||
@@ -1,314 +0,0 @@
|
||||
"""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"
|
||||
+40
-13
@@ -37,7 +37,7 @@ from openai.types.responses import (
|
||||
from openai.types.responses.response_usage import ResponseUsage
|
||||
from openai.types.shared import Reasoning
|
||||
|
||||
from strix.config import codex, grok
|
||||
from strix.config import codex, opencode
|
||||
from strix.config.loader import load_settings
|
||||
from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input
|
||||
from strix.config.tool_call_limits import TurnToolCallLimiter
|
||||
@@ -80,7 +80,12 @@ def _retry_statusless_provider_errors(context: RetryPolicyContext) -> bool:
|
||||
|
||||
|
||||
class _CodexResponsesModel(OpenAIResponsesModel):
|
||||
"""Responses model for the ChatGPT subscription backend (always streamed, stateless)."""
|
||||
"""Responses model for stateless subscription gateways (always streamed).
|
||||
|
||||
Used for the ChatGPT subscription backend and for Responses-served models on
|
||||
the OpenCode gateway: neither stores responses server-side, so reasoning is
|
||||
carried inline via ``reasoning.encrypted_content``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -472,6 +477,7 @@ class StrixProvider(MultiProvider):
|
||||
def get_model(self, model_name: str | None) -> Model:
|
||||
llm = load_settings().llm
|
||||
slug = codex.subscription_model(model_name)
|
||||
oc = opencode.subscription_model(model_name)
|
||||
idle_timeout = float(llm.stream_idle_timeout)
|
||||
if slug:
|
||||
# The ChatGPT subscription backend is always streamed; it has no
|
||||
@@ -482,10 +488,19 @@ class StrixProvider(MultiProvider):
|
||||
codex.get_subscription_client(),
|
||||
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())
|
||||
elif oc and oc.uses_responses:
|
||||
model = _CodexResponsesModel(
|
||||
oc.slug,
|
||||
opencode.get_subscription_client(oc.base_url),
|
||||
reasoning_effort=llm.reasoning_effort,
|
||||
)
|
||||
elif oc:
|
||||
model = OpenAIChatCompletionsModel(
|
||||
oc.slug, opencode.get_subscription_client(oc.base_url)
|
||||
)
|
||||
if llm.disable_streaming:
|
||||
model = _NonStreamingModel(model)
|
||||
idle_timeout = 0.0
|
||||
else:
|
||||
model = super().get_model(model_name)
|
||||
if llm.disable_streaming:
|
||||
@@ -545,15 +560,24 @@ RECOMMENDED_MODEL_NAMES = (
|
||||
_RECOMMENDED_MODEL_NAME_SET = frozenset(name.lower() for name in RECOMMENDED_MODEL_NAMES)
|
||||
|
||||
FRONTIER_MODEL_FAMILIES = (
|
||||
(("azure", "azure_ai", "bedrock_mantle", "chatgpt", "openai"), ("gpt-5",)),
|
||||
(("azure", "azure_ai", "bedrock_mantle", "chatgpt", "openai", "opencode"), ("gpt-5",)),
|
||||
(
|
||||
("anthropic", "azure_ai", "bedrock", "claude", "databricks", "snowflake", "vertex_ai"),
|
||||
(
|
||||
"anthropic",
|
||||
"azure_ai",
|
||||
"bedrock",
|
||||
"claude",
|
||||
"databricks",
|
||||
"opencode",
|
||||
"snowflake",
|
||||
"vertex_ai",
|
||||
),
|
||||
("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.8", "qwen3.7", "qwen3-max")),
|
||||
(("moonshot", "moonshotai", "kimi"), ("kimi-k3", "kimi-k2.7", "kimi-k2.6")),
|
||||
(("google", "gemini", "opencode", "vertex_ai"), ("gemini-3",)),
|
||||
(("deepseek", "opencode"), ("deepseek-v4", "deepseek-r1", "deepseek-reasoner")),
|
||||
(("alibaba", "dashscope", "opencode", "qwen"), ("qwen3.8", "qwen3.7", "qwen3-max")),
|
||||
(("kimi", "moonshot", "moonshotai", "opencode"), ("kimi-k3", "kimi-k2.7", "kimi-k2.6")),
|
||||
)
|
||||
|
||||
|
||||
@@ -561,7 +585,7 @@ 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) or grok.subscription_model(llm.model):
|
||||
if codex.subscription_model(llm.model) or opencode.subscription_model(llm.model):
|
||||
return
|
||||
_configure_litellm_compatibility()
|
||||
_configure_openrouter_attribution(llm.model)
|
||||
@@ -746,6 +770,9 @@ def uses_chat_completions_tool_schema(model_name: str, settings: Settings) -> bo
|
||||
"""Return whether the resolved SDK route can only receive JSON function tools."""
|
||||
if codex.subscription_model(model_name):
|
||||
return False
|
||||
oc = opencode.subscription_model(model_name)
|
||||
if oc:
|
||||
return not oc.uses_responses
|
||||
model = model_name.strip().lower()
|
||||
if "/" in model and not model.startswith("openai/"):
|
||||
return True
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
"""OpenCode subscription auth: API-key sign-in and the OpenAI clients that
|
||||
route inference through the OpenCode gateway.
|
||||
|
||||
Covers both OpenCode offerings — Zen (pay-as-you-go credits) and Go (the
|
||||
monthly subscription) — which share one account and API key but live behind
|
||||
different gateway base URLs. Unlike the ChatGPT subscription there is no
|
||||
OAuth: the user copies a plain API key from https://opencode.ai/auth, and
|
||||
using the gateway from other agents is officially supported.
|
||||
|
||||
Model routing follows the endpoint each model is served on (see
|
||||
https://opencode.ai/docs/zen/): GPT models use the Responses API, everything
|
||||
else the OpenAI-compatible Chat Completions API.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
import requests
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
from strix.config import codex
|
||||
|
||||
|
||||
PROVIDER = "opencode"
|
||||
|
||||
ZEN_BASE_URL = "https://opencode.ai/zen/v1"
|
||||
GO_BASE_URL = "https://opencode.ai/zen/go/v1"
|
||||
|
||||
# ``opencode/<model>`` runs on Zen credits; ``opencode-go/<model>`` on the Go
|
||||
# subscription (matching OpenCode's own ``opencode-go/`` model ids).
|
||||
ZEN_PREFIX = "opencode/"
|
||||
GO_PREFIX = "opencode-go/"
|
||||
|
||||
AUTH_CONSOLE_URL = "https://opencode.ai/auth"
|
||||
|
||||
_KEY_CHECK_TIMEOUT = 30
|
||||
|
||||
|
||||
class OpencodeAuthError(Exception):
|
||||
def __init__(self, code: str, message: str | None = None) -> None:
|
||||
self.code = code
|
||||
super().__init__(message or code)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SubscriptionModel:
|
||||
slug: str
|
||||
base_url: str
|
||||
uses_responses: bool
|
||||
|
||||
|
||||
def _uses_responses(slug: str, base_url: str) -> bool:
|
||||
lowered = slug.lower()
|
||||
if lowered.startswith("gpt-"):
|
||||
return True
|
||||
# Grok is served via Responses on Zen but Chat Completions on Go.
|
||||
return lowered.startswith("grok") and base_url == ZEN_BASE_URL
|
||||
|
||||
|
||||
def subscription_model(model_name: str | None) -> SubscriptionModel | None:
|
||||
"""The gateway model behind an ``opencode/`` or ``opencode-go/`` STRIX_LLM."""
|
||||
name = (model_name or "").strip()
|
||||
lowered = name.lower()
|
||||
for prefix, base_url in ((GO_PREFIX, GO_BASE_URL), (ZEN_PREFIX, ZEN_BASE_URL)):
|
||||
if lowered.startswith(prefix):
|
||||
slug = name[len(prefix) :]
|
||||
if not slug:
|
||||
return None
|
||||
return SubscriptionModel(slug, base_url, _uses_responses(slug, base_url))
|
||||
return None
|
||||
|
||||
|
||||
def read_record() -> dict[str, Any] | None:
|
||||
record = codex.read_provider_record(PROVIDER)
|
||||
if not isinstance(record, dict) or record.get("type") != "api_key":
|
||||
return None
|
||||
key = record.get("key")
|
||||
if not isinstance(key, str) or not key:
|
||||
return None
|
||||
return record
|
||||
|
||||
|
||||
def is_authenticated() -> bool:
|
||||
return read_record() is not None
|
||||
|
||||
|
||||
def save_api_key(key: str) -> None:
|
||||
codex.save_provider_record(PROVIDER, {"type": "api_key", "provider": PROVIDER, "key": key})
|
||||
|
||||
|
||||
def logout() -> None:
|
||||
codex.remove_provider_record(PROVIDER)
|
||||
|
||||
|
||||
def get_api_key() -> str:
|
||||
record = read_record()
|
||||
if record is None:
|
||||
raise OpencodeAuthError(
|
||||
"not_authenticated", "not signed in; run: strix auth login opencode"
|
||||
)
|
||||
return str(record["key"])
|
||||
|
||||
|
||||
def validate_api_key(key: str) -> None:
|
||||
"""Check the key against the gateway's models endpoint; raise if rejected."""
|
||||
try:
|
||||
response = requests.get(
|
||||
f"{ZEN_BASE_URL}/models",
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
timeout=_KEY_CHECK_TIMEOUT,
|
||||
)
|
||||
except requests.RequestException as exc:
|
||||
raise OpencodeAuthError("unavailable", str(exc)) from exc
|
||||
if response.status_code in (401, 403):
|
||||
raise OpencodeAuthError(
|
||||
"invalid_key", f"OpenCode rejected the API key (HTTP {response.status_code})"
|
||||
)
|
||||
if response.status_code >= 400:
|
||||
raise OpencodeAuthError("http_error", f"HTTP {response.status_code}: {response.text[:300]}")
|
||||
|
||||
|
||||
def build_openai_client(base_url: str) -> AsyncOpenAI:
|
||||
return AsyncOpenAI(
|
||||
api_key=get_api_key(),
|
||||
base_url=base_url,
|
||||
http_client=httpx.AsyncClient(timeout=httpx.Timeout(600.0, connect=30.0)),
|
||||
)
|
||||
|
||||
|
||||
_subscription_clients: dict[str, AsyncOpenAI] = {}
|
||||
|
||||
|
||||
def get_subscription_client(base_url: str) -> AsyncOpenAI:
|
||||
client = _subscription_clients.get(base_url)
|
||||
if client is None:
|
||||
client = build_openai_client(base_url)
|
||||
_subscription_clients[base_url] = client
|
||||
return client
|
||||
|
||||
|
||||
def auth_mode(model_name: str | None) -> str:
|
||||
"""Return "subscription" when STRIX_LLM runs on any subscription
|
||||
(OpenCode or ChatGPT), else "api_key"."""
|
||||
if subscription_model(model_name) or codex.subscription_model(model_name):
|
||||
return "subscription"
|
||||
return "api_key"
|
||||
|
||||
|
||||
def subscription_provider(model_name: str | None) -> str | None:
|
||||
"""The subscription behind STRIX_LLM: "opencode", "chatgpt", or None."""
|
||||
if subscription_model(model_name):
|
||||
return PROVIDER
|
||||
if codex.subscription_model(model_name):
|
||||
return "chatgpt"
|
||||
return None
|
||||
@@ -1,63 +0,0 @@
|
||||
"""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)}"
|
||||
@@ -1,152 +0,0 @@
|
||||
"""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
|
||||
@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any
|
||||
from agents.model_settings import ModelSettings
|
||||
from openai.types.shared import Reasoning
|
||||
|
||||
from strix.config import opencode
|
||||
from strix.config.models import (
|
||||
DEFAULT_MODEL_RETRY,
|
||||
OPENROUTER_ATTRIBUTION_HEADERS,
|
||||
@@ -272,6 +273,10 @@ def _prompt_cache_extra_args(model_name: str) -> dict[str, Any] | None:
|
||||
"""
|
||||
if not is_claude_model(model_name):
|
||||
return None
|
||||
# OpenCode routes use the raw OpenAI SDK, which rejects this LiteLLM-only
|
||||
# argument; the gateway applies Anthropic prompt caching itself.
|
||||
if opencode.subscription_model(model_name):
|
||||
return None
|
||||
if is_bedrock_route(model_name) and not bedrock_route_supports_prompt_caching(model_name):
|
||||
return None
|
||||
|
||||
|
||||
+148
-161
@@ -1,8 +1,9 @@
|
||||
"""`strix auth` — model-subscription sign-in (login / status / logout).
|
||||
"""`strix auth` — 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 a
|
||||
ChatGPT subscription and a ``grok/<model>`` one on a Grok/SuperGrok subscription.
|
||||
Signing in only stores credentials (``~/.strix/subscription-auth.json``); model
|
||||
selection stays with ``STRIX_LLM``. A ``chatgpt/<model>`` STRIX_LLM runs on the
|
||||
ChatGPT subscription; ``opencode/<model>`` (Zen credits) or
|
||||
``opencode-go/<model>`` (Go subscription) run on OpenCode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -12,7 +13,6 @@ import base64
|
||||
import logging
|
||||
import threading
|
||||
import webbrowser
|
||||
from dataclasses import dataclass
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
@@ -22,78 +22,33 @@ from rich.console import Console
|
||||
from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
|
||||
from strix.config import codex, grok, load_settings, subscription_store
|
||||
from strix.config import codex, load_settings, opencode
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
from types import ModuleType
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_CALLBACK_TIMEOUT_S = 300
|
||||
|
||||
|
||||
@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"
|
||||
# CLI-facing name for the default 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})
|
||||
_OPENCODE_PROVIDERS = frozenset({opencode.PROVIDER, "opencode-go", "zen"})
|
||||
|
||||
_USAGE = (
|
||||
"Usage:\n"
|
||||
" strix auth login [chatgpt|grok] [--manual]\n"
|
||||
" strix auth login chatgpt [--manual]\n"
|
||||
" strix auth login opencode\n"
|
||||
" strix auth status\n"
|
||||
" strix auth logout [chatgpt|grok]"
|
||||
" strix auth logout [chatgpt|opencode]"
|
||||
)
|
||||
|
||||
|
||||
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:
|
||||
"""Entry point for ``strix auth …``. Returns a process exit code."""
|
||||
console = Console()
|
||||
@@ -102,7 +57,7 @@ def run_auth(argv: list[str]) -> int:
|
||||
rest = argv[1:]
|
||||
|
||||
if subcommand in ("-h", "--help", "help"):
|
||||
console.print(_USAGE, markup=False)
|
||||
console.print(_USAGE)
|
||||
return 0
|
||||
|
||||
handlers: dict[str, Callable[[], int]] = {
|
||||
@@ -115,7 +70,7 @@ def run_auth(argv: list[str]) -> int:
|
||||
return handler()
|
||||
|
||||
console.print(f"[red]Unknown auth command:[/] {subcommand}\n")
|
||||
console.print(_USAGE, markup=False)
|
||||
console.print(_USAGE)
|
||||
return 2
|
||||
|
||||
|
||||
@@ -124,8 +79,8 @@ def _login(console: Console, argv: list[str]) -> int:
|
||||
parser.add_argument(
|
||||
"provider",
|
||||
nargs="?",
|
||||
default=_DEFAULT_PROVIDER,
|
||||
help="Model provider to sign in with (chatgpt or grok; default: chatgpt).",
|
||||
default=LOGIN_PROVIDER,
|
||||
help="Model provider to sign in with (default: chatgpt).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--manual",
|
||||
@@ -137,42 +92,100 @@ def _login(console: Console, argv: list[str]) -> int:
|
||||
except SystemExit as exc: # argparse already printed the message
|
||||
return int(exc.code or 2)
|
||||
|
||||
provider = _resolve_provider(args.provider)
|
||||
if provider is None:
|
||||
supported = ", ".join(f"'{name}'" for name in _PROVIDERS)
|
||||
console.print(f"[red]Unsupported provider:[/] {args.provider}. Supported: {supported}.")
|
||||
if args.provider.lower() in _OPENCODE_PROVIDERS:
|
||||
return _login_opencode(console)
|
||||
|
||||
if args.provider.lower() not in _ACCEPTED_PROVIDERS:
|
||||
console.print(
|
||||
f"[red]Unsupported provider:[/] {args.provider}. "
|
||||
f"Supported: '{LOGIN_PROVIDER}' (ChatGPT subscription) and "
|
||||
f"'{opencode.PROVIDER}' (OpenCode Zen/Go)."
|
||||
)
|
||||
return 2
|
||||
|
||||
module = provider.module
|
||||
verifier, challenge = module.generate_pkce()
|
||||
state = module.create_state()
|
||||
authorize_url = module.build_authorize_url(challenge, state)
|
||||
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(
|
||||
f"[bold]Signing in with {provider.display}[/] [dim](provider: {provider.name})[/]"
|
||||
"[dim]This uses your ChatGPT Plus/Pro plan for inference instead of a metered API key.[/]"
|
||||
)
|
||||
console.print(f"[dim]{provider.blurb}[/]")
|
||||
console.print()
|
||||
|
||||
try:
|
||||
record = _run_oauth_flow(
|
||||
console, provider, authorize_url, verifier, state, manual=args.manual
|
||||
)
|
||||
except provider.error as exc:
|
||||
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
|
||||
|
||||
module.save_record(record)
|
||||
_print_success(console, provider)
|
||||
codex.save_record(record)
|
||||
_print_success(console)
|
||||
return 0
|
||||
|
||||
|
||||
def _login_opencode(console: Console) -> int:
|
||||
console.print()
|
||||
console.print("[bold]Signing in with OpenCode[/] [dim](provider: opencode)[/]")
|
||||
console.print(
|
||||
"[dim]This uses your OpenCode Zen credits or Go subscription for inference.\n"
|
||||
f"Get your API key at {opencode.AUTH_CONSOLE_URL}[/]"
|
||||
)
|
||||
console.print()
|
||||
try:
|
||||
key = console.input("Paste your OpenCode API key: ", password=True).strip()
|
||||
except (EOFError, KeyboardInterrupt):
|
||||
console.print("\n[yellow]Sign-in cancelled.[/]")
|
||||
return 130
|
||||
if not key:
|
||||
console.print("[red]No API key provided.[/]")
|
||||
return 2
|
||||
try:
|
||||
opencode.validate_api_key(key)
|
||||
except opencode.OpencodeAuthError as exc:
|
||||
console.print(f"[red]SIGN-IN FAILED:[/] {exc}")
|
||||
return 1
|
||||
opencode.save_api_key(key)
|
||||
_print_opencode_success(console)
|
||||
return 0
|
||||
|
||||
|
||||
def _print_opencode_success(console: Console) -> None:
|
||||
text = Text()
|
||||
text.append("Signed in with your OpenCode account", style="bold #22c55e")
|
||||
text.append("\n\n", style="white")
|
||||
text.append("Set ", style="white")
|
||||
text.append("STRIX_LLM", style="bold white")
|
||||
text.append(" to an ", style="white")
|
||||
text.append("opencode/", style="bold cyan")
|
||||
text.append(" model (e.g. ", style="white")
|
||||
text.append("opencode/claude-sonnet-5", style="bold cyan")
|
||||
text.append(") to run on Zen credits, or ", style="white")
|
||||
text.append("opencode-go/", style="bold cyan")
|
||||
text.append(" (e.g. ", style="white")
|
||||
text.append("opencode-go/kimi-k3", style="bold cyan")
|
||||
text.append(") to run on the Go subscription.", 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()
|
||||
|
||||
|
||||
def _run_oauth_flow(
|
||||
console: Console,
|
||||
provider: _Provider,
|
||||
authorize_url: str,
|
||||
verifier: str,
|
||||
state: str,
|
||||
@@ -180,10 +193,7 @@ def _run_oauth_flow(
|
||||
manual: bool,
|
||||
) -> dict[str, Any]:
|
||||
"""Drive the browser (or manual) OAuth flow and return a token record."""
|
||||
module = provider.module
|
||||
server = (
|
||||
None if manual else _try_start_callback_server(module.CALLBACK_PORT, module.CALLBACK_PATH)
|
||||
)
|
||||
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}[/]")
|
||||
@@ -201,8 +211,8 @@ def _run_oauth_flow(
|
||||
if result is not None:
|
||||
code, returned_state, error = result
|
||||
if error:
|
||||
raise provider.error("oauth_error", error)
|
||||
return _finish(provider, code, returned_state, verifier, state, require_state=True)
|
||||
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
|
||||
@@ -212,13 +222,12 @@ def _run_oauth_flow(
|
||||
try:
|
||||
pasted = console.input("Paste the full redirect URL (or code#state): ").strip()
|
||||
except EOFError as exc:
|
||||
raise provider.error("no_input", "no redirect URL provided") from exc
|
||||
code, returned_state = module.parse_redirect_input(pasted)
|
||||
return _finish(provider, code, returned_state, verifier, state, require_state=False)
|
||||
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(
|
||||
provider: _Provider,
|
||||
code: str | None,
|
||||
returned_state: str | None,
|
||||
verifier: str,
|
||||
@@ -227,17 +236,16 @@ def _finish(
|
||||
require_state: bool,
|
||||
) -> dict[str, Any]:
|
||||
if not code:
|
||||
raise provider.error("no_code", "no authorization code found in the redirect")
|
||||
# The loopback callback from the provider 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.
|
||||
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 provider.error("state_mismatch", "missing state in callback; possible CSRF")
|
||||
raise codex.CodexAuthError("state_mismatch", "missing state in callback; possible CSRF")
|
||||
if returned_state is not None and returned_state != expected_state:
|
||||
raise provider.error("state_mismatch", "state did not match; possible CSRF")
|
||||
record: dict[str, Any] = provider.module.exchange_code(code, verifier)
|
||||
return record
|
||||
raise codex.CodexAuthError("state_mismatch", "state did not match; possible CSRF")
|
||||
return codex.exchange_code(code, verifier)
|
||||
|
||||
|
||||
class _CallbackServer:
|
||||
@@ -264,7 +272,7 @@ class _CallbackServer:
|
||||
self._httpd.server_close()
|
||||
|
||||
|
||||
def _try_start_callback_server(port: int, path: str) -> _CallbackServer | None:
|
||||
def _try_start_callback_server() -> _CallbackServer | None:
|
||||
event = threading.Event()
|
||||
holder: dict[str, Any] = {}
|
||||
|
||||
@@ -274,7 +282,7 @@ def _try_start_callback_server(port: int, path: str) -> _CallbackServer | None:
|
||||
|
||||
def do_GET(self) -> None:
|
||||
parsed = urlparse(self.path)
|
||||
if parsed.path != path:
|
||||
if parsed.path != codex.CALLBACK_PATH:
|
||||
self.send_response(404)
|
||||
self.end_headers()
|
||||
return
|
||||
@@ -291,9 +299,9 @@ def _try_start_callback_server(port: int, path: str) -> _CallbackServer | None:
|
||||
event.set()
|
||||
|
||||
try:
|
||||
httpd = HTTPServer(("127.0.0.1", port), Handler)
|
||||
httpd = HTTPServer(("127.0.0.1", codex.CALLBACK_PORT), Handler)
|
||||
except OSError:
|
||||
logger.debug("could not bind callback port %d", port, exc_info=True)
|
||||
logger.debug("could not bind callback port %d", codex.CALLBACK_PORT, exc_info=True)
|
||||
return None
|
||||
return _CallbackServer(httpd, event, holder)
|
||||
|
||||
@@ -304,67 +312,47 @@ def _first(query: dict[str, list[str]], key: str) -> str | None:
|
||||
|
||||
|
||||
def _status(console: Console) -> int:
|
||||
settings = load_settings()
|
||||
active_model = settings.llm.model
|
||||
signed_in_any = False
|
||||
for provider in _PROVIDERS.values():
|
||||
record = provider.module.read_record()
|
||||
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:
|
||||
record = codex.read_record()
|
||||
opencode_signed_in = opencode.is_authenticated()
|
||||
if record is None and not opencode_signed_in:
|
||||
console.print(
|
||||
"[yellow]Not signed in.[/] Run [cyan]strix auth login chatgpt[/] "
|
||||
"or [cyan]strix auth login grok[/] to sign in."
|
||||
"[yellow]Not signed in.[/] Run [cyan]strix auth login chatgpt[/] or "
|
||||
"[cyan]strix auth login opencode[/] to sign in."
|
||||
)
|
||||
return 1
|
||||
settings = load_settings()
|
||||
if record is not None:
|
||||
console.print("[green]Signed in[/] with a ChatGPT subscription.")
|
||||
console.print(f" Account: [bold]{record.get('account_id')}[/]")
|
||||
if opencode_signed_in:
|
||||
console.print("[green]Signed in[/] with an OpenCode account.")
|
||||
if codex.subscription_model(settings.llm.model) or opencode.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[/] or "
|
||||
"[cyan]opencode/claude-sonnet-5[/] to run on a subscription."
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
def _logout(console: Console, argv: list[str]) -> int:
|
||||
parser = argparse.ArgumentParser(prog="strix auth logout", add_help=True)
|
||||
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}.")
|
||||
def _logout(console: Console, argv: list[str] | None = None) -> int:
|
||||
target = (argv[0].lower() if argv else "") or "all"
|
||||
if target in _ACCEPTED_PROVIDERS or target == "all":
|
||||
codex.logout()
|
||||
if target in _OPENCODE_PROVIDERS or target == "all":
|
||||
opencode.logout()
|
||||
if target != "all" and target not in _ACCEPTED_PROVIDERS | _OPENCODE_PROVIDERS:
|
||||
console.print(f"[red]Unknown provider:[/] {target}\n")
|
||||
console.print(_USAGE)
|
||||
return 2
|
||||
target.module.logout()
|
||||
console.print(f"[green]Signed out of {target.display}.[/] Stored credentials removed.")
|
||||
console.print("[green]Signed out.[/] Stored subscription credentials removed.")
|
||||
return 0
|
||||
|
||||
|
||||
def _fail(console: Console, exc: Exception) -> int:
|
||||
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")
|
||||
@@ -382,18 +370,17 @@ def _fail(console: Console, exc: Exception) -> int:
|
||||
return 1
|
||||
|
||||
|
||||
def _print_success(console: Console, provider: _Provider) -> None:
|
||||
prefix = provider.module.SUBSCRIPTION_PREFIX
|
||||
def _print_success(console: Console) -> None:
|
||||
text = Text()
|
||||
text.append(f"Signed in with your {provider.display} subscription", style="bold #22c55e")
|
||||
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(prefix, style="bold cyan")
|
||||
text.append("chatgpt/", style="bold cyan")
|
||||
text.append(" model (e.g. ", style="white")
|
||||
text.append(provider.example_model, style="bold cyan")
|
||||
text.append(f") — runs are billed to your {provider.display} plan.", 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")
|
||||
|
||||
@@ -8,7 +8,7 @@ from rich.console import Console
|
||||
from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
|
||||
from strix.config import codex, grok, load_settings
|
||||
from strix.config import codex, load_settings, opencode
|
||||
from strix.interface.utils import (
|
||||
check_docker_connection,
|
||||
image_exists,
|
||||
@@ -37,14 +37,14 @@ def validate_environment() -> None:
|
||||
logger.info("Environment OK (ChatGPT subscription)")
|
||||
return
|
||||
|
||||
if grok.subscription_model(settings.llm.model):
|
||||
if not grok.is_authenticated():
|
||||
if opencode.subscription_model(settings.llm.model):
|
||||
if not opencode.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."
|
||||
f"[red]STRIX_LLM={settings.llm.model} uses your OpenCode subscription, "
|
||||
"but you're not signed in.[/] Run [cyan]strix auth login opencode[/] first."
|
||||
)
|
||||
sys.exit(1)
|
||||
logger.info("Environment OK (Grok subscription)")
|
||||
logger.info("Environment OK (OpenCode subscription)")
|
||||
return
|
||||
|
||||
if not settings.llm.model:
|
||||
|
||||
+10
-13
@@ -14,7 +14,7 @@ from rich.console import Console
|
||||
from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
|
||||
from strix.config import codex, load_settings, persist_current
|
||||
from strix.config import codex, load_settings, opencode, persist_current
|
||||
from strix.core.paths import run_dir_for
|
||||
from strix.interface.cli_args import parse_arguments
|
||||
from strix.interface.environment import (
|
||||
@@ -104,8 +104,14 @@ def _provider_import_hint(exc: BaseException, model: str) -> str | 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 an actionable hint for a known subscription error, or None."""
|
||||
model = load_settings().llm.model
|
||||
if opencode.subscription_model(model):
|
||||
joined = " ".join(_exception_messages(exc)).lower()
|
||||
if "error code: 401" in joined or "http 401" in joined or "unauthorized" in joined:
|
||||
return "Your OpenCode API key was rejected. Sign in again:\n strix auth login opencode"
|
||||
return None
|
||||
if not codex.subscription_model(model):
|
||||
return None
|
||||
joined = " ".join(_exception_messages(exc)).lower()
|
||||
if "not supported when using codex with a chatgpt account" in joined:
|
||||
@@ -436,16 +442,7 @@ def main() -> None:
|
||||
start_background_check()
|
||||
if not args.non_interactive and prompt_update_if_available(Console()):
|
||||
if is_binary_install() and sys.platform != "win32":
|
||||
# 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
|
||||
os.execv(sys.executable, sys.argv) # noqa: S606 # nosec B606
|
||||
sys.exit(0)
|
||||
|
||||
check_docker_installed()
|
||||
|
||||
@@ -14,7 +14,7 @@ import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from strix.config import Settings, load_settings, subscription
|
||||
from strix.config import Settings, load_settings, opencode
|
||||
from strix.core.paths import run_dir_for
|
||||
from strix.interface.utils import (
|
||||
assign_workspace_subdirs,
|
||||
@@ -226,7 +226,7 @@ def telemetry_start(args: argparse.Namespace) -> None:
|
||||
model = load_settings().llm.model
|
||||
kwargs = {
|
||||
"model": model,
|
||||
"auth_mode": subscription.auth_mode(model),
|
||||
"auth_mode": opencode.auth_mode(model),
|
||||
"scan_mode": args.scan_mode,
|
||||
"is_whitebox": is_whitebox_scan(args.targets_info),
|
||||
"interactive": not args.non_interactive,
|
||||
@@ -241,15 +241,14 @@ def _persist_run_record(args: argparse.Namespace) -> None:
|
||||
|
||||
run_dir = run_dir_for(args.run_name)
|
||||
run_dir.mkdir(parents=True, exist_ok=True)
|
||||
model = load_settings().llm.model
|
||||
run_record = {
|
||||
"run_id": args.run_name,
|
||||
"run_name": args.run_name,
|
||||
"status": "running",
|
||||
"start_time": datetime.now(UTC).isoformat(),
|
||||
"end_time": None,
|
||||
"auth_mode": subscription.auth_mode(model),
|
||||
"subscription_provider": subscription.provider_label(model),
|
||||
"auth_mode": opencode.auth_mode(load_settings().llm.model),
|
||||
"subscription_provider": opencode.subscription_provider(load_settings().llm.model),
|
||||
"targets_info": args.targets_info,
|
||||
"scan_mode": args.scan_mode,
|
||||
"instruction": args.instruction,
|
||||
|
||||
@@ -162,11 +162,12 @@ class TuiController:
|
||||
if self.report_state is not None:
|
||||
usage = dict(self.report_state.get_total_llm_usage())
|
||||
subscription = False
|
||||
subscription_name = ""
|
||||
with contextlib.suppress(Exception):
|
||||
subscription = is_subscription_run(self.report_state)
|
||||
if subscription:
|
||||
subscription_name = subscription_label(self.report_state)
|
||||
label = ""
|
||||
if subscription:
|
||||
with contextlib.suppress(Exception):
|
||||
label = subscription_label()
|
||||
model_warning = ""
|
||||
if model and not is_recommended_or_frontier_model(model):
|
||||
model_warning = (
|
||||
@@ -203,7 +204,7 @@ class TuiController:
|
||||
],
|
||||
"usage": terminal_projection(usage, max_string=256, max_items=20),
|
||||
"subscription": subscription,
|
||||
"subscription_label": terminal_projection(subscription_name, max_string=64),
|
||||
"subscription_label": label,
|
||||
"viewer_status": self.viewer_status,
|
||||
"viewer_url": terminal_projection(self.viewer_url, max_string=1024),
|
||||
"error": terminal_projection(self.error, max_string=2 * 1024),
|
||||
|
||||
@@ -175,7 +175,6 @@ def bounded_state_projection(state: dict[str, Any]) -> dict[str, Any]:
|
||||
"messages": [],
|
||||
"usage": {},
|
||||
"subscription": state["subscription"],
|
||||
"subscription_label": state["subscription_label"],
|
||||
"viewer_status": state["viewer_status"],
|
||||
"viewer_url": None,
|
||||
"error": terminal_projection(state["error"], max_string=256),
|
||||
|
||||
@@ -1067,12 +1067,11 @@ func TestBudgetPauseShowsOneWarningToastUntilResumed(t *testing.T) {
|
||||
|
||||
func TestStatsViewShowsSubscription(t *testing.T) {
|
||||
model := New(nil)
|
||||
model.snapshot.Model = "grok/grok-4"
|
||||
model.snapshot.Model = "gpt-5"
|
||||
model.snapshot.Subscription = true
|
||||
model.snapshot.SubscriptionLabel = "Grok subscription"
|
||||
model.snapshot.Usage = map[string]any{"total_tokens": float64(1200), "cost": 3.5}
|
||||
stats := ansi.Strip(model.statsView())
|
||||
if !strings.Contains(stats, "Grok subscription") {
|
||||
if !strings.Contains(stats, "ChatGPT subscription") {
|
||||
t.Fatalf("stats missing subscription line: %q", stats)
|
||||
}
|
||||
if strings.Contains(stats, "$") {
|
||||
@@ -1080,16 +1079,6 @@ 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) {
|
||||
report := vulnerabilityMarkdownReport(map[string]any{
|
||||
"title": "SQLi in login",
|
||||
|
||||
@@ -598,7 +598,7 @@ func (m Model) statsView() string {
|
||||
}
|
||||
label := m.snapshot.SubscriptionLabel
|
||||
if label == "" {
|
||||
label = "Subscription"
|
||||
label = "ChatGPT subscription"
|
||||
}
|
||||
b.WriteString(lipgloss.NewStyle().Foreground(green).Render(label))
|
||||
}
|
||||
|
||||
+11
-19
@@ -262,27 +262,19 @@ def is_subscription_run(report_state: Any) -> bool:
|
||||
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 subscription
|
||||
from strix.config import opencode
|
||||
|
||||
return subscription.auth_mode(load_settings().llm.model) == "subscription"
|
||||
return opencode.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").
|
||||
def subscription_label() -> str:
|
||||
"""Display name of the subscription behind the configured model."""
|
||||
from strix.config import opencode
|
||||
|
||||
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"
|
||||
model = load_settings().llm.model
|
||||
if opencode.subscription_model(model):
|
||||
return "OpenCode subscription"
|
||||
return "ChatGPT subscription"
|
||||
|
||||
|
||||
def _int_stat(usage: dict[str, Any], key: str) -> int:
|
||||
@@ -386,7 +378,7 @@ def build_live_stats_text(report_state: Any) -> Text:
|
||||
stats_text.append(str(model), style="white")
|
||||
if is_subscription_run(report_state):
|
||||
stats_text.append(" · ", style="dim white")
|
||||
stats_text.append(subscription_label(report_state), style="#22c55e")
|
||||
stats_text.append(subscription_label(), style="#22c55e")
|
||||
stats_text.append("\n")
|
||||
|
||||
vuln_count = len(report_state.vulnerability_reports)
|
||||
@@ -432,7 +424,7 @@ def build_tui_stats_text(report_state: Any) -> Text:
|
||||
subscription = is_subscription_run(report_state)
|
||||
if subscription:
|
||||
stats_text.append("\n")
|
||||
stats_text.append(subscription_label(report_state), style="#22c55e")
|
||||
stats_text.append(subscription_label(), style="#22c55e")
|
||||
|
||||
usage = _llm_usage(report_state)
|
||||
if usage and _int_stat(usage, "total_tokens") > 0:
|
||||
|
||||
@@ -101,7 +101,11 @@ export function RunDetails({
|
||||
const totalTokens = num(usage.total_tokens);
|
||||
const cost = num(usage.cost);
|
||||
const subscription = str(raw.auth_mode) === "subscription";
|
||||
const subscriptionProvider = str(raw.subscription_provider);
|
||||
const subscriptionProvider =
|
||||
str(raw.subscription_provider) ??
|
||||
(models.some((m) => m.toLowerCase().startsWith("opencode")) ? "opencode" : "chatgpt");
|
||||
const subscriptionLabel =
|
||||
subscriptionProvider === "opencode" ? "OpenCode subscription" : "ChatGPT subscription";
|
||||
|
||||
const sub = (n: number, word: string) => (
|
||||
<span className="text-[#666]"> ({formatNumber(n)} {word})</span>
|
||||
@@ -181,7 +185,7 @@ export function RunDetails({
|
||||
<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]">
|
||||
{subscriptionProvider ? `${subscriptionProvider} subscription` : "Subscription"}
|
||||
{subscriptionLabel}
|
||||
</span>
|
||||
</span>
|
||||
</Field>
|
||||
|
||||
+22
-22
File diff suppressed because one or more lines are too long
@@ -6,7 +6,7 @@
|
||||
<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-XDX3roAH.js"></script>
|
||||
<script type="module" crossorigin src="./assets/index-1LIW3rcB.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="./assets/index-DKbLYAbP.css">
|
||||
</head>
|
||||
<body>
|
||||
|
||||
@@ -6,7 +6,6 @@ import json
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from strix.config import subscription
|
||||
from strix.core.paths import run_record_path
|
||||
from strix.interface.tui.live_view import TuiLiveView
|
||||
|
||||
@@ -58,39 +57,7 @@ def read_run_summary(run_dir: Path) -> dict[str, Any]:
|
||||
record = {}
|
||||
status = record.get("status")
|
||||
finished = status in _TERMINAL_STATUSES and bool(record.get("end_time"))
|
||||
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
|
||||
return {**record, "finished": finished}
|
||||
|
||||
|
||||
def primary_target(record: dict[str, Any]) -> str | None:
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import Any
|
||||
|
||||
import litellm
|
||||
|
||||
from strix.config import load_settings, subscription
|
||||
from strix.config import load_settings
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -19,6 +19,9 @@ logger = logging.getLogger(__name__)
|
||||
# ``litellm/``, ``ollama/`` ...). Strip a leading provider segment on lookup.
|
||||
_STRIPPABLE_PREFIXES = (
|
||||
"openai/",
|
||||
"chatgpt/",
|
||||
"opencode-go/",
|
||||
"opencode/",
|
||||
"litellm/",
|
||||
"any-llm/",
|
||||
"ollama/",
|
||||
@@ -29,8 +32,6 @@ _DEFAULT_OUTPUT_TOKENS = 8_192
|
||||
|
||||
|
||||
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:
|
||||
if model.startswith(prefix):
|
||||
return model[len(prefix) :]
|
||||
@@ -47,11 +48,12 @@ def _safe_get_model_info(model: str) -> dict[str, Any] | None:
|
||||
@lru_cache(maxsize=128)
|
||||
def _model_info(model: str) -> dict[str, int]:
|
||||
lookup_key = _lookup_key(model)
|
||||
# Subscription prefixes are never LiteLLM keys, and a provider-qualified
|
||||
# ChatGPT lookup may start a synchronous device-login poll: only ask about
|
||||
# the resolved name.
|
||||
# Provider-qualified ChatGPT lookups may start a synchronous device-login
|
||||
# poll. LiteLLM keys the metadata by the underlying model slug.
|
||||
candidates = (
|
||||
(lookup_key,) if subscription.provider_for_model(model) is not None else (model, lookup_key)
|
||||
(lookup_key,)
|
||||
if model.startswith(("chatgpt/", "opencode/", "opencode-go/"))
|
||||
else (model, lookup_key)
|
||||
)
|
||||
for candidate in candidates:
|
||||
info = _safe_get_model_info(candidate)
|
||||
|
||||
@@ -11,7 +11,7 @@ from uuid import uuid4
|
||||
|
||||
from agents.usage import Usage
|
||||
|
||||
from strix.config import subscription
|
||||
from strix.config import opencode
|
||||
from strix.config.loader import load_settings
|
||||
from strix.core.paths import run_dir_for
|
||||
from strix.report.pricing import resolve_litellm_model
|
||||
@@ -123,8 +123,7 @@ class ReportState:
|
||||
self.scan_results: dict[str, Any] | None = None
|
||||
self.scan_config: dict[str, Any] | None = None
|
||||
self._llm_usage = LLMUsageLedger()
|
||||
model = load_settings().llm.model
|
||||
auth_mode = subscription.auth_mode(model)
|
||||
auth_mode = opencode.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,
|
||||
@@ -133,7 +132,7 @@ class ReportState:
|
||||
"end_time": None,
|
||||
"status": "running",
|
||||
"auth_mode": auth_mode,
|
||||
"subscription_provider": subscription.provider_label(model),
|
||||
"subscription_provider": opencode.subscription_provider(load_settings().llm.model),
|
||||
"targets_info": [],
|
||||
"llm_usage": self._build_llm_usage_record(),
|
||||
}
|
||||
|
||||
+74
-52
@@ -2,38 +2,28 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.config import codex, grok
|
||||
from strix.config import codex, opencode
|
||||
from strix.interface import auth_cli
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
_CHATGPT = auth_cli._PROVIDERS["chatgpt"]
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
store = tmp_path / "home" / ".strix" / "subscription-auth.json"
|
||||
monkeypatch.setattr(codex, "AUTH_PATH", store)
|
||||
monkeypatch.setattr(grok, "AUTH_PATH", store)
|
||||
monkeypatch.setattr(codex, "AUTH_PATH", tmp_path / "home" / ".strix" / "subscription-auth.json")
|
||||
|
||||
|
||||
def test_default_provider_is_chatgpt() -> None:
|
||||
assert auth_cli._DEFAULT_PROVIDER == "chatgpt"
|
||||
assert set(auth_cli._PROVIDERS) == {"chatgpt", "grok"}
|
||||
|
||||
|
||||
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_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:
|
||||
@@ -62,32 +52,32 @@ def test_finish_requires_state_on_loopback(monkeypatch: pytest.MonkeyPatch) -> N
|
||||
|
||||
# Loopback (require_state=True): missing or mismatched state is rejected.
|
||||
with pytest.raises(codex.CodexAuthError) as missing:
|
||||
auth_cli._finish(_CHATGPT, "code", None, "verifier", "expected", require_state=True)
|
||||
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(_CHATGPT, "code", "wrong", "verifier", "expected", require_state=True)
|
||||
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(
|
||||
_CHATGPT, "code", "expected", "verifier", "expected", require_state=True
|
||||
) == {"ok": True}
|
||||
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(
|
||||
_CHATGPT, "code", None, "verifier", "expected", require_state=False
|
||||
) == {"ok": True}
|
||||
assert auth_cli._finish("code", None, "verifier", "expected", require_state=False) == {
|
||||
"ok": True
|
||||
}
|
||||
with pytest.raises(codex.CodexAuthError):
|
||||
auth_cli._finish(_CHATGPT, "code", "wrong", "verifier", "expected", require_state=False)
|
||||
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(_CHATGPT, None, "expected", "verifier", "expected", require_state=True)
|
||||
auth_cli._finish(None, "expected", "verifier", "expected", require_state=True)
|
||||
assert exc.value.code == "no_code"
|
||||
|
||||
|
||||
@@ -95,31 +85,6 @@ def test_model_subcommand_removed() -> None:
|
||||
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"])
|
||||
def test_login_accepts_provider_aliases(provider: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
reached = {"flow": False}
|
||||
@@ -140,3 +105,60 @@ def test_login_accepts_provider_aliases(provider: str, monkeypatch: pytest.Monke
|
||||
|
||||
assert auth_cli.run_auth(["login", provider]) == 0
|
||||
assert reached["flow"] is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["opencode", "OpenCode", "opencode-go", "zen"])
|
||||
def test_login_accepts_opencode_aliases(provider: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
reached = {"login": False}
|
||||
|
||||
def _fake_login(_console: Any) -> int:
|
||||
reached["login"] = True
|
||||
return 0
|
||||
|
||||
monkeypatch.setattr(auth_cli, "_login_opencode", _fake_login)
|
||||
assert auth_cli.run_auth(["login", provider]) == 0
|
||||
assert reached["login"] is True
|
||||
|
||||
|
||||
def test_login_opencode_validates_and_saves(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
saved: dict[str, str] = {}
|
||||
monkeypatch.setattr("rich.console.Console.input", lambda _self, *_a, **_k: " sk-oc-test ")
|
||||
monkeypatch.setattr(opencode, "validate_api_key", lambda key: saved.setdefault("checked", key))
|
||||
monkeypatch.setattr(opencode, "save_api_key", lambda key: saved.setdefault("key", key))
|
||||
|
||||
assert auth_cli.run_auth(["login", "opencode"]) == 0
|
||||
assert saved == {"checked": "sk-oc-test", "key": "sk-oc-test"}
|
||||
|
||||
|
||||
def test_login_opencode_rejects_bad_key(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr("rich.console.Console.input", lambda _self, *_a, **_k: "bad")
|
||||
|
||||
def _reject(_key: str) -> None:
|
||||
raise opencode.OpencodeAuthError("invalid_key")
|
||||
|
||||
monkeypatch.setattr(opencode, "validate_api_key", _reject)
|
||||
assert auth_cli.run_auth(["login", "opencode"]) == 1
|
||||
assert opencode.is_authenticated() is False
|
||||
|
||||
|
||||
def test_logout_provider_scoped() -> None:
|
||||
codex.save_record(
|
||||
{
|
||||
"type": "oauth",
|
||||
"provider": "codex",
|
||||
"access": "a",
|
||||
"refresh": "r",
|
||||
"account_id": "acct",
|
||||
"expires_at": time.time() + 3600,
|
||||
}
|
||||
)
|
||||
opencode.save_api_key("sk-oc-test")
|
||||
|
||||
assert auth_cli.run_auth(["logout", "opencode"]) == 0
|
||||
assert opencode.is_authenticated() is False
|
||||
assert codex.is_authenticated() is True
|
||||
|
||||
assert auth_cli.run_auth(["logout"]) == 0
|
||||
assert codex.is_authenticated() is False
|
||||
|
||||
assert auth_cli.run_auth(["logout", "bogus"]) == 2
|
||||
|
||||
@@ -39,17 +39,6 @@ def test_context_window_chatgpt_prefix_skips_provider_auth(
|
||||
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:
|
||||
context_budget._model_info.cache_clear()
|
||||
|
||||
|
||||
@@ -1,267 +0,0 @@
|
||||
"""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"
|
||||
@@ -1,113 +0,0 @@
|
||||
"""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"
|
||||
@@ -111,6 +111,13 @@ def test_make_model_settings_no_prompt_cache_for_non_claude(model_name: str) ->
|
||||
assert make_model_settings(None, model_name=model_name).extra_args is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_name", ["opencode/claude-sonnet-5", "opencode-go/claude-sonnet-5"])
|
||||
def test_no_prompt_cache_for_opencode_claude(model_name: str) -> None:
|
||||
# The OpenCode route uses the raw OpenAI SDK, whose create() rejects the
|
||||
# LiteLLM-only cache_control_injection_points argument.
|
||||
assert _cache_points(model_name) is None
|
||||
|
||||
|
||||
def test_no_prompt_cache_for_unmapped_bedrock_claude_model(monkeypatch: Any) -> None:
|
||||
# A Bedrock Claude model LiteLLM hasn't mapped must run uncached, not crash.
|
||||
unmapped = "bedrock/global.anthropic.claude-brand-new-9"
|
||||
|
||||
@@ -66,6 +66,11 @@ def test_recommended_models_are_matched_case_insensitively() -> None:
|
||||
"moonshot/kimi-k2.6",
|
||||
"kimi-k2.7-code",
|
||||
"moonshot/kimi-k3",
|
||||
"opencode/gpt-5.4",
|
||||
"opencode/claude-sonnet-5",
|
||||
"opencode-go/kimi-k3",
|
||||
"opencode-go/deepseek-v4-flash",
|
||||
"opencode-go/qwen3.8-max",
|
||||
],
|
||||
)
|
||||
def test_frontier_model_families_are_accepted(model_name: str) -> None:
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
"""Tests for OpenCode (Zen/Go) subscription auth: prefix parsing and key store."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
from strix.config import codex, opencode
|
||||
|
||||
|
||||
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(codex, "AUTH_PATH", path)
|
||||
return path
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "slug", "base_url", "uses_responses"),
|
||||
[
|
||||
("opencode/claude-sonnet-5", "claude-sonnet-5", opencode.ZEN_BASE_URL, False),
|
||||
("opencode/gpt-5.4", "gpt-5.4", opencode.ZEN_BASE_URL, True),
|
||||
("opencode/grok-4.5", "grok-4.5", opencode.ZEN_BASE_URL, True),
|
||||
("OpenCode/Kimi-K3", "Kimi-K3", opencode.ZEN_BASE_URL, False),
|
||||
("opencode-go/kimi-k3", "kimi-k3", opencode.GO_BASE_URL, False),
|
||||
("opencode-go/gpt-5.6-luna", "gpt-5.6-luna", opencode.GO_BASE_URL, True),
|
||||
("opencode-go/grok-4.5", "grok-4.5", opencode.GO_BASE_URL, False),
|
||||
],
|
||||
)
|
||||
def test_subscription_model_parses_prefixes(
|
||||
model: str, slug: str, base_url: str, uses_responses: bool
|
||||
) -> None:
|
||||
parsed = opencode.subscription_model(model)
|
||||
assert parsed is not None
|
||||
assert parsed.slug == slug
|
||||
assert parsed.base_url == base_url
|
||||
assert parsed.uses_responses == uses_responses
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["openai/gpt-5.4", "chatgpt/gpt-5.4", "opencode/", "opencode-go/", "opencode", "", None],
|
||||
)
|
||||
def test_subscription_model_rejects_non_opencode(model: str | None) -> None:
|
||||
assert opencode.subscription_model(model) is None
|
||||
|
||||
|
||||
def test_store_roundtrip_and_logout() -> None:
|
||||
assert opencode.read_record() is None
|
||||
assert opencode.is_authenticated() is False
|
||||
|
||||
opencode.save_api_key("sk-oc-test")
|
||||
record = opencode.read_record()
|
||||
assert record is not None
|
||||
assert record["key"] == "sk-oc-test"
|
||||
assert opencode.is_authenticated() is True
|
||||
assert opencode.get_api_key() == "sk-oc-test"
|
||||
|
||||
opencode.logout()
|
||||
assert opencode.read_record() is None
|
||||
opencode.logout() # no-op when already gone
|
||||
|
||||
|
||||
def test_store_coexists_with_chatgpt_record() -> None:
|
||||
codex.save_record({"type": "oauth", "access": "a", "refresh": "r", "account_id": "acct"})
|
||||
opencode.save_api_key("sk-oc-test")
|
||||
|
||||
assert codex.read_record() is not None
|
||||
assert opencode.get_api_key() == "sk-oc-test"
|
||||
|
||||
opencode.logout()
|
||||
assert codex.read_record() is not None
|
||||
assert opencode.read_record() is None
|
||||
|
||||
|
||||
def test_get_api_key_raises_when_not_signed_in() -> None:
|
||||
with pytest.raises(opencode.OpencodeAuthError) as exc:
|
||||
opencode.get_api_key()
|
||||
assert exc.value.code == "not_authenticated"
|
||||
|
||||
|
||||
def test_auth_mode_covers_both_subscriptions() -> None:
|
||||
assert opencode.auth_mode("opencode/claude-sonnet-5") == "subscription"
|
||||
assert opencode.auth_mode("opencode-go/kimi-k3") == "subscription"
|
||||
assert opencode.auth_mode("chatgpt/gpt-5.4") == "subscription"
|
||||
assert opencode.auth_mode("openai/gpt-5.4") == "api_key"
|
||||
assert opencode.auth_mode(None) == "api_key"
|
||||
|
||||
|
||||
def test_subscription_provider() -> None:
|
||||
assert opencode.subscription_provider("opencode/claude-sonnet-5") == "opencode"
|
||||
assert opencode.subscription_provider("opencode-go/kimi-k3") == "opencode"
|
||||
assert opencode.subscription_provider("chatgpt/gpt-5.4") == "chatgpt"
|
||||
assert opencode.subscription_provider("openai/gpt-5.4") is None
|
||||
assert opencode.subscription_provider(None) is None
|
||||
|
||||
|
||||
def _response(status_code: int, text: str = "") -> mock.MagicMock:
|
||||
response = mock.MagicMock()
|
||||
response.status_code = status_code
|
||||
response.text = text
|
||||
return response
|
||||
|
||||
|
||||
def test_validate_api_key_accepts_ok() -> None:
|
||||
with mock.patch.object(requests, "get", return_value=_response(200)) as get:
|
||||
opencode.validate_api_key("sk-oc-test")
|
||||
assert get.call_args.kwargs["headers"]["Authorization"] == "Bearer sk-oc-test"
|
||||
|
||||
|
||||
def test_validate_api_key_rejects_unauthorized() -> None:
|
||||
with (
|
||||
mock.patch.object(requests, "get", return_value=_response(401)),
|
||||
pytest.raises(opencode.OpencodeAuthError) as exc,
|
||||
):
|
||||
opencode.validate_api_key("bad-key")
|
||||
assert exc.value.code == "invalid_key"
|
||||
|
||||
|
||||
def test_validate_api_key_maps_network_errors() -> None:
|
||||
with (
|
||||
mock.patch.object(requests, "get", side_effect=requests.ConnectionError("boom")),
|
||||
pytest.raises(opencode.OpencodeAuthError) as exc,
|
||||
):
|
||||
opencode.validate_api_key("sk-oc-test")
|
||||
assert exc.value.code == "unavailable"
|
||||
@@ -1,107 +0,0 @@
|
||||
"""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()
|
||||
@@ -110,28 +110,6 @@ def test_state_populates_model_warning_for_non_frontier_model() -> None:
|
||||
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:
|
||||
setup_args = args()
|
||||
setup_args.targets_info = [
|
||||
|
||||
@@ -70,53 +70,6 @@ def test_read_run_summary_finished_flag(tmp_path: Path) -> None:
|
||||
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:
|
||||
run_dir = _make_run(tmp_path, "empty", status="running", end_time=None)
|
||||
assert read_vulnerabilities(run_dir) == []
|
||||
|
||||
Reference in New Issue
Block a user