diff --git a/modules/users/tests/test_session_version_hardening.py b/modules/users/tests/test_session_version_hardening.py new file mode 100644 index 00000000..6eb02090 --- /dev/null +++ b/modules/users/tests/test_session_version_hardening.py @@ -0,0 +1,242 @@ +"""The revocation cache's rough edges: a racy read, and an early invalidation. + +The 30-second cache in front of ``User.session_version`` is a deliberate trade +(``session_version_cache``), but two things around it were not. +""" + +from __future__ import annotations + +import httpx +import pytest +from simple_module_db import RequestSession +from sqlalchemy import select +from users.models import User +from users.session_version_cache import ( + SESSION_VERSION_TTL_SECONDS, + clear_session_version_cache, + configure_session_version_cache, + read_session_version, + session_version_ttl, + store_session_version, +) + +_CREDS = {"username": "admin@example.com", "password": "AdminPass1!"} + + +@pytest.fixture(autouse=True) +def _cold_cache(): + clear_session_version_cache() + yield + configure_session_version_cache(SESSION_VERSION_TTL_SECONDS) + clear_session_version_cache() + + +class TestTheReadIsNotRacy: + """``if k in cache: return cache[k]`` on a TTLCache can raise between the + two — the entry expires in the gap. The caller has no ``except KeyError``, + so landing in that window turned a revocation check into a 500.""" + + def test_a_miss_reports_no_hit(self): + assert read_session_version("nobody") == (False, None) + + def test_a_cached_none_is_a_hit(self): + """``None`` means "no such row" and is a legitimate cached answer; + collapsing it with "not cached" costs a DB read on every request.""" + store_session_version("gone", None) + assert read_session_version("gone") == (True, None) + + def test_a_cached_value_is_a_hit(self): + store_session_version("dana", 4) + assert read_session_version("dana") == (True, 4) + + def test_an_entry_expiring_mid_read_does_not_raise(self, monkeypatch): + """The race, made deterministic: the membership test succeeds and the + subscript then finds nothing. One `.get` cannot be torn this way.""" + + class _Vanishing(dict): + """A cache whose entry expires between the test and the subscript.""" + + ttl = SESSION_VERSION_TTL_SECONDS # so the fixture can restore afterwards + + def __contains__(self, key): + return True + + def get(self, key, default=None): + return default + + def __getitem__(self, key): + raise KeyError(key) + + monkeypatch.setattr( + "users.session_version_cache._SESSION_VERSIONS", _Vanishing(), raising=True + ) + assert read_session_version("dana") == (False, None) + + +class TestTheWindowIsAnOperatorChoice: + def test_the_default_is_the_documented_thirty_seconds(self): + assert session_version_ttl() == SESSION_VERSION_TTL_SECONDS + + def test_it_can_be_shortened(self): + configure_session_version_cache(5) + assert session_version_ttl() == 5 + + def test_zero_disables_caching(self): + """The honest knob for a deployment that will not accept any window in + which one worker has not seen another's revocation.""" + configure_session_version_cache(0) + store_session_version("dana", 4) + assert read_session_version("dana") == (False, None) + + def test_reconfiguring_to_the_same_window_keeps_the_warm_cache(self): + store_session_version("dana", 4) + configure_session_version_cache(int(session_version_ttl())) + assert read_session_version("dana") == (True, 4) + + def test_a_negative_window_is_floored_rather_than_rejected(self): + configure_session_version_cache(-1) + assert session_version_ttl() == 0 + + +class TestInvalidationWaitsForTheCommit: + """Clearing before the row is durable means a failed commit leaves the cache + empty and the counter unchanged — the next read repopulates the *old* value + and quietly re-admits everything the bump was meant to strand.""" + + @staticmethod + async def _admin_id(app): + async with app.state.sm.db.session_factory() as session: + return ( + await session.execute(select(User.id).where(User.email == "admin@example.com")) + ).scalar_one() + + @pytest.mark.anyio + async def test_a_password_change_clears_the_entry(self, users_app, anon_client): + assert (await anon_client.post("/api/users/auth/login", data=_CREDS)).status_code == 204 + user_id = await self._admin_id(users_app) + store_session_version(user_id, 0) + assert read_session_version(user_id)[0] is True + + resp = await anon_client.post( + "/api/users/me/password", + json={"current_password": "AdminPass1!", "new_password": "NewAdminPass1!"}, + ) + assert resp.status_code == 204, resp.text + assert read_session_version(user_id)[0] is False + + @pytest.mark.anyio + async def test_revoke_all_clears_the_entry(self, users_app, anon_client): + assert (await anon_client.post("/api/users/auth/login", data=_CREDS)).status_code == 204 + user_id = await self._admin_id(users_app) + store_session_version(user_id, 0) + assert read_session_version(user_id)[0] is True + + resp = await anon_client.post("/api/users/me/sessions/revoke-all") + assert resp.status_code == 204, resp.text + assert read_session_version(user_id)[0] is False + + @pytest.mark.anyio + async def test_the_cache_still_holds_when_the_commit_begins(self, users_app, anon_client): + """The handler wiring, not just the hook. + + Fault injection cannot show this: patching ``commit`` to raise trips + ``user_manager.update``'s own commit first, before either invalidation + runs. So observe the ordering instead — record whether the entry is + still cached each time a commit begins. Inline invalidation empties it + during the handler, so the *last* observation is a miss; hung on the + commit, the entry is still there when the final commit starts. + """ + from sqlalchemy import event + from sqlalchemy.orm import Session + + assert (await anon_client.post("/api/users/auth/login", data=_CREDS)).status_code == 204 + user_id = await self._admin_id(users_app) + store_session_version(user_id, 0) + + cached_at_commit: list[bool] = [] + + def observe(session): + cached_at_commit.append(read_session_version(user_id)[0]) + + event.listen(Session, "before_commit", observe) + try: + resp = await anon_client.post( + "/api/users/me/password", + json={"current_password": "AdminPass1!", "new_password": "NewAdminPass1!"}, + ) + finally: + event.remove(Session, "before_commit", observe) + + assert resp.status_code == 204, resp.text + assert cached_at_commit, "no commit was observed — the probe did not fire" + assert cached_at_commit[-1] is True, ( + "the cache was cleared before the row was durable: a failed commit would " + "leave the old counter to be read back as current" + ) + # And it is gone once the commit has landed. + assert read_session_version(user_id)[0] is False + + @pytest.mark.anyio + async def test_a_rolled_back_bump_leaves_the_cache_intact(self, users_app): + """The whole point of the move: no commit, no invalidation. Otherwise + the cache is emptied for a bump that never happened, and the next read + caches the old counter as if it were current.""" + user_id = await self._admin_id(users_app) + store_session_version(user_id, 3) + + from users.session_version_cache import forget_session_version + + async with users_app.state.sm.db.session_factory() as session: + assert isinstance(session, RequestSession) + session.on_commit(lambda: forget_session_version(user_id)) + await session.rollback() + + assert read_session_version(user_id) == (True, 3) + + +class TestTheProviderStillSeesRevocations: + @pytest.mark.anyio + async def test_a_bump_signs_the_other_browser_out(self, users_app, anon_client): + """End to end, through the cache: the guarantee the hardening must not + have broken.""" + from simple_module_test import forge_session_cookie + + assert (await anon_client.post("/api/users/auth/login", data=_CREDS)).status_code == 204 + user_id = await self._admin_id_of(users_app) + + cookie = forge_session_cookie( + str(users_app.state.sm.settings.secret_key), + {"user_id": str(user_id), "session_version": 0}, + ) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=users_app), + base_url="http://testserver", + cookies={"session": cookie}, + ) as other: + assert (await other.get("/admin/users/")).status_code == 200 + + assert (await anon_client.post("/api/users/me/sessions/revoke-all")).status_code == 204 + assert (await other.get("/admin/users/")).status_code in (401, 302, 303) + + @staticmethod + async def _admin_id_of(app): + async with app.state.sm.db.session_factory() as session: + return ( + await session.execute(select(User.id).where(User.email == "admin@example.com")) + ).scalar_one() + + +def test_the_module_level_cache_is_the_one_the_helpers_use(): + """Guards the rebuild in ``configure_session_version_cache``: the helpers + look the global up by name, so reassigning it has to reach them. + + Read through the module rather than the name imported at collection time — + that name still points at the cache the rebuild replaced, which is exactly + the trap this test exists to describe. + """ + import users.session_version_cache as cache_module + + configure_session_version_cache(7) + store_session_version("dana", 1) + assert cache_module._SESSION_VERSIONS.get("dana") == 1 + assert cache_module._SESSION_VERSIONS.ttl == 7 diff --git a/modules/users/users/auth_local/self_account.py b/modules/users/users/auth_local/self_account.py index 7dc7ec55..8b5acc02 100644 --- a/modules/users/users/auth_local/self_account.py +++ b/modules/users/users/auth_local/self_account.py @@ -15,9 +15,9 @@ from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi_users import exceptions as fu_exceptions +from simple_module_db import RequestSession from simple_module_db.deps import get_db from sqlalchemy import delete, update -from sqlalchemy.ext.asyncio import AsyncSession from users.auth_local.rate_limit import enforce_auth_throughput_limit from users.constants import SESSION_VERSION_KEY @@ -44,7 +44,7 @@ async def change_my_password( request: Request, user=Depends(fastapi_users.current_user(active=True)), user_manager: UserManager = Depends(get_user_manager), - db: AsyncSession = Depends(get_db), + db: RequestSession = Depends(get_db), ) -> Response: """Change your own password, proving you know the current one. @@ -91,9 +91,15 @@ async def change_my_password( # UPDATE would be rolled back as a read-only request. user.session_version = int(user.session_version or 0) + 1 request.session[SESSION_VERSION_KEY] = user.session_version - # This worker caches the counter for up to 30s; without this it would keep - # honouring the sessions this change was meant to strand. - forget_session_version(user.id) + # This worker caches the counter; without this it would keep honouring the + # sessions this change was meant to strand. Hung on the commit rather than + # run inline: clearing before the row is durable means a failed commit + # leaves the cache empty and the counter unchanged, so the next read + # repopulates the *old* value and quietly re-admits everything. The + # callback runs inside the response cycle, so the browser that pressed the + # button is never told it worked while this worker still lets the old + # sessions in. + db.on_commit(lambda: forget_session_version(user.id)) # The bearer half of the same revocation. The stamped ``session_version`` # on each row already strands them, but leaving the rows behind means a # stolen token keeps resolving a row until its own deadline passes — and @@ -113,7 +119,7 @@ async def change_my_password( async def revoke_all_my_sessions( request: Request, user=Depends(fastapi_users.current_user(active=True)), - db: AsyncSession = Depends(get_db), + db: RequestSession = Depends(get_db), ) -> Response: """Sign this account out of every browser and client, including this one. @@ -132,8 +138,10 @@ async def revoke_all_my_sessions( # ``get_db`` dependency — so the assignment is enough. user.session_version = int(user.session_version or 0) + 1 # See the password change: the cached counter has to go with the bump, or - # this worker keeps letting the sessions it just revoked back in. - forget_session_version(user.id) + # this worker keeps letting the sessions it just revoked back in — and it + # has to go *after* the commit, or a rollback restores the old counter + # behind an emptied cache. + db.on_commit(lambda: forget_session_version(user.id)) await db.execute( update(RefreshToken) .where(RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None)) diff --git a/modules/users/users/module.py b/modules/users/users/module.py index c7464347..bbc99ae7 100644 --- a/modules/users/users/module.py +++ b/modules/users/users/module.py @@ -224,6 +224,7 @@ async def on_startup(self, app: FastAPI) -> None: from users.mailer import build_mailer, default_app_name from users.oauth.providers import build_client_map, provider_buttons from users.roles_cache import refresh_roles_cache + from users.session_version_cache import configure_session_version_cache state = app.state.users s = state.settings @@ -287,6 +288,10 @@ def _app_name() -> str: s.login_redirect_url = f"{first_view}/" if first_view else "/" reconfigure_cookie_transport(auth_backend, s) + # The revocation cache's staleness window is an operator choice: the + # default trades one indexed read per request for a bounded window in + # which another worker's revocation is not yet seen here. + configure_session_version_cache(s.session_version_cache_ttl_seconds) await asyncio.gather( bootstrap_admin_from_env(app), diff --git a/modules/users/users/session_version_cache.py b/modules/users/users/session_version_cache.py index a356d9a0..165bb484 100644 --- a/modules/users/users/session_version_cache.py +++ b/modules/users/users/session_version_cache.py @@ -14,9 +14,11 @@ __all__ = [ "SESSION_VERSION_TTL_SECONDS", "clear_session_version_cache", + "configure_session_version_cache", "forget_session_version", "peek_session_version", "read_session_version", + "session_version_ttl", "store_session_version", ] @@ -32,10 +34,17 @@ still being let in. 30 seconds is chosen to be shorter than any plausible "did it work?" retry and -long enough to collapse a page's worth of requests into one read. +long enough to collapse a page's worth of requests into one read. It is the +*default*, not a constant: an operator who considers any cross-process lag +unacceptable for a password change made because an account is believed +compromised can shorten it — to 0, which disables the cache and pays the read on +every request — via ``SM_USERS_SESSION_VERSION_TTL_SECONDS``. See +:func:`configure_session_version_cache`. """ -_SESSION_VERSIONS: TTLCache = TTLCache(maxsize=10_000, ttl=SESSION_VERSION_TTL_SECONDS) +_CACHE_MAXSIZE = 10_000 + +_SESSION_VERSIONS: TTLCache = TTLCache(maxsize=_CACHE_MAXSIZE, ttl=SESSION_VERSION_TTL_SECONDS) """``user_id -> session_version`` (or ``None`` for "no such row"). Bounded and per-process. An LRU eviction is not a correctness problem: a miss @@ -43,6 +52,10 @@ """ +_MISS = object() +"""Distinguishes "not cached" from a cached ``None`` in a single lookup.""" + + def forget_session_version(user_id) -> None: """Drop this account's cached revocation counter. @@ -55,9 +68,9 @@ def forget_session_version(user_id) -> None: def peek_session_version(user_id): """The cached counter without reading the DB, or ``None`` when not cached. - For tests and diagnostics: the cache stores ``None`` for a missing row, so - a caller that needs to tell "absent" from "cached as missing" apart should - use ``user_id in _SESSION_VERSIONS``. + For tests and diagnostics. The cache stores ``None`` for a missing row, so a + caller that needs to tell "absent" from "cached as missing" apart wants + :func:`read_session_version`, which reports both in one read. """ return _SESSION_VERSIONS.get(user_id) @@ -70,15 +83,47 @@ def clear_session_version_cache() -> None: def read_session_version(user_id) -> tuple[bool, int | None]: """``(hit, stored)`` — ``hit`` is False when nothing is cached. - Two return values rather than a sentinel because ``None`` is a legitimate - cached answer ("no such row"), and collapsing it with "not cached" would - turn a deleted account into a DB read on every request. + Two return values rather than a sentinel in the *signature* because ``None`` + is a legitimate cached answer ("no such row"), and collapsing it with "not + cached" would turn a deleted account into a DB read on every request. + + One ``get`` rather than ``in`` then ``[]``: this is a ``TTLCache``, and an + entry can expire between the two. The window is sub-microsecond and the + caller has no ``except KeyError``, so landing in it turned a revocation + check into a 500 on an authenticated request. """ - if user_id in _SESSION_VERSIONS: - return True, _SESSION_VERSIONS[user_id] - return False, None + cached = _SESSION_VERSIONS.get(user_id, _MISS) + if cached is _MISS: + return False, None + return True, cached def store_session_version(user_id, stored: int | None) -> None: """Record what the DB answered for this account.""" _SESSION_VERSIONS[user_id] = None if stored is None else int(stored) + + +def session_version_ttl() -> float: + """The window currently in effect, in seconds.""" + return _SESSION_VERSIONS.ttl + + +def configure_session_version_cache(ttl_seconds: int) -> None: + """Rebuild the cache with an operator-chosen staleness window. + + ``TTLCache.ttl`` is read-only, so a different window means a new cache. Called + once from ``UsersModule.on_startup``; a no-op when the window already matches, + so a settings reload does not throw away a warm cache for nothing. + + ``0`` disables caching — every entry expires the moment it is written, so the + revocation check goes back to one indexed read per request. That is the honest + knob for a deployment that will not accept *any* window in which one worker + has not yet seen another's revocation. The cross-process fix proper is a + shared invalidation channel, which this layer cannot reach: Redis belongs to + the ``background_tasks`` plugin, and the framework ``EventBus`` is in-process. + """ + global _SESSION_VERSIONS + ttl = max(0, int(ttl_seconds)) + if ttl == _SESSION_VERSIONS.ttl: + return + _SESSION_VERSIONS = TTLCache(maxsize=_CACHE_MAXSIZE, ttl=ttl) diff --git a/modules/users/users/settings.py b/modules/users/users/settings.py index 11f5ea23..25825d44 100644 --- a/modules/users/users/settings.py +++ b/modules/users/users/settings.py @@ -19,6 +19,8 @@ from simple_module_core.redirect_safety import non_empty_redirect from simple_module_core.settings_base import DbBackedSettings +from users.session_version_cache import SESSION_VERSION_TTL_SECONDS + _PLACEHOLDER_RESET_SECRET = "dev-reset-token-secret-change-me" _PLACEHOLDER_VERIFY_SECRET = "dev-verify-token-secret-change-me" @@ -81,6 +83,12 @@ def _non_empty_redirect(cls, value: str) -> str: # number out of this rather than spelling it in the copy, so an operator # who shortens the window does not leave the checkbox lying. remember_me_max_age_seconds: int = 60 * 60 * 24 * 30 # 30 days + # How long a worker reuses a read of ``User.session_version`` before going + # back to the DB. The cost of the cache is a window in which *another* + # worker's revocation has not been seen here; 0 disables it and pays one + # indexed read per request instead. See ``users.session_version_cache``. + session_version_cache_ttl_seconds: int = SESSION_VERSION_TTL_SECONDS + cookie_secure: bool = True # flipped False in dev by the module at startup cookie_samesite: str = "lax"