test: rewrite to exercise actual production _refresh_token_cache

The previous test used a _SharedCache simulation that proved the atomic
swap pattern works but didn't exercise the real app.py code. This rewrite
imports app.py with AUTH_ENABLED=true, mocks SessionLocal and logger,
creates a real AuthManager user, and calls the actual _refresh_token_cache()
while concurrent readers access the actual _token_cache global.

7 tests: single row, multiple prefixes, empty DB, app.state sync, dirty
flag cleared, 4 concurrent readers x 100 refreshes (zero empty reads),
and 50 create/revoke churn cycles with concurrent readers.
This commit is contained in:
isharak7m
2026-09-13 19:19:12 +05:30
parent 0b6d44890f
commit 6001f82019
+149 -137
View File
@@ -1,176 +1,188 @@
"""Token cache atomic swap prevents race condition during refresh.
"""Regression test for token cache race condition.
_refresh_token_cache() in app.py previously mutated the shared _token_cache
dict in two steps: .clear() then .update(). Between those calls the dict was
empty, so any concurrent reader (line 428) saw zero candidates and returned
401 for a valid token.
The fix replaces the two-step mutation with an atomic reference swap
(_token_cache = dict(new_map)). Python's GIL makes the assignment atomic,
so readers always see either the old fully-populated dict or the new one.
Exercises the actual _refresh_token_cache() and _token_cache in app.py
to verify the atomic swap fix eliminates the race window.
"""
import os
import sys
import threading
import time
from collections import defaultdict
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@pytest.fixture
def app_module(monkeypatch):
"""Import app.py with AUTH_ENABLED=true and minimal mocked deps.
def _make_token_row(prefix, token_id="t1", token_hash="h1", owner="admin", scopes="chat"):
Sets up a real AuthManager user ('admin') so normalize_known_username
resolves the token owner. Replaces SessionLocal with a MagicMock so
_refresh_token_cache() can run without a real DB.
"""
monkeypatch.setenv("AUTH_ENABLED", "true")
monkeypatch.setenv("DATABASE_URL", "sqlite:///:memory:")
# Clear cached app module
monkeypatch.delitem(sys.modules, "app", raising=False)
import app as app_mod # noqa: E402
app_mod.SessionLocal = MagicMock()
app_mod.logger = MagicMock()
app_mod.auth_manager.setup("admin", "TestPass123!")
return app_mod
def _seed(app_mod, rows):
"""Set up the mocked SessionLocal to return *rows* on next query."""
app_mod.SessionLocal.return_value.query.return_value.filter.return_value.all.return_value = rows
def _row(prefix, tid="t1", th="h1", owner="admin", scopes="chat"):
return SimpleNamespace(
token_prefix=prefix,
id=token_id,
token_hash=token_hash,
owner=owner,
scopes=scopes,
is_active=True,
token_prefix=prefix, id=tid, token_hash=th,
owner=owner, scopes=scopes, is_active=True,
)
class _SharedCache:
"""Mimics the module-level _token_cache global in app.py.
Both reader and writer access .current — the same reference object.
The writer atomically reassigns .current to a new dict; the GIL
ensures the reader never sees a half-built reference.
"""
def __init__(self, initial=None):
self.current = initial or {}
def _build_refresh_fn(shared, rows):
"""Build a _refresh_token_cache closure mirroring app.py's fixed logic."""
def _refresh():
new_map = defaultdict(list)
for r in rows:
scope_list = [s.strip() for s in (r.scopes or "chat").split(",") if s.strip()]
new_map[r.token_prefix].append((r.id, r.token_hash, r.owner, scope_list))
shared.current = dict(new_map)
return _refresh
# ---------------------------------------------------------------------------
# Tests
# Tests — all use the REAL app._refresh_token_cache and REAL app._token_cache
# ---------------------------------------------------------------------------
class TestAtomicSwapNoEmptyWindow:
"""The atomic swap must never leave _token_cache empty between frames."""
class TestRefreshPopulatesCache:
"""Single refresh call should populate _token_cache from DB rows."""
def test_swap_replaces_dict_content(self):
shared = _SharedCache({"old_prefix": [("old_id", "old_hash", "admin", ["chat"])]})
rows = [_make_token_row("new_prfx", "t2", "h2", "admin", "chat")]
refresh = _build_refresh_fn(shared, rows)
def test_single_row(self, app_module):
_seed(app_module, [_row("ody_abc")])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
assert "ody_abc" in app_module._token_cache
assert app_module._token_cache["ody_abc"][0][0] == "t1"
refresh()
def test_multiple_prefixes(self, app_module):
_seed(app_module, [
_row("ody_aaaa", "t1", "h1", "admin", "chat"),
_row("ody_bbbb", "t2", "h2", "admin", "chat,tools"),
_row("ody_aaaa", "t3", "h3", "admin", "memory"),
])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
assert len(app_module._token_cache) == 2
assert len(app_module._token_cache["ody_aaaa"]) == 2
assert app_module._token_cache["ody_bbbb"][0][3] == ["chat", "tools"]
assert "old_prefix" not in shared.current
assert "new_prfx" in shared.current
assert shared.current["new_prfx"][0] == ("t2", "h2", "admin", ["chat"])
def test_empty_db_clears_cache(self, app_module):
app_module._token_cache["stale"] = [("x", "y", "z", ["chat"])]
_seed(app_module, [])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
assert len(app_module._token_cache) == 0
def test_swap_is_atomic_under_concurrent_readers(self):
"""Concurrent readers never see an empty dict during refresh."""
shared = _SharedCache({"ody_test12": [("t1", "hash1", "admin", ["chat"])]})
rows = [_make_token_row("ody_newtok", "t2", "hash2", "admin", "chat")]
refresh = _build_refresh_fn(shared, rows)
class TestAppStateSync:
"""app.state._token_cache must stay synchronized with _token_cache."""
def test_state_ref_matches_after_refresh(self, app_module):
_seed(app_module, [_row("ody_sync")])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
assert app_module.app.state._token_cache is app_module._token_cache
assert "ody_sync" in app_module.app.state._token_cache
def test_state_dirty_cleared(self, app_module):
_seed(app_module, [_row("ody_x")])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
assert app_module.app.state._token_cache_dirty is False
class TestConcurrentReaders:
"""The core regression: concurrent readers must never see an empty cache."""
def test_no_empty_reads_during_refresh(self, app_module):
"""4 reader threads + 100 refreshes on the real _token_cache global."""
_seed(app_module, [_row("ody_race")])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
stop = threading.Event()
reader_results = {"empty": 0, "ok": 0}
results = {"empty": 0, "ok": 0}
def reader_loop():
def reader():
while not stop.is_set():
# Mimic app.py line 429: read the shared global directly
snapshot = shared.current
if len(snapshot) == 0:
reader_results["empty"] += 1
if len(app_module._token_cache) == 0:
results["empty"] += 1
else:
reader_results["ok"] += 1
results["ok"] += 1
threads = [threading.Thread(target=reader_loop, daemon=True) for _ in range(4)]
for t in threads:
def churner():
for _ in range(100):
app_module._refresh_token_cache()
readers = [threading.Thread(target=reader, daemon=True) for _ in range(4)]
for t in readers:
t.start()
time.sleep(0.01)
for _ in range(100):
refresh()
time.sleep(0.01)
churn = threading.Thread(target=churner)
churn.start()
time.sleep(0.1)
churn.join(timeout=5)
stop.set()
for t in threads:
for t in readers:
t.join(timeout=2)
assert reader_results["empty"] == 0, (
assert results["empty"] == 0, (
"Readers saw empty cache %d times (ok=%d)"
% (reader_results["empty"], reader_results["ok"])
% (results["empty"], results["ok"])
)
assert reader_results["ok"] > 0, "readers should have seen data at least once"
assert results["ok"] > 0
def test_app_state_ref_stays_in_sync(self):
"""app.state._token_cache must point to the same dict."""
shared = _SharedCache()
rows = [_make_token_row("pfx_a")]
refresh = _build_refresh_fn(shared, rows)
def test_no_empty_reads_with_token_churn(self, app_module):
"""Simulate token create/revoke churn while reading."""
_seed(app_module, [_row("ody_keep", "t1", "h1", "admin", "chat")])
app_module.app.state._token_cache_dirty = True
app_module._refresh_token_cache()
app_state = SimpleNamespace(_token_cache=shared.current)
refresh()
app_state._token_cache = shared.current
stop = threading.Event()
results = {"empty": 0, "ok": 0}
assert app_state._token_cache is shared.current
assert "pfx_a" in app_state._token_cache
def reader():
while not stop.is_set():
if len(app_module._token_cache) == 0:
results["empty"] += 1
else:
results["ok"] += 1
def churner():
for i in range(50):
_seed(app_module, [
_row("ody_keep", "t1", "h1", "admin", "chat"),
_row("ody_new_%d" % i, "t%d" % (i + 10), "h%d" % (i + 10), "admin", "chat"),
])
app_module._refresh_token_cache()
_seed(app_module, [_row("ody_keep", "t1", "h1", "admin", "chat")])
app_module._refresh_token_cache()
class TestRefreshFromDB:
"""Verify the refresh logic handles DB rows correctly."""
def test_multiple_prefixes(self):
shared = _SharedCache()
rows = [
_make_token_row("ody_aaaa", "t1", "h1", "admin", "chat"),
_make_token_row("ody_bbbb", "t2", "h2", "admin", "chat,tools"),
_make_token_row("ody_aaaa", "t3", "h3", "admin", "memory"),
]
refresh = _build_refresh_fn(shared, rows)
refresh()
assert len(shared.current) == 2
assert len(shared.current["ody_aaaa"]) == 2
assert len(shared.current["ody_bbbb"]) == 1
assert shared.current["ody_bbbb"][0][3] == ["chat", "tools"]
def test_empty_db_clears_cache(self):
shared = _SharedCache({"stale": [("x", "y", "z", ["chat"])]})
refresh = _build_refresh_fn(shared, rows=[])
refresh()
assert len(shared.current) == 0
def test_concurrent_refreshes_dont_corrupt(self):
"""Multiple threads refreshing simultaneously don't corrupt cache."""
shared = _SharedCache()
rows = [
_make_token_row("pfx_%d" % i, "t%d" % i, "h%d" % i, "admin", "chat")
for i in range(20)
]
refresh = _build_refresh_fn(shared, rows)
errors = []
def refresh_loop():
try:
for _ in range(50):
refresh()
except Exception as e:
errors.append(e)
threads = [threading.Thread(target=refresh_loop) for _ in range(4)]
for t in threads:
readers = [threading.Thread(target=reader, daemon=True) for _ in range(4)]
for t in readers:
t.start()
for t in threads:
t.join(timeout=5)
assert errors == []
assert len(shared.current) == 20
churn = threading.Thread(target=churner)
churn.start()
churn.join(timeout=10)
stop.set()
for t in readers:
t.join(timeout=2)
assert results["empty"] == 0, (
"Readers saw empty cache %d times during churn (ok=%d)"
% (results["empty"], results["ok"])
)
assert results["ok"] > 0
assert "ody_keep" in app_module._token_cache
assert len(app_module._token_cache) == 1