"""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 from unittest import mock import pytest import requests from strix.config import codex if TYPE_CHECKING: from pathlib import Path def _fake_jwt(account_id: str) -> str: def seg(obj: dict[str, Any]) -> str: return base64.urlsafe_b64encode(json.dumps(obj).encode()).rstrip(b"=").decode() header = seg({"alg": "none"}) payload = seg({"https://api.openai.com/auth": {"chatgpt_account_id": account_id}}) return f"{header}.{payload}.sig" @pytest.fixture(autouse=True) def _tmp_store(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path: path = tmp_path / "home" / ".strix" / "subscription-auth.json" monkeypatch.setattr(codex, "AUTH_PATH", path) return path def test_pkce_challenge_matches_verifier_and_is_unpadded() -> None: verifier, challenge = codex.generate_pkce() expected = ( base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode() ) assert challenge == expected assert "=" not in verifier assert "=" not in challenge def test_authorize_url_carries_pkce_and_client() -> None: url = codex.build_authorize_url("chal", "st8") assert codex.AUTHORIZE_URL in url assert "code_challenge=chal" in url assert "code_challenge_method=S256" in url assert f"client_id={codex.CLIENT_ID}" in url assert "state=st8" in url def test_post_form_returns_parsed_body() -> None: resp = mock.MagicMock() resp.status_code = 200 resp.content = b'{"access_token": "tok"}' with mock.patch.object(requests, "post", return_value=resp) as post: data = codex._post_form({"grant_type": "refresh_token"}) assert data == {"access_token": "tok"} assert post.call_args.kwargs["timeout"] == codex._TOKEN_TIMEOUT @pytest.mark.parametrize( ("value", "expected"), [ ("http://localhost:1455/auth/callback?code=AAA&state=BBB", ("AAA", "BBB")), ("AAA#BBB", ("AAA", "BBB")), ("code=AAA&state=BBB", ("AAA", "BBB")), ("AAA", ("AAA", None)), ("", (None, None)), ], ) def test_parse_redirect_input(value: str, expected: tuple[str | None, str | None]) -> None: assert codex.parse_redirect_input(value) == expected @pytest.mark.parametrize( ("model", "expected"), [ ("chatgpt/gpt-5.4", "gpt-5.4"), ("ChatGPT/GPT-5.5", "GPT-5.5"), (" chatgpt/gpt-5.4 ", "gpt-5.4"), ("openai/gpt-5.4", None), # metered API path ("anthropic/claude-opus-4-8", None), ("gpt-5.4", None), ("chatgpt/", None), ("", None), (None, None), ], ) def test_subscription_model(model: str | None, expected: str | None) -> None: assert codex.subscription_model(model) == expected def test_auth_mode() -> None: assert codex.auth_mode("chatgpt/gpt-5.4") == "subscription" assert codex.auth_mode("openai/gpt-5.4") == "api_key" assert codex.auth_mode("anthropic/claude-opus-4-8") == "api_key" assert codex.auth_mode(None) == "api_key" def test_is_content_guardrail_error() -> None: # The backend's real wording (from a live gpt-5.6-sol block). raw = RuntimeError( "This content was flagged for possible cybersecurity risk. If this seems " "wrong, try rephrasing. To get authorized, join the Trusted Access for Cyber program." ) assert codex.is_content_guardrail_error(raw) is True # The already-typed error is recognized regardless of its message wording. assert codex.is_content_guardrail_error(codex.CodexContentGuardrailError("gpt-5.6-sol")) is True # Unrelated errors are not misclassified. assert codex.is_content_guardrail_error(RuntimeError("rate limit exceeded")) is False def test_content_guardrail_error_message() -> None: err = codex.CodexContentGuardrailError("gpt-5.6-sol") assert err.model == "gpt-5.6-sol" assert "gpt-5.6-sol" in str(err) assert "STRIX_LLM" in str(err) def test_account_id_from_jwt() -> None: assert codex._account_id_from_jwt(_fake_jwt("acct-42")) == "acct-42" assert codex._account_id_from_jwt("not-a-jwt") is None assert codex._account_id_from_jwt("") is None def test_store_roundtrip_and_logout() -> None: assert codex.read_record() is None assert codex.is_authenticated() is False codex.save_record( { "type": "oauth", "provider": "codex", "access": _fake_jwt("acct-42"), "refresh": "r1", "account_id": "acct-42", "expires_at": time.time() + 3600, } ) record = codex.read_record() assert record is not None assert record["account_id"] == "acct-42" assert codex.is_authenticated() is True codex.logout() assert codex.read_record() is None codex.logout() # no-op when already gone def test_read_record_rejects_incomplete_records() -> None: codex.save_record({"type": "oauth", "access": "a"}) # missing refresh/account assert codex.read_record() is None assert codex.is_authenticated() is False def test_get_valid_token_returns_stored_when_fresh(monkeypatch: pytest.MonkeyPatch) -> None: def _boom(_payload: dict[str, str]) -> dict[str, Any]: msg = "should not refresh a fresh token" raise AssertionError(msg) monkeypatch.setattr(codex, "_post_form", _boom) codex.save_record( { "type": "oauth", "provider": "codex", "access": "access-fresh", "refresh": "r1", "account_id": "acct-42", "expires_at": time.time() + 3600, } ) assert codex.get_valid_token() == ("access-fresh", "acct-42") def test_get_valid_token_refreshes_and_persists_rotation(monkeypatch: pytest.MonkeyPatch) -> None: calls = {"n": 0} def _fake_post(payload: dict[str, str]) -> dict[str, Any]: calls["n"] += 1 assert payload["grant_type"] == "refresh_token" assert payload["refresh_token"] == "r1" return {"access_token": _fake_jwt("acct-42"), "refresh_token": "r2", "expires_in": 3600} monkeypatch.setattr(codex, "_post_form", _fake_post) codex.save_record( { "type": "oauth", "provider": "codex", "access": "stale", "refresh": "r1", "account_id": "acct-42", "expires_at": time.time() - 10, # already expired } ) _access, account_id = codex.get_valid_token() assert calls["n"] == 1 assert account_id == "acct-42" # Rotated refresh token was written back to the store. record = codex.read_record() assert record is not None assert record["refresh"] == "r2" def test_get_valid_token_uses_token_rotated_by_another_process( monkeypatch: pytest.MonkeyPatch, ) -> None: # Simulate a parallel Strix process rotating the token while we wait for the # refresh guard: the pre-guard read sees the stale token, the in-guard read # sees the winner's fresh one, so we must NOT exchange the now-dead refresh. records = [ { "type": "oauth", "provider": "codex", "access": "stale", "refresh": "r1", "account_id": "acct", "expires_at": time.time() - 10, }, { "type": "oauth", "provider": "codex", "access": "fresh-from-other-process", "refresh": "r2", "account_id": "acct", "expires_at": time.time() + 3600, }, ] calls = {"n": 0} def _fake_read() -> dict[str, Any]: record = records[min(calls["n"], len(records) - 1)] calls["n"] += 1 return record def _boom(_payload: dict[str, str]) -> dict[str, Any]: msg = "must not refresh a token another process already rotated" raise AssertionError(msg) monkeypatch.setattr(codex, "read_record", _fake_read) monkeypatch.setattr(codex, "_post_form", _boom) access, account_id = codex.get_valid_token() assert access == "fresh-from-other-process" assert account_id == "acct" def _expired_record(refresh: str, access: str) -> dict[str, Any]: return { "type": "oauth", "provider": "codex", "access": access, "refresh": refresh, "account_id": "acct-42", "expires_at": time.time() - 10, } def test_get_valid_token_recovers_when_refresh_loses_race( monkeypatch: pytest.MonkeyPatch, ) -> None: # Lock failed open: our in-guard read still saw the stale token, so we tried to # refresh and lost the race (invalid_grant). By then a peer has saved a fresh # token — recover from it instead of failing the scan on the dead one. codex.save_record(_expired_record("r1", "stale")) def _fake_post(_payload: dict[str, str]) -> dict[str, Any]: codex.save_record( { "type": "oauth", "provider": "codex", "access": "fresh-from-peer", "refresh": "r2", "account_id": "acct-42", "expires_at": time.time() + 3600, } ) raise codex.CodexAuthError("token_http_error", "HTTP 400: invalid_grant") monkeypatch.setattr(codex, "_post_form", _fake_post) assert codex.get_valid_token() == ("fresh-from-peer", "acct-42") def test_get_valid_token_reraises_refresh_error_without_rotation( monkeypatch: pytest.MonkeyPatch, ) -> None: # Refresh fails and no peer rotated the token: surface the error, don't mask it. codex.save_record(_expired_record("r1", "stale")) def _fake_post(_payload: dict[str, str]) -> dict[str, Any]: raise codex.CodexAuthError("token_http_error", "HTTP 400: invalid_grant") monkeypatch.setattr(codex, "_post_form", _fake_post) with pytest.raises(codex.CodexAuthError): codex.get_valid_token() def test_get_valid_token_raises_when_not_signed_in() -> None: with pytest.raises(codex.CodexAuthError) as exc: codex.get_valid_token() assert exc.value.code == "not_authenticated"