Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 13 additions & 3 deletions api/app/app/core/token_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,17 @@ def revoke_refresh_jti(jti: str) -> None:
logger.error("Failed to revoke refresh jti: %s", exc)


def revoke_all_user_refresh_tokens(user_id: str) -> int:
def revoke_all_user_refresh_tokens(user_id: str) -> int | None:
"""Revoke every refresh token issued to a user.

Returns the number of tokens revoked (0 if none existed), or None if
the operation could not be completed (e.g. Redis unreachable) -- the
caller MUST treat None as "revocation not guaranteed", not as
"nothing to revoke". This matters most for password-reset, which
relies on this call to actually invalidate any session an attacker
may hold; silently reporting success on a Redis hiccup would leave
those sessions valid with no indication anything went wrong.
"""
try:
redis = get_redis()
index_key = f"{USER_INDEX_PREFIX}{user_id}"
Expand All @@ -74,5 +84,5 @@ def revoke_all_user_refresh_tokens(user_id: str) -> int:
pipe.execute()
return len(jtis)
except Exception as exc:
logger.error("Failed to bulk-revoke refresh tokens: %s", exc)
return 0
logger.error("Failed to bulk-revoke refresh tokens for %s: %s", user_id, exc)
return None
16 changes: 15 additions & 1 deletion api/app/app/modules/auth/router.py
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,21 @@ def password_reset_confirm_endpoint(
detail="Invalid or expired reset token.",
)
if result.user_id:
revoke_all_user_refresh_tokens(result.user_id)
revoked = revoke_all_user_refresh_tokens(result.user_id)
if revoked is None:
# The password itself was already changed -- don't pretend
# that didn't happen -- but we can't guarantee any refresh
# token an attacker holds was actually invalidated, which is
# the whole point of revoking on reset. Surface that instead
# of silently reporting full success.
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=(
"Your password was changed, but we couldn't confirm "
"your other sessions were signed out. Please sign "
"out of other devices manually, or try again shortly."
),
)
return SimpleStatusResponse()


Expand Down
68 changes: 68 additions & 0 deletions api/app/tests/test_token_store_revocation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
"""revoke_all_user_refresh_tokens must distinguish "nothing to revoke"
(0) from "couldn't complete the revocation" (None) -- callers like
password-reset confirm rely on this to know whether it's safe to report
success."""

from typing import Any


from app.core import token_store


class FakePipeline:
def __init__(self, redis: "FakeRedis") -> None:
self._redis = redis
self._ops: list[tuple[str, tuple[Any, ...]]] = []

def delete(self, key: str) -> "FakePipeline":
self._ops.append(("delete", (key,)))
return self

def execute(self) -> list[Any]:
for op, args in self._ops:
getattr(self._redis, f"_do_{op}")(*args)
return [None] * len(self._ops)


class FakeRedis:
def __init__(
self, members: set[bytes] | None = None, fail: bool = False
) -> None:
self._members = members or set()
self._fail = fail
self._store: dict[str, bytes] = {}

def smembers(self, _key: str) -> set[bytes]:
if self._fail:
raise ConnectionError("redis unreachable")
return self._members

def pipeline(self) -> FakePipeline:
return FakePipeline(self)

def _do_delete(self, key: str) -> None:
self._store.pop(key, None)


def test_returns_zero_when_user_has_no_tokens(monkeypatch: Any) -> None:
fake = FakeRedis(members=set())
monkeypatch.setattr(token_store, "get_redis", lambda: fake)
assert token_store.revoke_all_user_refresh_tokens("user-1") == 0


def test_returns_count_when_tokens_revoked(monkeypatch: Any) -> None:
fake = FakeRedis(members={b"jti-a", b"jti-b", b"jti-c"})
monkeypatch.setattr(token_store, "get_redis", lambda: fake)
assert token_store.revoke_all_user_refresh_tokens("user-1") == 3


def test_returns_none_not_zero_when_redis_fails(monkeypatch: Any) -> None:
fake = FakeRedis(fail=True)
monkeypatch.setattr(token_store, "get_redis", lambda: fake)
result = token_store.revoke_all_user_refresh_tokens("user-1")
# The critical assertion: a Redis failure must be distinguishable
# from "there was nothing to revoke". Returning 0 here would let a
# caller (e.g. password-reset confirm) believe revocation succeeded
# when it didn't run at all.
assert result is None
assert result != 0
Loading