From 9fd11eedec364ff45a45a2dfe10182b3514cf7ab Mon Sep 17 00:00:00 2001 From: yoni Date: Wed, 29 Jul 2026 17:39:26 +0000 Subject: [PATCH] fix(llm): harden subscription credential store and provider labeling Addresses PR review: - Extract the shared ~/.strix/subscription-auth.json handling into subscription_store, writing tokens owner-only (0600) from creation via os.open instead of chmod-after-write, closing the window where credentials were briefly world/group-readable. - Serialize read-modify-write across providers and processes with a reentrant lock, so overlapping ChatGPT/Grok save/logout/refresh operations no longer clobber each other. - Live/TUI stats label the subscription from the persisted run record (falling back to provider-aware settings), so resumed runs no longer mislabel the provider when STRIX_LLM changes. --- pyproject.toml | 2 + strix/config/codex.py | 72 +++++------------- strix/config/grok.py | 72 +++++------------- strix/config/subscription_store.py | 117 +++++++++++++++++++++++++++++ strix/interface/utils.py | 25 ++++-- tests/test_grok_routing.py | 18 +++++ tests/test_subscription_store.py | 55 ++++++++++++++ 7 files changed, 247 insertions(+), 114 deletions(-) create mode 100644 strix/config/subscription_store.py create mode 100644 tests/test_subscription_store.py diff --git a/pyproject.toml b/pyproject.toml index f25fc5f5..97b58a31 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -273,6 +273,8 @@ ignore = [ # 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"] diff --git a/strix/config/codex.py b/strix/config/codex.py index dcd277a3..d6741205 100644 --- a/strix/config/codex.py +++ b/strix/config/codex.py @@ -16,7 +16,6 @@ import hashlib import json import logging import secrets -import threading import time import urllib.parse from pathlib import Path @@ -24,6 +23,8 @@ from typing import TYPE_CHECKING, Any import requests +from strix.config import subscription_store + if TYPE_CHECKING: from collections.abc import Iterator @@ -52,33 +53,12 @@ _ACCOUNT_CLAIM = "https://api.openai.com/auth" _TOKEN_TIMEOUT = 30 _EXPIRY_SKEW_S = 300 -_refresh_lock = threading.Lock() - # Kept separate from cli-config.json so OAuth tokens never land in the env-var config. AUTH_PATH = Path.home() / ".strix" / "subscription-auth.json" -def _read_store() -> dict[str, Any]: - try: - data = json.loads(AUTH_PATH.read_text(encoding="utf-8")) - except (OSError, json.JSONDecodeError): - return {} - return data if isinstance(data, dict) else {} - - -def _write_store(data: dict[str, Any]) -> None: - AUTH_PATH.parent.mkdir(parents=True, exist_ok=True) - tmp = AUTH_PATH.with_suffix(".json.tmp") - tmp.write_text(json.dumps(data, indent=2), encoding="utf-8") - with contextlib.suppress(OSError): - tmp.chmod(0o600) - tmp.replace(AUTH_PATH) - with contextlib.suppress(OSError): - AUTH_PATH.chmod(0o600) - - def read_record() -> dict[str, Any] | None: - record = _read_store().get(PROVIDER) + 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") and record.get("account_id")): @@ -91,45 +71,31 @@ def is_authenticated() -> bool: def save_record(record: dict[str, Any]) -> None: - data = _read_store() - data[PROVIDER] = record - _write_store(data) + with subscription_store.guard(AUTH_PATH): + data = subscription_store.read(AUTH_PATH) + data[PROVIDER] = record + subscription_store.write(AUTH_PATH, data) def logout() -> None: - data = _read_store() - if PROVIDER not in data: - return - del data[PROVIDER] - if data: - _write_store(data) - return - with contextlib.suppress(OSError): - AUTH_PATH.unlink() + 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 _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() + with subscription_store.guard(AUTH_PATH): + yield class CodexAuthError(Exception): diff --git a/strix/config/grok.py b/strix/config/grok.py index 2350bc7d..e047a5e8 100644 --- a/strix/config/grok.py +++ b/strix/config/grok.py @@ -17,7 +17,6 @@ import hashlib import json import logging import secrets -import threading import time import urllib.parse from pathlib import Path @@ -25,6 +24,8 @@ from typing import TYPE_CHECKING, Any import requests +from strix.config import subscription_store + if TYPE_CHECKING: from collections.abc import Iterator @@ -51,34 +52,13 @@ XAI_BASE_URL = "https://api.x.ai/v1" _TOKEN_TIMEOUT = 30 _EXPIRY_SKEW_S = 300 -_refresh_lock = threading.Lock() - # 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_store() -> dict[str, Any]: - try: - data = json.loads(AUTH_PATH.read_text(encoding="utf-8")) - except (OSError, json.JSONDecodeError): - return {} - return data if isinstance(data, dict) else {} - - -def _write_store(data: dict[str, Any]) -> None: - AUTH_PATH.parent.mkdir(parents=True, exist_ok=True) - tmp = AUTH_PATH.with_suffix(".json.tmp") - tmp.write_text(json.dumps(data, indent=2), encoding="utf-8") - with contextlib.suppress(OSError): - tmp.chmod(0o600) - tmp.replace(AUTH_PATH) - with contextlib.suppress(OSError): - AUTH_PATH.chmod(0o600) - - def read_record() -> dict[str, Any] | None: - record = _read_store().get(PROVIDER) + 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")): @@ -91,45 +71,31 @@ def is_authenticated() -> bool: def save_record(record: dict[str, Any]) -> None: - data = _read_store() - data[PROVIDER] = record - _write_store(data) + with subscription_store.guard(AUTH_PATH): + data = subscription_store.read(AUTH_PATH) + data[PROVIDER] = record + subscription_store.write(AUTH_PATH, data) def logout() -> None: - data = _read_store() - if PROVIDER not in data: - return - del data[PROVIDER] - if data: - _write_store(data) - return - with contextlib.suppress(OSError): - AUTH_PATH.unlink() + 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 _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() + with subscription_store.guard(AUTH_PATH): + yield class GrokAuthError(Exception): diff --git a/strix/config/subscription_store.py b/strix/config/subscription_store.py new file mode 100644 index 00000000..3e68d205 --- /dev/null +++ b/strix/config/subscription_store.py @@ -0,0 +1,117 @@ +"""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 threading +from typing import TYPE_CHECKING, Any + + +if TYPE_CHECKING: + from collections.abc import Iterator + from io import TextIOWrapper + from pathlib import Path + + +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.""" + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_suffix(".json.tmp") + fd = os.open(str(tmp), os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600) + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + json.dump(data, handle, indent=2) + except BaseException: + with contextlib.suppress(OSError): + tmp.unlink() + raise + tmp.replace(path) + 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 + 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() + self._flock_handle = None + + +_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 | None: + try: + import fcntl + except ImportError: + return None + lock_path = path.with_suffix(".lock") + try: + lock_path.parent.mkdir(parents=True, exist_ok=True) + handle = lock_path.open("w") + except OSError: + return None + with contextlib.suppress(OSError): + fcntl.flock(handle.fileno(), fcntl.LOCK_EX) + return handle diff --git a/strix/interface/utils.py b/strix/interface/utils.py index 8aa3c0b3..716ed66e 100644 --- a/strix/interface/utils.py +++ b/strix/interface/utils.py @@ -266,13 +266,22 @@ def _is_subscription(report_state: Any) -> bool: return subscription.auth_mode(load_settings().llm.model) == "subscription" -def _subscription_label() -> str: - """Human label for the active model subscription (e.g. "Grok subscription").""" - from strix.config import grok +def _subscription_label(report_state: Any) -> str: + """Human label for the active model subscription (e.g. "Grok subscription"). - if grok.subscription_model(load_settings().llm.model): - return "Grok subscription" - return "ChatGPT subscription" + Prefers the persisted run record so a resumed run keeps its original provider + even if STRIX_LLM later points at a different one; falls back to current + settings. + """ + record = getattr(report_state, "run_record", None) + if isinstance(record, dict): + provider = record.get("subscription_provider") + if isinstance(provider, str) and provider: + return f"{provider} subscription" + from strix.config import subscription + + label = subscription.provider_label(load_settings().llm.model) + return f"{label} subscription" if label else "Subscription" def _int_stat(usage: dict[str, Any], key: str) -> int: @@ -371,7 +380,7 @@ def build_live_stats_text(report_state: Any) -> Text: stats_text.append(str(model), style="white") if _is_subscription(report_state): stats_text.append(" ยท ", style="dim white") - stats_text.append(_subscription_label(), style="#22c55e") + stats_text.append(_subscription_label(report_state), style="#22c55e") stats_text.append("\n") vuln_count = len(report_state.vulnerability_reports) @@ -417,7 +426,7 @@ def build_tui_stats_text(report_state: Any) -> Text: subscription = _is_subscription(report_state) if subscription: stats_text.append("\n") - stats_text.append(_subscription_label(), style="#22c55e") + stats_text.append(_subscription_label(report_state), style="#22c55e") usage = _llm_usage(report_state) if usage and _int_stat(usage, "total_tokens") > 0: diff --git a/tests/test_grok_routing.py b/tests/test_grok_routing.py index c4456a60..26ff796d 100644 --- a/tests/test_grok_routing.py +++ b/tests/test_grok_routing.py @@ -8,6 +8,7 @@ from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel from strix.config import grok, subscription from strix.config.models import StrixProvider +from strix.interface import utils from strix.report import state as state_mod @@ -51,3 +52,20 @@ def test_run_record_reports_grok_provider(monkeypatch) -> None: # type: ignore[ 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" diff --git a/tests/test_subscription_store.py b/tests/test_subscription_store.py new file mode 100644 index 00000000..522c8fcb --- /dev/null +++ b/tests/test_subscription_store.py @@ -0,0 +1,55 @@ +"""Shared subscription credential store: secure writes and cross-provider locking.""" + +from __future__ import annotations + +import stat +from typing import TYPE_CHECKING + +from strix.config import codex, grok, subscription_store + + +if TYPE_CHECKING: + from pathlib import Path + + import pytest + + +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_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"