mirror of
https://github.com/usestrix/strix.git
synced 2026-08-21 18:52:47 +02:00
feat(auth): sign in with a ChatGPT subscription for inference
Add an OAuth-based path to run Strix on a user's ChatGPT Plus/Pro subscription instead of a metered API key, modeled on OpenAI's Codex CLI. Auth: - strix/auth: Codex OAuth login (authorization-code + PKCE), a 0600 token store, refresh-on-expiry, and an AsyncOpenAI client that routes inference through the ChatGPT backend (chatgpt.com/backend-api/codex) with a per-request auth hook so long scans survive token expiry. - `strix auth login|logout|status` CLI (browser loopback on :1455 with a manual-paste fallback); STRIX_AUTH_MODE=subscription persisted to config. Inference wiring: - Subscription branch in configure_sdk_model_defaults installs the Codex client and the Responses API. - _CodexResponsesModel always streams (the backend rejects non-streamed requests) and aggregates back for the non-streaming get_response path. - store=false + encrypted reasoning for the stateless backend; models coerced to plan-available names (default gpt-5.4 — 5.5+ apply stricter content moderation that interferes with security testing). UX / reporting: - Track tokens but report $0.00 in the TUI, completion panel, and web viewer run details; record auth_mode in run.json and PostHog/Scarf. - Graceful, actionable errors for unavailable models and expired sign-in. - Restyled OAuth callback page (Strix branding + link to strix.ai). Tests: PKCE/URL/redirect parsing, token refresh + account-id, streaming aggregation, cost zeroing, CLI routing/provider aliasing. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
89a707ff51
commit
d35af02e47
@@ -0,0 +1,69 @@
|
||||
"""Tests for the `strix auth` CLI: subcommand routing and provider naming."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.auth import codex, store
|
||||
from strix.interface import auth_cli
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(store, "AUTH_PATH", tmp_path / "home" / ".strix" / "subscription-auth.json")
|
||||
|
||||
|
||||
def test_login_provider_is_chatgpt() -> None:
|
||||
assert auth_cli.LOGIN_PROVIDER == "chatgpt"
|
||||
assert codex.PROVIDER in auth_cli._ACCEPTED_PROVIDERS
|
||||
assert "chatgpt" in auth_cli._ACCEPTED_PROVIDERS
|
||||
|
||||
|
||||
def test_unknown_subcommand_returns_usage_error() -> None:
|
||||
assert auth_cli.run_auth(["bogus"]) == 2
|
||||
|
||||
|
||||
def test_help_returns_zero() -> None:
|
||||
assert auth_cli.run_auth(["--help"]) == 0
|
||||
|
||||
|
||||
def test_status_not_signed_in() -> None:
|
||||
assert auth_cli.run_auth(["status"]) == 1
|
||||
|
||||
|
||||
def test_login_rejects_unsupported_provider(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def _should_not_run(*_args: Any, **_kwargs: Any) -> dict[str, Any]:
|
||||
msg = "OAuth flow must not start for an unsupported provider"
|
||||
raise AssertionError(msg)
|
||||
|
||||
monkeypatch.setattr(auth_cli, "_run_oauth_flow", _should_not_run)
|
||||
assert auth_cli.run_auth(["login", "gemini"]) == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["chatgpt", "codex", "ChatGPT"])
|
||||
def test_login_accepts_provider_aliases(provider: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
reached = {"flow": False}
|
||||
|
||||
def _fake_flow(*_args: Any, **_kwargs: Any) -> dict[str, Any]:
|
||||
reached["flow"] = True
|
||||
return {
|
||||
"type": "oauth",
|
||||
"provider": "codex",
|
||||
"access": "a",
|
||||
"refresh": "r",
|
||||
"account_id": "acct",
|
||||
"expires_at": 0,
|
||||
}
|
||||
|
||||
monkeypatch.setattr(auth_cli, "_run_oauth_flow", _fake_flow)
|
||||
monkeypatch.setattr(codex, "save_record", lambda _record: None)
|
||||
monkeypatch.setattr(auth_cli, "_persist_subscription_config", lambda _model: None)
|
||||
|
||||
assert auth_cli.run_auth(["login", provider]) == 0
|
||||
assert reached["flow"] is True
|
||||
@@ -0,0 +1,189 @@
|
||||
"""Tests for ChatGPT (Codex) subscription auth: PKCE, token handling, store."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
|
||||
from strix.auth import codex, store
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _fake_jwt(account_id: str) -> str:
|
||||
def seg(obj: dict[str, Any]) -> str:
|
||||
return base64.urlsafe_b64encode(json.dumps(obj).encode()).rstrip(b"=").decode()
|
||||
|
||||
header = seg({"alg": "none"})
|
||||
payload = seg({"https://api.openai.com/auth": {"chatgpt_account_id": account_id}})
|
||||
return f"{header}.{payload}.sig"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path:
|
||||
path = tmp_path / "home" / ".strix" / "subscription-auth.json"
|
||||
monkeypatch.setattr(store, "AUTH_PATH", path)
|
||||
return path
|
||||
|
||||
|
||||
def test_pkce_challenge_matches_verifier_and_is_unpadded() -> None:
|
||||
verifier, challenge = codex.generate_pkce()
|
||||
expected = (
|
||||
base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode()
|
||||
)
|
||||
assert challenge == expected
|
||||
assert "=" not in verifier
|
||||
assert "=" not in challenge
|
||||
|
||||
|
||||
def test_authorize_url_carries_pkce_and_client() -> None:
|
||||
url = codex.build_authorize_url("chal", "st8")
|
||||
assert codex.AUTHORIZE_URL in url
|
||||
assert "code_challenge=chal" in url
|
||||
assert "code_challenge_method=S256" in url
|
||||
assert f"client_id={codex.CLIENT_ID}" in url
|
||||
assert "state=st8" in url
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("value", "expected"),
|
||||
[
|
||||
("http://localhost:1455/auth/callback?code=AAA&state=BBB", ("AAA", "BBB")),
|
||||
("AAA#BBB", ("AAA", "BBB")),
|
||||
("code=AAA&state=BBB", ("AAA", "BBB")),
|
||||
("AAA", ("AAA", None)),
|
||||
("", (None, None)),
|
||||
],
|
||||
)
|
||||
def test_parse_redirect_input(value: str, expected: tuple[str | None, str | None]) -> None:
|
||||
assert codex.parse_redirect_input(value) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected"),
|
||||
[
|
||||
("openai/gpt-5.5", "gpt-5.5"),
|
||||
("gpt-5.4", "gpt-5.4"),
|
||||
# A configured-but-unlisted OpenAI/bare name is passed through; the backend validates.
|
||||
("openai/gpt-5.6", "gpt-5.6"),
|
||||
# Another provider can't be served by a ChatGPT subscription → coerced to default.
|
||||
("anthropic/claude-opus-4-8", codex.DEFAULT_CODEX_MODEL),
|
||||
("deepseek/deepseek-v4-pro", codex.DEFAULT_CODEX_MODEL),
|
||||
("vertex_ai/gemini-3-pro", codex.DEFAULT_CODEX_MODEL),
|
||||
(None, codex.DEFAULT_CODEX_MODEL),
|
||||
("", codex.DEFAULT_CODEX_MODEL),
|
||||
],
|
||||
)
|
||||
def test_normalize_model(model: str | None, expected: str) -> None:
|
||||
assert codex.normalize_model(model) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "compatible"),
|
||||
[
|
||||
("openai/gpt-5.5", True),
|
||||
("gpt-5.4", True),
|
||||
("gpt-5.1-codex", True), # bare name: backend is the authority
|
||||
("anthropic/claude-opus-4-8", False),
|
||||
("deepseek/deepseek-v4-pro", False),
|
||||
("", False),
|
||||
(None, False),
|
||||
],
|
||||
)
|
||||
def test_is_backend_compatible(model: str | None, compatible: bool) -> None:
|
||||
assert codex.is_backend_compatible(model) is compatible
|
||||
|
||||
|
||||
def test_account_id_from_jwt() -> None:
|
||||
assert codex._account_id_from_jwt(_fake_jwt("acct-42")) == "acct-42"
|
||||
assert codex._account_id_from_jwt("not-a-jwt") is None
|
||||
assert codex._account_id_from_jwt("") is None
|
||||
|
||||
|
||||
def test_store_roundtrip_and_logout() -> None:
|
||||
assert codex.read_record() is None
|
||||
assert codex.is_authenticated() is False
|
||||
|
||||
codex.save_record(
|
||||
{
|
||||
"type": "oauth",
|
||||
"provider": "codex",
|
||||
"access": _fake_jwt("acct-42"),
|
||||
"refresh": "r1",
|
||||
"account_id": "acct-42",
|
||||
"expires_at": time.time() + 3600,
|
||||
}
|
||||
)
|
||||
record = codex.read_record()
|
||||
assert record is not None
|
||||
assert record["account_id"] == "acct-42"
|
||||
assert codex.is_authenticated() is True
|
||||
|
||||
codex.logout()
|
||||
assert codex.read_record() is None
|
||||
codex.logout() # no-op when already gone
|
||||
|
||||
|
||||
def test_read_record_rejects_incomplete_records() -> None:
|
||||
store.write_provider("codex", {"type": "oauth", "access": "a"}) # missing refresh/account
|
||||
assert codex.read_record() is None
|
||||
assert codex.is_authenticated() is False
|
||||
|
||||
|
||||
def test_get_valid_token_returns_stored_when_fresh(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def _boom(_payload: dict[str, str]) -> dict[str, Any]:
|
||||
msg = "should not refresh a fresh token"
|
||||
raise AssertionError(msg)
|
||||
|
||||
monkeypatch.setattr(codex, "_post_form", _boom)
|
||||
codex.save_record(
|
||||
{
|
||||
"type": "oauth",
|
||||
"provider": "codex",
|
||||
"access": "access-fresh",
|
||||
"refresh": "r1",
|
||||
"account_id": "acct-42",
|
||||
"expires_at": time.time() + 3600,
|
||||
}
|
||||
)
|
||||
assert codex.get_valid_token() == ("access-fresh", "acct-42")
|
||||
|
||||
|
||||
def test_get_valid_token_refreshes_and_persists_rotation(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
calls = {"n": 0}
|
||||
|
||||
def _fake_post(payload: dict[str, str]) -> dict[str, Any]:
|
||||
calls["n"] += 1
|
||||
assert payload["grant_type"] == "refresh_token"
|
||||
assert payload["refresh_token"] == "r1"
|
||||
return {"access_token": _fake_jwt("acct-42"), "refresh_token": "r2", "expires_in": 3600}
|
||||
|
||||
monkeypatch.setattr(codex, "_post_form", _fake_post)
|
||||
codex.save_record(
|
||||
{
|
||||
"type": "oauth",
|
||||
"provider": "codex",
|
||||
"access": "stale",
|
||||
"refresh": "r1",
|
||||
"account_id": "acct-42",
|
||||
"expires_at": time.time() - 10, # already expired
|
||||
}
|
||||
)
|
||||
_access, account_id = codex.get_valid_token()
|
||||
assert calls["n"] == 1
|
||||
assert account_id == "acct-42"
|
||||
# Rotated refresh token was written back to the store.
|
||||
assert codex.read_record()["refresh"] == "r2"
|
||||
|
||||
|
||||
def test_get_valid_token_raises_when_not_signed_in() -> None:
|
||||
with pytest.raises(codex.CodexAuthError) as exc:
|
||||
codex.get_valid_token()
|
||||
assert exc.value.code == "not_authenticated"
|
||||
@@ -0,0 +1,138 @@
|
||||
"""Regression test for the ChatGPT Codex backend's streaming requirement.
|
||||
|
||||
The backend rejects non-streamed requests with ``{"detail": "Stream must be set
|
||||
to true"}``. ``_CodexResponsesModel`` must therefore issue a streamed request
|
||||
even from the non-streaming ``get_response`` path and aggregate the events into
|
||||
a single response. A local server that mimics that behaviour proves the wrapper
|
||||
works where the stock responses model would fail.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import threading
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
from agents.model_settings import ModelSettings
|
||||
from agents.models.interface import ModelTracing
|
||||
from agents.models.openai_responses import OpenAIResponsesModel
|
||||
from openai import AsyncOpenAI, BadRequestError
|
||||
|
||||
from strix.config.models import _CodexResponsesModel
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
|
||||
|
||||
def _response_payload() -> dict[str, Any]:
|
||||
return {
|
||||
"id": "resp_1",
|
||||
"object": "response",
|
||||
"created_at": 0,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.5",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "m1",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "OK", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"input_tokens": 1,
|
||||
"output_tokens": 1,
|
||||
"total_tokens": 2,
|
||||
"input_tokens_details": {"cached_tokens": 0},
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
},
|
||||
"parallel_tool_calls": False,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"metadata": {},
|
||||
"temperature": 1.0,
|
||||
"top_p": 1.0,
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"instructions": None,
|
||||
"max_output_tokens": None,
|
||||
}
|
||||
|
||||
|
||||
class _Handler(BaseHTTPRequestHandler):
|
||||
def log_message(self, *args: Any) -> None:
|
||||
pass
|
||||
|
||||
def do_POST(self) -> None:
|
||||
length = int(self.headers.get("Content-Length", 0))
|
||||
body = json.loads(self.rfile.read(length) or b"{}")
|
||||
if not body.get("stream"):
|
||||
payload = json.dumps({"detail": "Stream must be set to true"}).encode()
|
||||
self.send_response(400)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(payload)))
|
||||
self.end_headers()
|
||||
self.wfile.write(payload)
|
||||
return
|
||||
event = {
|
||||
"type": "response.completed",
|
||||
"sequence_number": 0,
|
||||
"response": _response_payload(),
|
||||
}
|
||||
frame = f"event: response.completed\ndata: {json.dumps(event)}\n\n".encode()
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", "text/event-stream")
|
||||
self.end_headers()
|
||||
self.wfile.write(frame)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def backend_url() -> Iterator[str]:
|
||||
server = HTTPServer(("127.0.0.1", 0), _Handler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_address[1]}/backend-api/codex"
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
|
||||
|
||||
def _client(base_url: str) -> AsyncOpenAI:
|
||||
return AsyncOpenAI(api_key="tok", base_url=base_url)
|
||||
|
||||
|
||||
def _call_kwargs() -> dict[str, Any]:
|
||||
return {
|
||||
"system_instructions": "s",
|
||||
"input": "hi",
|
||||
"model_settings": ModelSettings(
|
||||
store=False, response_include=["reasoning.encrypted_content"]
|
||||
),
|
||||
"tools": [],
|
||||
"output_schema": None,
|
||||
"handoffs": [],
|
||||
"tracing": ModelTracing.DISABLED,
|
||||
"previous_response_id": None,
|
||||
"conversation_id": None,
|
||||
"prompt": None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stock_model_fails_on_non_streamed_backend(backend_url: str) -> None:
|
||||
model = OpenAIResponsesModel(model="gpt-5.5", openai_client=_client(backend_url))
|
||||
with pytest.raises(BadRequestError, match="Stream must be set to true"):
|
||||
await model.get_response(**_call_kwargs())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_model_streams_and_aggregates(backend_url: str) -> None:
|
||||
model = _CodexResponsesModel(model="gpt-5.5", openai_client=_client(backend_url))
|
||||
response = await model.get_response(**_call_kwargs())
|
||||
assert response.output[0].content[0].text == "OK"
|
||||
assert response.usage.total_tokens == 2
|
||||
@@ -35,6 +35,7 @@ async def test_persistent_rate_limit_stops_gracefully(
|
||||
settings = types.SimpleNamespace(
|
||||
llm=types.SimpleNamespace(
|
||||
model="openai/gpt-4o",
|
||||
auth_mode="api_key",
|
||||
reasoning_effort="high",
|
||||
force_required_tool_choice=False,
|
||||
timeout=300,
|
||||
|
||||
@@ -43,10 +43,12 @@ def _patch_engine_scaffold(
|
||||
settings = types.SimpleNamespace(
|
||||
llm=types.SimpleNamespace(
|
||||
model="openai/gpt-4o",
|
||||
auth_mode="api_key",
|
||||
reasoning_effort="high",
|
||||
force_required_tool_choice=False,
|
||||
timeout=300,
|
||||
)
|
||||
),
|
||||
runtime=types.SimpleNamespace(max_context_images=3),
|
||||
)
|
||||
monkeypatch.setattr(runner, "load_settings", lambda: settings)
|
||||
monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _settings: None)
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
"""Subscription runs track tokens but report zero cost."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from agents.usage import Usage
|
||||
|
||||
from strix.report.usage import LLMUsageLedger
|
||||
|
||||
|
||||
def _usage() -> Usage:
|
||||
usage = Usage()
|
||||
usage.requests = 1
|
||||
usage.input_tokens = 1000
|
||||
usage.output_tokens = 200
|
||||
usage.total_tokens = 1200
|
||||
return usage
|
||||
|
||||
|
||||
def test_zero_cost_ledger_keeps_tokens_but_reports_no_cost() -> None:
|
||||
ledger = LLMUsageLedger()
|
||||
ledger.zero_cost = True
|
||||
ledger.record(agent_id="a", usage=_usage(), agent_name="strix", model="gpt-5.5")
|
||||
|
||||
record = ledger.to_record()
|
||||
assert record["cost"] == 0.0
|
||||
assert record["total_tokens"] == 1200
|
||||
assert record["input_tokens"] == 1000
|
||||
assert record["output_tokens"] == 200
|
||||
assert ledger.total_cost == 0.0
|
||||
|
||||
|
||||
def test_zero_cost_ledger_ignores_observed_cost() -> None:
|
||||
ledger = LLMUsageLedger()
|
||||
ledger.zero_cost = True
|
||||
ledger.record_observed_cost(4.20)
|
||||
assert ledger.total_cost == 0.0
|
||||
|
||||
|
||||
def test_normal_ledger_still_estimates_cost() -> None:
|
||||
# Sanity check the flag is opt-in: without it, an OpenAI-native model still
|
||||
# accrues an estimated cost (proves zeroing is what suppresses it).
|
||||
ledger = LLMUsageLedger()
|
||||
ledger.record(agent_id="a", usage=_usage(), agent_name="strix", model="gpt-5.5")
|
||||
assert ledger.to_record()["total_tokens"] == 1200
|
||||
# Cost estimation depends on litellm's cost map; it should be >= 0 and not error.
|
||||
assert ledger.total_cost >= 0.0
|
||||
Reference in New Issue
Block a user