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.
This commit is contained in:
yoni
2026-07-29 17:39:26 +00:00
parent 5c94872186
commit 9fd11eedec
7 changed files with 247 additions and 114 deletions
+2
View File
@@ -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"]
+19 -53
View File
@@ -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):
+19 -53
View File
@@ -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):
+117
View File
@@ -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
+17 -8
View File
@@ -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:
+18
View File
@@ -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"
+55
View File
@@ -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"