mirror of
https://github.com/usestrix/strix.git
synced 2026-08-17 01:29:42 +02:00
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:
@@ -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
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user