fix(auth): swap the API token cache atomically instead of clearing it

`_refresh_token_cache` rebuilt the bearer-token map in two steps, `clear()`
then `update()`. Between them the dict a concurrent reader was already holding
was empty, so a valid token landing in that window found zero candidates and
got a 401. The refresh runs on a worker thread via `to_thread`, so the window
is real rather than theoretical.

Build the new map and rebind the name. The reader dereferences the global once
and then holds a map that is complete — the previous one if it read early, the
new one if it read late, never a half-built one. `app.state._token_cache` is
rebound with it, because it was bound once at startup and would otherwise point
at the abandoned dict.

Ported from public `dev` (`984337b3`), with its test.

The upstream test does not pin the fix: both of its concurrency cases pass
against the pre-fix code, because landing a GIL switch inside a window a few
bytecodes wide does not happen across 100 refreshes. They are kept as written
and `TestRefreshLeavesTheReadersMapAlone` is added next to them, stating the
same invariant at object level — a map a reader already holds is not mutated
by a later refresh — which fails on the pre-fix code without depending on
thread scheduling.
This commit is contained in:
Léo
2026-09-30 16:07:12 +02:00
parent 6105702901
commit 5e3e153fba
2 changed files with 236 additions and 2 deletions
+3 -2
View File
@@ -316,6 +316,7 @@ if AUTH_ENABLED:
def _refresh_token_cache():
"""Rebuild the prefix→[(id,hash)] map from the DB."""
global _token_cache
from collections import defaultdict
new_map = defaultdict(list)
db = SessionLocal()
@@ -334,8 +335,8 @@ if AUTH_ENABLED:
new_map[r.token_prefix].append((r.id, r.token_hash, owner_key, scopes))
finally:
db.close()
_token_cache.clear()
_token_cache.update(new_map)
_token_cache = dict(new_map)
app.state._token_cache = _token_cache
app.state._token_cache_dirty = False
# Headers that prove a request was forwarded by a proxy/tunnel (cloudflared,
+233
View File
@@ -0,0 +1,233 @@
"""Regression test for token cache race condition.
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 types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
@pytest.fixture
def app_module(monkeypatch):
"""Import app.py with AUTH_ENABLED=true and minimal mocked deps.
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=tid, token_hash=th,
owner=owner, scopes=scopes, is_active=True,
)
# ---------------------------------------------------------------------------
# Tests — all use the REAL app._refresh_token_cache and REAL app._token_cache
# ---------------------------------------------------------------------------
class TestRefreshPopulatesCache:
"""Single refresh call should populate _token_cache from DB 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"
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"]
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
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()
results = {"empty": 0, "ok": 0}
def reader():
while not stop.is_set():
if len(app_module._token_cache) == 0:
results["empty"] += 1
else:
results["ok"] += 1
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()
churn = threading.Thread(target=churner)
churn.start()
time.sleep(0.1)
churn.join(timeout=5)
stop.set()
for t in readers:
t.join(timeout=2)
assert results["empty"] == 0, (
"Readers saw empty cache %d times (ok=%d)"
% (results["empty"], results["ok"])
)
assert results["ok"] > 0
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()
stop = threading.Event()
results = {"empty": 0, "ok": 0}
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()
readers = [threading.Thread(target=reader, daemon=True) for _ in range(4)]
for t in readers:
t.start()
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
class TestRefreshLeavesTheReadersMapAlone:
"""The deterministic form of the invariant the timing tests above only sample.
`test_no_empty_reads_during_refresh` and `test_no_empty_reads_with_token_churn`
both pass against the pre-fix `.clear()` + `.update()` code: hitting a window
that is a handful of bytecodes wide needs a GIL switch to land inside it, and
across 100 and 50 refreshes it does not. They document the intent; they do not
pin it.
What the fix actually guarantees is object-level: a reader dereferences the
global once, and the dict it ends up holding is never mutated afterwards. The
rebuild builds a new map and rebinds the name, so a reader holding the old one
sees a complete stale map rather than a half-built current one. Asserting that
fails on the pre-fix code without depending on thread scheduling.
"""
def test_a_held_map_is_not_emptied_by_a_later_refresh(self, app_module):
_seed(app_module, [_row("ody_before", "t1", "h1", "admin", "chat")])
app_module._refresh_token_cache()
# What a reader would have dereferenced, and what it held at that moment.
held = app_module._token_cache
snapshot = dict(held)
assert snapshot, "fixture did not populate the cache"
_seed(app_module, [_row("ody_after", "t2", "h2", "admin", "chat")])
app_module._refresh_token_cache()
assert held == snapshot, (
"the refresh mutated the map a reader was already holding: %r" % (held,)
)
assert "ody_after" in app_module._token_cache
assert "ody_before" not in app_module._token_cache
def test_app_state_follows_the_rebind(self, app_module):
"""app.state._token_cache is bound once at startup and must not go stale."""
_seed(app_module, [_row("ody_one")])
app_module._refresh_token_cache()
_seed(app_module, [_row("ody_two")])
app_module._refresh_token_cache()
assert app_module.app.state._token_cache is app_module._token_cache
assert "ody_two" in app_module.app.state._token_cache