From ce0dc12dc290f7a5e74496b8452b396ad171d669 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Mon, 20 Jul 2026 13:09:23 +0100 Subject: [PATCH 01/25] feat(session): add EvictOverflowSessions and LockUserSessions queries --- db/queries/session.sql | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/db/queries/session.sql b/db/queries/session.sql index 83d7a2b..8699a66 100644 --- a/db/queries/session.sql +++ b/db/queries/session.sql @@ -57,3 +57,23 @@ WHERE expires_at < NOW(); -- name: CountUserSessions :one SELECT COUNT(*) FROM user_sessions WHERE user_id = $1; + +-- name: lock_user_sessions :exec +SELECT pg_advisory_xact_lock(hashtext(sqlc.arg(user_id)::text)::bigint); + +-- name: evict_overflow_sessions :many +WITH overflow AS ( + SELECT GREATEST(0, COUNT(*) - sqlc.arg(session_limit)) AS n + FROM user_sessions AS count_s + WHERE count_s.user_id = sqlc.arg(user_id) +) +DELETE FROM user_sessions AS outer_s +WHERE outer_s.id IN ( + SELECT inner_s.id + FROM user_sessions AS inner_s + WHERE inner_s.user_id = sqlc.arg(user_id) AND inner_s.id != sqlc.arg(id) + ORDER BY inner_s.last_active ASC, inner_s.created_at ASC + LIMIT (SELECT n FROM overflow) + FOR UPDATE SKIP LOCKED +) +RETURNING outer_s.id; \ No newline at end of file From 54cb0754607eed5bcef2843c619a4cfcb3f47bf5 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Mon, 20 Jul 2026 13:09:55 +0100 Subject: [PATCH 02/25] chore(gen): regenerate session querier from updated SQL --- db/generated/session.py | 34 +++++++++++++++++++++++++++++++++- 1 file changed, 33 insertions(+), 1 deletion(-) diff --git a/db/generated/session.py b/db/generated/session.py index fb9528a..d12c95a 100644 --- a/db/generated/session.py +++ b/db/generated/session.py @@ -4,7 +4,7 @@ # source: session.sql import dataclasses import datetime -from typing import AsyncIterator, Optional +from typing import Any, AsyncIterator, Optional import uuid import sqlalchemy @@ -43,6 +43,25 @@ """ +EVICT_OVERFLOW_SESSIONS = """-- name: evict_overflow_sessions \\:many +WITH overflow AS ( + SELECT GREATEST(0, COUNT(*) - :p3) AS n + FROM user_sessions AS count_s + WHERE count_s.user_id = :p1 +) +DELETE FROM user_sessions AS outer_s +WHERE outer_s.id IN ( + SELECT inner_s.id + FROM user_sessions AS inner_s + WHERE inner_s.user_id = :p1 AND inner_s.id != :p2 + ORDER BY inner_s.last_active ASC, inner_s.created_at ASC + LIMIT (SELECT n FROM overflow) + FOR UPDATE SKIP LOCKED +) +RETURNING outer_s.id +""" + + GET_SESSION_BY_DEVICE_FOR_USER = """-- name: get_session_by_device_for_user \\:one SELECT id, user_id, device_id, created_at, last_active, expires_at FROM user_sessions @@ -64,6 +83,11 @@ """ +LOCK_USER_SESSIONS = """-- name: lock_user_sessions \\:exec +SELECT pg_advisory_xact_lock(hashtext(:p1\\:\\:text)\\:\\:bigint) +""" + + UPDATE_SESSION_ACTIVITY = """-- name: update_session_activity \\:exec UPDATE user_sessions SET last_active = NOW() @@ -125,6 +149,11 @@ async def delete_session_by_device(self, *, device_id: uuid.UUID, user_id: uuid. async def delete_session_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(DELETE_SESSION_BY_ID), {"p1": id, "p2": user_id}) + async def evict_overflow_sessions(self, *, user_id: uuid.UUID, id: uuid.UUID, session_limit: Optional[Any]) -> AsyncIterator[uuid.UUID]: + result = await self._conn.stream(sqlalchemy.text(EVICT_OVERFLOW_SESSIONS), {"p1": user_id, "p2": id, "p3": session_limit}) + async for row in result: + yield row[0] + async def get_session_by_device_for_user(self, *, device_id: uuid.UUID, user_id: uuid.UUID) -> Optional[models.UserSession]: row = (await self._conn.execute(sqlalchemy.text(GET_SESSION_BY_DEVICE_FOR_USER), {"p1": device_id, "p2": user_id})).first() if row is None: @@ -163,6 +192,9 @@ async def list_sessions_by_user(self, *, user_id: uuid.UUID) -> AsyncIterator[mo expires_at=row[5], ) + async def lock_user_sessions(self, *, user_id: str) -> None: + await self._conn.execute(sqlalchemy.text(LOCK_USER_SESSIONS), {"p1": user_id}) + async def update_session_activity(self, *, id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(UPDATE_SESSION_ACTIVITY), {"p1": id}) From 660879f22b445aecda7c48f15a67e15a99e6b3bc Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Mon, 20 Jul 2026 13:10:13 +0100 Subject: [PATCH 03/25] feat(auth): session cap enforcement with SQL-native eviction --- app/service/users.py | 36 +++++++++++++----------------------- 1 file changed, 13 insertions(+), 23 deletions(-) diff --git a/app/service/users.py b/app/service/users.py index f5c8227..9979ca1 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -31,7 +31,7 @@ from db.generated import user as user_queries from db.generated import devices as device_queries from db.generated import session as session_queries -from db.generated.models import User, UserDevice, UserSession +from db.generated.models import User, UserDevice from app.core.logger import logger from app.service.face_embedding import FaceImagePayload, FaceEmbeddingService from app.schema.internal.single_face_match import ClosestUserMatch @@ -276,26 +276,9 @@ async def _create_mobile_session( ) -> MobileAuthResponse: user_id: uuid.UUID = user.id - device = await self._ensure_device_for_login(user_id, req) - - existing_session = await self.session_querier.get_session_by_device_for_user( - device_id=device.id, - user_id=user_id, - ) + await self.session_querier.lock_user_sessions(user_id=str(user_id)) - if existing_session is None: - sessions: list[UserSession] = [] - async for s in self.session_querier.list_sessions_by_user(user_id=user_id): - sessions.append(s) - - if len(sessions) >= AuthService.SESSION_LIMIT: - oldest = min(sessions, key=lambda s: (s.last_active, s.created_at)) - await SessionService.delete_session_cache(redis, oldest.id) - await self.session_querier.delete_session_by_id(id=oldest.id, user_id=user_id) - logger.warning( - "session_evicted user_id=%s evicted_session_id=%s", - user_id, oldest.id, - ) + device = await self._ensure_device_for_login(user_id, req) expires_at = datetime.now(timezone.utc) + timedelta( days=settings.MOBILE_SESSION_DAYS @@ -306,10 +289,18 @@ async def _create_mobile_session( device_id=device.id, expires_at=expires_at, ) - if not session: raise AppException.internal_error("Failed to create session") + async for evicted_id in self.session_querier.evict_overflow_sessions( + user_id=user_id, id=session.id, session_limit=AuthService.SESSION_LIMIT + ): + await SessionService.delete_session_cache(redis, evicted_id) + logger.warning( + "session_evicted user_id=%s evicted_session_id=%s", + user_id, evicted_id, + ) + access_token = create_acces_mobile_token(str(session.id)) refresh_token = create_refresh_mobile_token(str(session.id)) expiry = Get_expiry_time() @@ -323,7 +314,7 @@ async def _create_mobile_session( expires_at=session.expires_at, blocked=user.blocked, ttl=AuthService.REDIS_SESSION_TTL, - last_active=session.last_active + last_active=session.last_active, ) return MobileAuthResponse( @@ -334,7 +325,6 @@ async def _create_mobile_session( user_id=user_id, is_new_user=is_new_user, ) - async def refresh_token( self, redis: RedisClient, From b94409df493b957820f6af8957e63999bf978923 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Mon, 20 Jul 2026 13:10:44 +0100 Subject: [PATCH 04/25] test(integration): add concurrent login stress test --- .../test_session_device_management.py | 100 ++++++++++++++++++ 1 file changed, 100 insertions(+) diff --git a/tests/integration/test_session_device_management.py b/tests/integration/test_session_device_management.py index fc25fa8..d4c9381 100644 --- a/tests/integration/test_session_device_management.py +++ b/tests/integration/test_session_device_management.py @@ -8,6 +8,7 @@ actually being ON DELETE CASCADE. Both were previously verified by hand via psql; these tests make that verification automatic and regression-proof. """ +import asyncio import uuid from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock @@ -22,6 +23,7 @@ from db.generated import session as session_queries from db.generated import user as user_queries +pytestmark = pytest.mark.integration # =========================================================================== @@ -221,3 +223,101 @@ async def test_revoke_device_cascades_delete_session_real_db( ) await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) await db_conn.commit() + +@pytest.mark.asyncio +async def test_concurrent_new_device_logins_settle_at_cap_real_db( + auth_service: AuthService, + db_conn, +) -> None: + """Stress test for EvictOldestSessions with SKIP LOCKED: multiple + simultaneous logins from distinct new devices must never overshoot the + session cap, and the final session set must be exactly SESSION_LIMIT rows. + This is the only test that exercises the real Postgres locking behavior + that the design depends on — a fake cannot verify this.""" + from app.core.config import settings + from sqlalchemy.ext.asyncio import create_async_engine + + password = "ValidPass@123" + email = f"test-concurrent-{uuid.uuid4()}@multai.com" + physical_device_id = uuid.uuid4() + + user = await user_queries.AsyncQuerier(db_conn).create_user( + email=email, + hashed_password=hash_password(password), + ) + assert user is not None + user_id = user.id + + # Pre-seed one session so we start exactly at cap-1. + cap = AuthService.SESSION_LIMIT + pre_seed_device = await device_queries.AsyncQuerier(db_conn).create_device( + arg=device_queries.CreateDeviceParams( + column_1=None, + user_id=user_id, + device_name="Pre-seed Device", + device_type="android", + totp_secret=None, + physical_device_id=physical_device_id, + ) + ) + await session_queries.AsyncQuerier(db_conn).upsert_session( + user_id=user_id, + device_id=pre_seed_device.id, + expires_at=datetime.now(timezone.utc) + timedelta(days=7), + ) + + # CRITICAL: Commit the setup transaction so the user/device/session rows + # are visible to the separate connections used by concurrent tasks. + await db_conn.commit() + + assert cap >= 2, "SESSION_LIMIT must be >= 2 for this test to be meaningful" + concurrent_logins = cap + + # Need separate connections for true concurrency — asyncpg can't multiplex + # on a single connection. Each task gets its own connection from the engine. + url = ( + f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}" + f"@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" + ) + engine = create_async_engine(url, pool_pre_ping=True) + + async def _login_task(task_idx: int) -> None: + async with engine.connect() as conn: + task_auth = AuthService( + user_querier=user_queries.AsyncQuerier(conn), + session_querier=session_queries.AsyncQuerier(conn), + device_querier=device_queries.AsyncQuerier(conn), + face_embedding_service=auth_service.face_embedding_service, + ) + req = MobileLoginRequest( + email=email, + password=password, + device_name=f"Concurrent Device {task_idx}", + device_type="ios", + physical_device_id=uuid.uuid4(), + ) + await task_auth.mobile_login(_FakeRedis(), req) + await conn.commit() + + try: + await asyncio.gather(*(_login_task(i) for i in range(concurrent_logins))) + + count = ( + await db_conn.execute( + text("SELECT COUNT(*) FROM user_sessions WHERE user_id = :uid"), + {"uid": user_id}, + ) + ).scalar() + assert count == cap, ( + f"Expected exactly {cap} sessions after concurrent logins, got {count}" + ) + finally: + await engine.dispose() + await db_conn.execute( + text("DELETE FROM user_sessions WHERE user_id = :uid"), {"uid": user_id} + ) + await db_conn.execute( + text("DELETE FROM user_devices WHERE user_id = :uid"), {"uid": user_id} + ) + await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await db_conn.commit() From cb1254c0ef4d79074b42eab08854435aebe4cb10 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Mon, 20 Jul 2026 13:10:53 +0100 Subject: [PATCH 05/25] test(auth): update all fixtures/fakes for new eviction contract --- tests/unit/test_auth_email_otp.py | 11 ++- tests/unit/test_auth_service.py | 72 +++++++++++++------ tests/unit/test_mobile_auth_email_logging.py | 12 +++- .../test_mobile_auth_intent_validation.py | 22 +++++- 4 files changed, 87 insertions(+), 30 deletions(-) diff --git a/tests/unit/test_auth_email_otp.py b/tests/unit/test_auth_email_otp.py index 7736892..b1850e8 100644 --- a/tests/unit/test_auth_email_otp.py +++ b/tests/unit/test_auth_email_otp.py @@ -1,6 +1,6 @@ import uuid import json -from unittest.mock import AsyncMock, patch, ANY +from unittest.mock import AsyncMock, MagicMock, patch, ANY import pytest from app.service.users import AuthService @@ -100,7 +100,14 @@ async def test_verify_mobile_register_success( mock_user.blocked = False mock_user_querier.create_user.return_value = mock_user - mock_session_querier.count_user_sessions.return_value = 0 + mock_session_querier.lock_user_sessions = AsyncMock(return_value=None) + + async def _empty_evict(*, user_id, id, session_limit): + return + yield # pragma: no cover + + mock_session_querier.evict_overflow_sessions = MagicMock(side_effect=_empty_evict) + mock_session = AsyncMock() mock_session.id = uuid.uuid4() mock_session_querier.upsert_session.return_value = mock_session diff --git a/tests/unit/test_auth_service.py b/tests/unit/test_auth_service.py index 5fd342f..73e98d1 100644 --- a/tests/unit/test_auth_service.py +++ b/tests/unit/test_auth_service.py @@ -128,7 +128,13 @@ def device_querier() -> AsyncMock: def session_querier() -> AsyncMock: from db.generated import session as session_queries q = MagicMock(spec=session_queries.AsyncQuerier) - q.count_user_sessions = AsyncMock(return_value=0) + q.lock_user_sessions = AsyncMock(return_value=None) + + async def _default_empty_evict(*, user_id, id, session_limit): + return + yield # pragma: no cover + + q.evict_overflow_sessions = MagicMock(side_effect=_default_empty_evict) q.get_session_by_device_for_user = AsyncMock(return_value=None) q.list_sessions_by_user = MagicMock(return_value=_empty_async_iter()) q.delete_session_by_id = AsyncMock() @@ -287,38 +293,23 @@ async def test_blocked_user_raises_403( class TestSessionLimit: @pytest.mark.asyncio async def test_at_cap_evicts_oldest_and_succeeds( - self, - auth_service: AuthService, - user_querier: AsyncMock, - session_querier: AsyncMock, - redis: AsyncMock, + self, auth_service, user_querier, session_querier, redis, ) -> None: user = _make_user() user_querier.get_user_by_email.return_value = user - oldest = _make_session(user_id=user.id) - middle = _make_session(user_id=user.id) - newest = _make_session(user_id=user.id) + evicted_id = uuid.uuid4() - # Stagger so oldest is clearly the minimum - base = datetime.now(timezone.utc) - oldest.last_active = base - timedelta(seconds=2) - middle.last_active = base - timedelta(seconds=1) - newest.last_active = base + async def _evict(*, user_id, id, session_limit): + assert session_limit == AuthService.SESSION_LIMIT + yield evicted_id - async def _sessions(*, user_id): - yield oldest - yield middle - yield newest - - session_querier.list_sessions_by_user = _sessions + session_querier.evict_overflow_sessions = MagicMock(side_effect=_evict) result = await auth_service.mobile_login(redis, _make_login_request()) assert result.access_token - session_querier.delete_session_by_id.assert_called_once() - called_id = session_querier.delete_session_by_id.call_args.kwargs.get("id") - assert called_id == oldest.id + session_querier.evict_overflow_sessions.assert_called_once() @pytest.mark.asyncio async def test_within_session_limit_succeeds( @@ -337,6 +328,41 @@ async def test_within_session_limit_succeeds( session_querier.delete_session_by_id.assert_not_called() + @pytest.mark.asyncio + async def test_multiple_new_devices_at_cap_evict_exact_overflow( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """evict_overflow_sessions must be called with session_limit=SESSION_LIMIT, + and every session id it yields must trigger a Redis cache eviction.""" + user = _make_user() + user_querier.get_user_by_email.return_value = user + + evicted_ids = [uuid.uuid4(), uuid.uuid4(), uuid.uuid4()] + + async def _evict(*, user_id, id, session_limit): + assert session_limit == AuthService.SESSION_LIMIT + for eid in evicted_ids: + yield eid + + session_querier.evict_overflow_sessions = MagicMock(side_effect=_evict) + + result = await auth_service.mobile_login(redis, _make_login_request()) + + assert result.access_token + session_querier.evict_overflow_sessions.assert_called_once() + call_kwargs = session_querier.evict_overflow_sessions.call_args.kwargs + assert call_kwargs["session_limit"] == AuthService.SESSION_LIMIT + assert call_kwargs["user_id"] == user.id + # Redis delete must be called for each evicted session + assert redis.delete.call_count == 3 + + + + # =========================================================================== # 4. Logout # =========================================================================== diff --git a/tests/unit/test_mobile_auth_email_logging.py b/tests/unit/test_mobile_auth_email_logging.py index 3a59281..e1bb4eb 100644 --- a/tests/unit/test_mobile_auth_email_logging.py +++ b/tests/unit/test_mobile_auth_email_logging.py @@ -4,6 +4,7 @@ # mypy: disable-error-code=arg-type import asyncio +from collections.abc import AsyncIterator import logging import uuid from datetime import datetime, timezone @@ -70,8 +71,15 @@ class FakeSessionQuerier: def __init__(self, session: FakeSession) -> None: self._session = session - async def count_user_sessions(self, user_id: uuid.UUID) -> int: - return 0 + + async def lock_user_sessions(self, *, user_id: str) -> None: + return None + + async def evict_overflow_sessions( + self, *, user_id: uuid.UUID, id: uuid.UUID, session_limit: int + ) -> AsyncIterator[uuid.UUID]: + return + yield # pragma: no cover async def get_session_by_device_for_user( self, *, device_id: uuid.UUID, user_id: uuid.UUID diff --git a/tests/unit/test_mobile_auth_intent_validation.py b/tests/unit/test_mobile_auth_intent_validation.py index 99bc07a..b5b746e 100644 --- a/tests/unit/test_mobile_auth_intent_validation.py +++ b/tests/unit/test_mobile_auth_intent_validation.py @@ -114,9 +114,6 @@ class FakeSessionQuerier: def __init__(self) -> None: self._sessions: dict[tuple[uuid.UUID, uuid.UUID], FakeSession] = {} - async def count_user_sessions(self, user_id: uuid.UUID) -> int: - return sum(1 for (u, _d) in self._sessions if u == user_id) - async def get_session_by_device_for_user( self, *, device_id: uuid.UUID, user_id: uuid.UUID ) -> FakeSession | None: @@ -142,6 +139,25 @@ async def delete_session_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> No if key_to_remove: del self._sessions[key_to_remove] + async def lock_user_sessions(self, *, user_id: str) -> None: + return None + + async def evict_overflow_sessions( + self, *, user_id: uuid.UUID, id: uuid.UUID, session_limit: int + ) -> AsyncIterator[uuid.UUID]: + candidates = [ + s for (u, _d), s in list(self._sessions.items()) + if u == user_id and s.id != id + ] + # +1 accounts for the current session itself, which isn't in `candidates` + # but does count toward the real COUNT(*) the SQL version computes. + overflow = max(0, (len(candidates) + 1) - session_limit) + candidates.sort(key=lambda s: (s.last_active, s.created_at)) + for s in candidates[:overflow]: + key = next(k for k, v in self._sessions.items() if v is s) + del self._sessions[key] + yield s.id + async def upsert_session( self, *, From 9d78135c882b8304dddfe2b9cf51c59163a14dc5 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Fri, 24 Jul 2026 02:13:18 +0100 Subject: [PATCH 06/25] Move blocked-user re-check under row lock to close login race with block_user --- app/service/users.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/app/service/users.py b/app/service/users.py index 9979ca1..6d3768f 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -124,10 +124,18 @@ async def mobile_login( if not verify_password(req.password, existing_user.hashed_password or ""): logger.warning("login attempt: invalid_credentials user_id=%s", existing_user.id) raise AppException.unauthorized("Invalid credentials") - logger.info("login success user_id=%s", existing_user.id) + + locked_user = await self.user_querier.get_user_by_id_for_update(id=existing_user.id) + if not locked_user: + raise AppException.unauthorized("User not found") + if locked_user.blocked: + logger.warning("login attempt: user_blocked_at_commit user_id=%s", locked_user.id) + raise AppException.forbidden("User is blocked") + + logger.info("login success user_id=%s", locked_user.id) return await self._create_mobile_session( redis=redis, - user=existing_user, + user=locked_user, req=req, is_new_user=False, ) From f36fa113ab8af44ea551cc1c8b38d0c3311c4315 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Fri, 24 Jul 2026 02:13:18 +0100 Subject: [PATCH 07/25] Add tests for blocked-user race fix: unit branch coverage and real concurrent block-vs-login trial --- .../test_session_device_management.py | 96 +++++++++++++++++++ tests/unit/test_auth_service.py | 70 ++++++++++++++ .../test_mobile_auth_intent_validation.py | 6 ++ tests/unit/test_mobile_auth_rate_limiting.py | 10 +- 4 files changed, 179 insertions(+), 3 deletions(-) diff --git a/tests/integration/test_session_device_management.py b/tests/integration/test_session_device_management.py index d4c9381..243f658 100644 --- a/tests/integration/test_session_device_management.py +++ b/tests/integration/test_session_device_management.py @@ -321,3 +321,99 @@ async def _login_task(task_idx: int) -> None: ) await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) await db_conn.commit() + +@pytest.mark.asyncio +async def test_concurrent_block_and_login_never_leaves_blocked_user_with_session( + db_conn, +) -> None: + """The race this test exists for: block_user and mobile_login racing on + the same user. Regardless of which wins the timing, a user that ends up + blocked must never retain an active session — that would mean a login + slipped through the row-lock re-check and created a session after + block_user's cleanup already ran. Repeated because this is a genuine + timing race, not deterministic on a single run.""" + from app.core.config import settings + from app.service.users import AuthService + from sqlalchemy.ext.asyncio import create_async_engine + + url = ( + f"postgresql+asyncpg://{settings.POSTGRES_USER}:{settings.POSTGRES_PASSWORD}" + f"@{settings.POSTGRES_HOST}:{settings.POSTGRES_PORT}/{settings.POSTGRES_DB}" + ) + engine = create_async_engine(url, pool_pre_ping=True) + + async def _run_one_trial() -> None: + password = "ValidPass@123" + email = f"test-block-race-{uuid.uuid4()}@multai.com" + + user = await user_queries.AsyncQuerier(db_conn).create_user( + email=email, hashed_password=hash_password(password), + ) + assert user is not None + user_id = user.id + await db_conn.commit() + + async def _login() -> None: + async with engine.connect() as conn: + svc = AuthService( + user_querier=user_queries.AsyncQuerier(conn), + session_querier=session_queries.AsyncQuerier(conn), + device_querier=device_queries.AsyncQuerier(conn), + face_embedding_service=MagicMock(), + ) + req = MobileLoginRequest( + email=email, password=password, + device_name="Race Device", device_type="android", + physical_device_id=uuid.uuid4(), + ) + try: + await svc.mobile_login(_FakeRedis(), req) + except Exception: + pass + await conn.commit() + + async def _block() -> None: + async with engine.connect() as conn: + svc = AuthService( + user_querier=user_queries.AsyncQuerier(conn), + session_querier=session_queries.AsyncQuerier(conn), + device_querier=device_queries.AsyncQuerier(conn), + face_embedding_service=MagicMock(), + ) + await svc.block_user(redis=_FakeRedis(), user_id=user_id) + await conn.commit() + + await asyncio.gather(_login(), _block()) + + blocked = ( + await db_conn.execute( + text("SELECT blocked FROM users WHERE id = :uid"), {"uid": user_id} + ) + ).scalar() + session_count = ( + await db_conn.execute( + text("SELECT COUNT(*) FROM user_sessions WHERE user_id = :uid"), + {"uid": user_id}, + ) + ).scalar() + + await db_conn.execute( + text("DELETE FROM user_sessions WHERE user_id = :uid"), {"uid": user_id} + ) + await db_conn.execute( + text("DELETE FROM user_devices WHERE user_id = :uid"), {"uid": user_id} + ) + await db_conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + await db_conn.commit() + + if blocked: + assert session_count == 0, ( + "A blocked user retained an active session — the row-lock " + "re-check in mobile_login did not close the race." + ) + + try: + for _ in range(20): + await _run_one_trial() + finally: + await engine.dispose() diff --git a/tests/unit/test_auth_service.py b/tests/unit/test_auth_service.py index 73e98d1..d4ea84a 100644 --- a/tests/unit/test_auth_service.py +++ b/tests/unit/test_auth_service.py @@ -105,6 +105,7 @@ def user_querier() -> AsyncMock: from db.generated import user as user_queries q = MagicMock(spec=user_queries.AsyncQuerier) q.get_user_by_email = AsyncMock(return_value=None) + q.get_user_by_id_for_update = AsyncMock(return_value=None) q.create_user = AsyncMock() q.get_user_by_id = AsyncMock() q.find_closest_user_by_embedding = AsyncMock(return_value=None) @@ -246,6 +247,7 @@ async def test_valid_credentials_return_tokens( ) -> None: existing = _make_user(password="Correctpass1!") user_querier.get_user_by_email.return_value = existing + user_querier.get_user_by_id_for_update.return_value = existing result = await auth_service.mobile_login( redis, _make_login_request(password="Correctpass1!") @@ -263,6 +265,7 @@ async def test_wrong_password_raises_401( ) -> None: existing = _make_user(password="Rightpassword1!") user_querier.get_user_by_email.return_value = existing + user_querier.get_user_by_id_for_update.return_value = existing with pytest.raises(HTTPException) as exc_info: await auth_service.mobile_login( @@ -297,6 +300,7 @@ async def test_at_cap_evicts_oldest_and_succeeds( ) -> None: user = _make_user() user_querier.get_user_by_email.return_value = user + user_querier.get_user_by_id_for_update.return_value = user evicted_id = uuid.uuid4() @@ -321,6 +325,7 @@ async def test_within_session_limit_succeeds( ) -> None: user = _make_user() user_querier.get_user_by_email.return_value = user + user_querier.get_user_by_id_for_update.return_value = user session_querier.list_sessions_by_user = MagicMock(return_value=_empty_async_iter()) result = await auth_service.mobile_login(redis, _make_login_request()) @@ -340,6 +345,7 @@ async def test_multiple_new_devices_at_cap_evict_exact_overflow( and every session id it yields must trigger a Redis cache eviction.""" user = _make_user() user_querier.get_user_by_email.return_value = user + user_querier.get_user_by_id_for_update.return_value = user evicted_ids = [uuid.uuid4(), uuid.uuid4(), uuid.uuid4()] @@ -509,3 +515,67 @@ async def test_returns_closest_user_match( assert result is not None assert result.user_id == row.id assert result.distance == 0.25 + +class TestBlockedUserRaceCondition: + @pytest.mark.asyncio + async def test_blocked_between_initial_check_and_lock_is_caught( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """Simulates the exact race: the first read sees an unblocked user, + but the row-locked re-read (as if block_user committed in between) + sees blocked=True. Login must still be rejected, and no session + may be created.""" + unblocked_snapshot = _make_user(blocked=False) + blocked_after_lock = _make_user( + user_id=unblocked_snapshot.id, blocked=True + ) + user_querier.get_user_by_email.return_value = unblocked_snapshot + user_querier.get_user_by_id_for_update.return_value = blocked_after_lock + + with pytest.raises(HTTPException) as exc_info: + await auth_service.mobile_login(redis, _make_login_request()) + + assert exc_info.value.status_code == 403 + session_querier.upsert_session.assert_not_called() + + @pytest.mark.asyncio + async def test_locked_row_read_is_used_for_session_creation( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """The locked re-read's user object must be what actually gets + passed forward — not the earlier, possibly-stale read.""" + stale = _make_user(email="stale@test.com") + fresh = _make_user(user_id=stale.id, email="fresh@test.com") + user_querier.get_user_by_email.return_value = stale + user_querier.get_user_by_id_for_update.return_value = fresh + + await auth_service.mobile_login(redis, _make_login_request()) + + cache_kwargs = redis.set.call_args + # cache_session_for_auth writes via redis.set with the payload as JSON; + # simplest reliable check is via the querier call itself: + user_querier.get_user_by_id_for_update.assert_called_once_with(id=stale.id) + + @pytest.mark.asyncio + async def test_missing_user_at_lock_time_raises_401( + self, + auth_service: AuthService, + user_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """Defensive case: user vanished between the two reads (e.g. deleted).""" + user_querier.get_user_by_email.return_value = _make_user() + user_querier.get_user_by_id_for_update.return_value = None + + with pytest.raises(HTTPException) as exc_info: + await auth_service.mobile_login(redis, _make_login_request()) + + assert exc_info.value.status_code == 401 \ No newline at end of file diff --git a/tests/unit/test_mobile_auth_intent_validation.py b/tests/unit/test_mobile_auth_intent_validation.py index b5b746e..a53ae74 100644 --- a/tests/unit/test_mobile_auth_intent_validation.py +++ b/tests/unit/test_mobile_auth_intent_validation.py @@ -65,8 +65,14 @@ async def get_user_by_email(self, email: str) -> FakeUser | None: async def get_user_by_id(self, id: uuid.UUID) -> FakeUser | None: if self._user.id == id: return self._user + for created in self._created_users.values(): + if created.id == id: + return created return None + async def get_user_by_id_for_update(self, id: uuid.UUID) -> FakeUser | None: + return await self.get_user_by_id(id=id) + async def create_user(self, *, email: str, hashed_password: str) -> FakeUser: new_user = FakeUser(email=email, exists=True) new_user.hashed_password = hashed_password diff --git a/tests/unit/test_mobile_auth_rate_limiting.py b/tests/unit/test_mobile_auth_rate_limiting.py index 960485b..7477f22 100644 --- a/tests/unit/test_mobile_auth_rate_limiting.py +++ b/tests/unit/test_mobile_auth_rate_limiting.py @@ -35,12 +35,16 @@ def __init__(self) -> None: self.hashed_password = hash_password("ValidPass@123") self.blocked = False - class FakeUserQuerier: - async def get_user_by_email(self, email: str) -> FakeUser: - return FakeUser() + def __init__(self) -> None: + self._user = FakeUser() + async def get_user_by_email(self, email: str) -> FakeUser: + return self._user + async def get_user_by_id_for_update(self, id: uuid.UUID) -> FakeUser: + return self._user + class FakeDeviceQuerier: pass From d8c69bbb08b6e2fbd41b6fc9406d64adc6f43e8f Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Fri, 24 Jul 2026 02:20:10 +0100 Subject: [PATCH 08/25] chore: removed unused test variables --- tests/unit/test_auth_service.py | 5 +---- tests/unit/test_mobile_auth_rate_limiting.py | 2 +- 2 files changed, 2 insertions(+), 5 deletions(-) diff --git a/tests/unit/test_auth_service.py b/tests/unit/test_auth_service.py index d4ea84a..6a78ae6 100644 --- a/tests/unit/test_auth_service.py +++ b/tests/unit/test_auth_service.py @@ -559,9 +559,6 @@ async def test_locked_row_read_is_used_for_session_creation( await auth_service.mobile_login(redis, _make_login_request()) - cache_kwargs = redis.set.call_args - # cache_session_for_auth writes via redis.set with the payload as JSON; - # simplest reliable check is via the querier call itself: user_querier.get_user_by_id_for_update.assert_called_once_with(id=stale.id) @pytest.mark.asyncio @@ -578,4 +575,4 @@ async def test_missing_user_at_lock_time_raises_401( with pytest.raises(HTTPException) as exc_info: await auth_service.mobile_login(redis, _make_login_request()) - assert exc_info.value.status_code == 401 \ No newline at end of file + assert exc_info.value.status_code == 401 diff --git a/tests/unit/test_mobile_auth_rate_limiting.py b/tests/unit/test_mobile_auth_rate_limiting.py index 7477f22..19c8f0c 100644 --- a/tests/unit/test_mobile_auth_rate_limiting.py +++ b/tests/unit/test_mobile_auth_rate_limiting.py @@ -44,7 +44,7 @@ async def get_user_by_email(self, email: str) -> FakeUser: async def get_user_by_id_for_update(self, id: uuid.UUID) -> FakeUser: return self._user - + class FakeDeviceQuerier: pass From b15d6ff90b99ad942cec97b163a7b9e0fc2ac191 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sat, 25 Jul 2026 23:26:28 +0100 Subject: [PATCH 09/25] chore: seperated client ip check --- app/deps/client_ip.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) create mode 100644 app/deps/client_ip.py diff --git a/app/deps/client_ip.py b/app/deps/client_ip.py new file mode 100644 index 0000000..cf5a668 --- /dev/null +++ b/app/deps/client_ip.py @@ -0,0 +1,15 @@ +from fastapi import Request +from app.core.config import settings + + +def get_client_ip(request: Request) -> str | None: + if settings.TRUST_PROXY_HEADERS: + forwarded_for = request.headers.get("x-forwarded-for") + if forwarded_for: + return forwarded_for.split(",", maxsplit=1)[0].strip() or None + + real_ip = request.headers.get("x-real-ip") + if real_ip: + return real_ip.strip() or None + + return request.client.host if request.client else None \ No newline at end of file From eb7b2f2b3b76fb804f9a0623c4a20313590f8100 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sat, 25 Jul 2026 23:27:02 +0100 Subject: [PATCH 10/25] chore: rewired client ip check --- app/deps/rate_limit.py | 31 ++++++++++++------------------- 1 file changed, 12 insertions(+), 19 deletions(-) diff --git a/app/deps/rate_limit.py b/app/deps/rate_limit.py index a21e205..b146adf 100644 --- a/app/deps/rate_limit.py +++ b/app/deps/rate_limit.py @@ -1,34 +1,27 @@ from fastapi import Request, HTTPException from typing import Callable +from app.deps.client_ip import get_client_ip from app.infra.redis import RedisClient -from app.core.config import settings - -def _get_client_ip(request: Request) -> str: - if settings.TRUST_PROXY_HEADERS: - forwarded_for = request.headers.get("x-forwarded-for") - if forwarded_for: - return forwarded_for.split(",", maxsplit=1)[0].strip() - real_ip = request.headers.get("x-real-ip") - if real_ip: - return real_ip.strip() - return request.client.host if request.client else "127.0.0.1" - +from app.core.logger import logger def RateLimiter(requests: int, window: int) -> Callable: async def _rate_limit_dependency(request: Request) -> None: - client_ip = _get_client_ip(request) - # For simplicity, IP based rate limit on the endpoint + client_ip = get_client_ip(request) or "127.0.0.1" path = request.url.path key = f"rate_limit:{path}:{client_ip}" redis = RedisClient.get_instance() - # Increment request count - current = await redis.incr(key) - if current == 1: - # Set expiry for the window if it's the first request - await redis.expire(key, window) + try: + current = await redis.incr(key) + if current == 1: + await redis.expire(key, window) + except HTTPException: + raise + except Exception: + logger.warning("rate_limit: redis unavailable, failing open for key=%s", key) + return if current > requests: raise HTTPException(status_code=429, detail="Too Many Requests") From 240cab2cc327b32a65f0dc744a571960a5edfed4 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sat, 25 Jul 2026 23:27:34 +0100 Subject: [PATCH 11/25] chore: rewired client up check --- app/router/mobile/auth.py | 29 ++++++++++------------------- 1 file changed, 10 insertions(+), 19 deletions(-) diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index b0aefb3..d91a9b3 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -9,6 +9,7 @@ from app.container import get_container, Container from app.core.config import settings from app.core.constant import AuditEventType +from app.deps.client_ip import get_client_ip from app.deps.token_auth import MobileUserSchema, get_current_mobile_user from app.deps.rate_limit import RateLimiter @@ -25,27 +26,13 @@ router = APIRouter(prefix="/auth") - -def _get_client_ip(request: Request) -> str | None: - if settings.TRUST_PROXY_HEADERS: - forwarded_for = request.headers.get("x-forwarded-for") - if forwarded_for: - return forwarded_for.split(",", maxsplit=1)[0].strip() or None - - real_ip = request.headers.get("x-real-ip") - if real_ip: - return real_ip.strip() or None - - return request.client.host if request.client else None - - @router.post("/register", response_model=RegisterPendingResponse, dependencies=[Depends(RateLimiter(requests=5, window=60))]) async def mobile_register( req: MobileRegisterRequest, request: Request, container: Container = Depends(get_container), ) -> RegisterPendingResponse: - client_ip = _get_client_ip(request) + client_ip = get_client_ip(request) result = await container.auth_service.mobile_register(container.redis, req, client_ip=client_ip) return result @@ -56,7 +43,7 @@ async def mobile_register_resend_otp( request: Request, container: Container = Depends(get_container), ) -> RegisterPendingResponse: - client_ip = _get_client_ip(request) + client_ip = get_client_ip(request) result = await container.auth_service.mobile_register_resend_otp(container.redis, req.email, client_ip=client_ip) return result @@ -67,7 +54,7 @@ async def mobile_register_verify( request: Request, container: Container = Depends(get_container), ) -> MobileAuthResponse: - client_ip = _get_client_ip(request) + client_ip = get_client_ip(request) result = await container.auth_service.verify_mobile_register(container.redis, req, client_ip=client_ip) await container.audit_service.create_record( event_type=AuditEventType.USER_SIGNUP, @@ -83,7 +70,7 @@ async def mobile_login( request: Request, container: Container = Depends(get_container), ) -> MobileAuthResponse: - client_ip = _get_client_ip(request) + client_ip = get_client_ip(request) result = await container.auth_service.mobile_login(container.redis, req, client_ip=client_ip) await container.audit_service.create_record( event_type=AuditEventType.USER_LOGIN, @@ -93,7 +80,11 @@ async def mobile_login( return result -@router.post("/refresh", response_model=MobileAuthResponse) +@router.post( + "/refresh", + response_model=MobileAuthResponse, + dependencies=[Depends(RateLimiter(requests=10, window=60))], +) async def refresh_token( req: RefreshTokenRequest, container: Container = Depends(get_container), From 7ba856c45c40bb1e2f1f7e8fb97f0f2a67d02ff4 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sat, 25 Jul 2026 23:29:40 +0100 Subject: [PATCH 12/25] fix: fixed ruff errors --- app/deps/client_ip.py | 2 +- app/router/mobile/auth.py | 1 - 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/app/deps/client_ip.py b/app/deps/client_ip.py index cf5a668..9bae50e 100644 --- a/app/deps/client_ip.py +++ b/app/deps/client_ip.py @@ -12,4 +12,4 @@ def get_client_ip(request: Request) -> str | None: if real_ip: return real_ip.strip() or None - return request.client.host if request.client else None \ No newline at end of file + return request.client.host if request.client else None diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index d91a9b3..a4e41ef 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -7,7 +7,6 @@ from uuid import UUID from app.container import get_container, Container -from app.core.config import settings from app.core.constant import AuditEventType from app.deps.client_ip import get_client_ip from app.deps.token_auth import MobileUserSchema, get_current_mobile_user From e231855ff8989da78f235ae4e3a2b77ba6660f51 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sat, 25 Jul 2026 23:30:37 +0100 Subject: [PATCH 13/25] feat: added mobile access and refresh token life time configs --- app/core/config.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/app/core/config.py b/app/core/config.py index 8726cd6..d6ab412 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -41,6 +41,10 @@ class Settings(BaseSettings): MOBILE_SESSION_DAYS: int = 7 SESSION_ACTIVITY_THROTTLE_SECONDS: int = 60 + # Mobile access/refresh token lifetimes + MOBILE_ACCESS_TOKEN_TTL_SECONDS: int = 900 + MOBILE_REFRESH_TOKEN_REUSE_GRACE_SECONDS: int = 30 + # Mobile auth validation defaults MOBILE_AUTH_PASSWORD_MIN_LEN: int = 8 MOBILE_AUTH_PASSWORD_MAX_LEN: int = 128 From bbec6580bcc437aa0c495c1cdb1b36ae917eb9c1 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sat, 25 Jul 2026 23:32:16 +0100 Subject: [PATCH 14/25] feat: added refresh token create and hash function + refresh cache payload encryption/decryption --- app/core/securite.py | 45 ++++++++++++++++++++++++++------------------ 1 file changed, 27 insertions(+), 18 deletions(-) diff --git a/app/core/securite.py b/app/core/securite.py index 77342a3..b9bc82a 100644 --- a/app/core/securite.py +++ b/app/core/securite.py @@ -1,15 +1,17 @@ import base64 import hashlib +import os from datetime import datetime, timedelta, timezone +import secrets from typing import Any, Literal import jwt +from cryptography.hazmat.primitives.ciphers.aead import AESGCM from passlib.context import CryptContext from pydantic import BaseModel, ConfigDict import pyotp from app.core.config import settings from app.core.exceptions import AppException from app.core.logger import logger - pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") @@ -56,25 +58,12 @@ def decode_access_mobile_token(token: str) -> dict[str, Any]: raise AppException.unauthorized("Invalid token") -def create_refresh_mobile_token(session_id: str) -> str: - payload: dict[str, Any] = { - "session_id": session_id, - "exp": int( - (datetime.now(timezone.utc) + timedelta(seconds=Get_expiry_time() * 4)).timestamp() - ), - } - return jwt.encode(payload, key=settings.jwt_secret, algorithm=settings.jwt_algorithm) - +def create_raw_refresh_token() -> str: + return secrets.token_urlsafe(32) -def decode_refresh_mobile_token(token: str) -> dict[str, Any]: - try: - payload = jwt.decode(token, key=settings.jwt_secret, algorithms=[settings.jwt_algorithm]) - return payload - except jwt.ExpiredSignatureError: - raise AppException.unauthorized("Token has expired") - except jwt.InvalidTokenError: - raise AppException.unauthorized("Invalid token") +def hash_refresh_token(raw_token: str) -> str: + return hashlib.sha256(raw_token.encode("utf-8")).hexdigest() def create_totp_secret() -> str: return pyotp.random_base32() @@ -100,6 +89,26 @@ def generate_Acces_token_stuff(user_id: str, role: str) -> str: } return jwt.encode(payload, key=settings.jwt_secret, algorithm=settings.jwt_algorithm) +def _get_refresh_cache_aesgcm() -> AESGCM: + key = base64.b64decode(settings.encryption_key) + return AESGCM(key) + +def encrypt_refresh_cache_payload(plaintext: str) -> str: + """Encrypt a JSON string for storage in Redis. Returns a base64 string + safe to store directly (nonce + ciphertext packed together).""" + aes = _get_refresh_cache_aesgcm() + nonce = os.urandom(12) + ciphertext = aes.encrypt(nonce, plaintext.encode("utf-8"), None) + return base64.b64encode(nonce + ciphertext).decode("utf-8") + +def decrypt_refresh_cache_payload(encoded: str) -> str: + """Reverse of encrypt_refresh_cache_payload. Raises on tampering or + wrong key — treat any exception as 'cache miss'.""" + aes = _get_refresh_cache_aesgcm() + raw = base64.b64decode(encoded) + nonce, ciphertext = raw[:12], raw[12:] + plaintext = aes.decrypt(nonce, ciphertext, None) + return plaintext.decode("utf-8") # class EmbeddingCrypto: From 24e5c6588b144d6b385cd0db4bc10b143fd6d18d Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 15/25] refactor(container): constructor-inject SessionService, wire in refresh_token_querier --- app/container.py | 10 +++---- app/service/session.py | 62 +++++++++++++++++++++++++----------------- 2 files changed, 42 insertions(+), 30 deletions(-) diff --git a/app/container.py b/app/container.py index 86e0d46..c845258 100644 --- a/app/container.py +++ b/app/container.py @@ -31,7 +31,7 @@ from db.generated import upload_request_photos as upload_request_photo_queries from db.generated import upload_requests as upload_request_queries from db.generated import user as user_queries - +from db.generated import refresh_token as refresh_token_queries from db.generated import events as event_queries from db.generated import event_participant as participant_queries from db.generated import notifications as notification_queries @@ -72,11 +72,10 @@ def __init__( self.event_querier = event_queries.AsyncQuerier(conn) self.participant_querier = participant_queries.AsyncQuerier(conn) self.stats_querier = stats_queries.AsyncQuerier(conn) + self.refresh_token_querier = refresh_token_queries.AsyncQuerier(conn) - # services - self.session_service = SessionService() - self.session_service.init( - session=self.session_querier, + self.session_service = SessionService( + session_querier=self.session_querier, redis=self.redis, ) @@ -90,6 +89,7 @@ def __init__( user_querier=self.user_querier, device_querier=self.device_querier, session_querier=self.session_querier, + refresh_token_querier=self.refresh_token_querier, face_embedding_service=self.face_embedding_service, ) diff --git a/app/service/session.py b/app/service/session.py index cdfb7f4..d09f86e 100644 --- a/app/service/session.py +++ b/app/service/session.py @@ -6,6 +6,7 @@ from datetime import datetime from app.infra.redis import RedisClient from app.core.constant import RedisKey +from app.core.logger import logger class MobileSessionCache(BaseModel): @@ -18,14 +19,13 @@ class MobileSessionCache(BaseModel): class SessionService: - session_querier: session_queries.AsyncQuerier - redis: RedisClient - - def init(self, session: session_queries.AsyncQuerier, redis: RedisClient) -> None: - self.session_querier = session + def __init__( + self, + session_querier: session_queries.AsyncQuerier, + redis: RedisClient, + ) -> None: + self.session_querier = session_querier self.redis = redis - SessionService.session_querier = session - SessionService.redis = redis @staticmethod async def cache_session_for_auth( @@ -36,7 +36,7 @@ async def cache_session_for_auth( expires_at: datetime, blocked: bool, ttl: int, - last_active: datetime + last_active: datetime, ) -> None: key = RedisKey.MobileSessionCache.value.format(session_id=session_id) payload = MobileSessionCache( @@ -45,9 +45,14 @@ async def cache_session_for_auth( email=email, expires_at=expires_at, blocked=blocked, - last_active=last_active + last_active=last_active, ) - await redis.set(key=key, value=payload.model_dump_json(), expire=ttl) + try: + await redis.set(key=key, value=payload.model_dump_json(), expire=ttl) + except Exception: + logger.warning( + "cache_session_for_auth: redis unavailable, session_id=%s", session_id + ) @staticmethod async def get_cached_session( @@ -55,7 +60,13 @@ async def get_cached_session( session_id: uuid.UUID, ) -> MobileSessionCache | None: key = RedisKey.MobileSessionCache.value.format(session_id=session_id) - raw = await redis.get(key) + try: + raw = await redis.get(key) + except Exception: + logger.warning( + "get_cached_session: redis unavailable, session_id=%s", session_id + ) + return None # caller falls through to Postgres if raw is None: return None return MobileSessionCache.model_validate_json(raw) @@ -66,32 +77,33 @@ async def delete_session_cache( session_id: uuid.UUID, ) -> None: key = RedisKey.MobileSessionCache.value.format(session_id=session_id) - await redis.delete(key) + try: + await redis.delete(key) + except Exception: + logger.warning( + "delete_session_cache: redis unavailable, session_id=%s", session_id + ) - @staticmethod - async def get_session_by_id(session_id: uuid.UUID) -> UserSession: + async def get_session_by_id(self, session_id: uuid.UUID) -> UserSession: try: - session = await SessionService.session_querier.get_session_by_id(id=session_id) + session = await self.session_querier.get_session_by_id(id=session_id) if session is None: - raise AppException.not_found("session Not found ") + raise AppException.not_found("session not found") return session except Exception as e: raise DBExceptionImpl.handle(e) - @staticmethod - async def delete_expired_sessions() -> None: + async def delete_expired_sessions(self) -> None: try: - await SessionService.session_querier.delete_expired_sessions() + await self.session_querier.delete_expired_sessions() except Exception as e: raise DBExceptionImpl.handle(e) - @staticmethod - async def count_user_sessions(user_id: uuid.UUID) -> int: + async def count_user_sessions(self, user_id: uuid.UUID) -> int: try: - count = await SessionService.session_querier.count_user_sessions(user_id=user_id) + count = await self.session_querier.count_user_sessions(user_id=user_id) if count is None: - raise AppException.internal_error("failed to count ") - else: - return count + raise AppException.internal_error("failed to count") + return count except Exception as e: raise DBExceptionImpl.handle(e) From e0c8dda9c207c88d52832c0fc917bb63f1e0ac88 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 16/25] test: pass refresh_token_querier through auth_service fixture --- tests/integration/test_enrollment_flow.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/integration/test_enrollment_flow.py b/tests/integration/test_enrollment_flow.py index 836f1ff..1968946 100644 --- a/tests/integration/test_enrollment_flow.py +++ b/tests/integration/test_enrollment_flow.py @@ -14,6 +14,7 @@ from app.service.users import AuthService from app.service.face_embedding import FaceImagePayload from db.generated import user as user_queries +from db.generated import refresh_token as refresh_token_queries # =========================================================================== @@ -50,6 +51,7 @@ def auth_service(mock_face_embedding: AsyncMock, db_conn) -> AuthService: session_querier=session_queries.AsyncQuerier(db_conn), device_querier=device_queries.AsyncQuerier(db_conn), face_embedding_service=mock_face_embedding, + refresh_token_querier=refresh_token_queries.AsyncQuerier(db_conn), ) From 6df9c610373ac5bb1d18933000b1c54c613db217 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 17/25] test: wire refresh_token_querier into all AuthService construction sites, fix over-escaped docstring --- tests/integration/test_session_device_management.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/tests/integration/test_session_device_management.py b/tests/integration/test_session_device_management.py index 243f658..9055ed1 100644 --- a/tests/integration/test_session_device_management.py +++ b/tests/integration/test_session_device_management.py @@ -20,6 +20,7 @@ from app.schema.request.mobile.auth import MobileLoginRequest from app.service.users import AuthService from db.generated import devices as device_queries +from db.generated import refresh_token as refresh_token_queries from db.generated import session as session_queries from db.generated import user as user_queries @@ -61,6 +62,7 @@ def auth_service(mock_face_embedding: AsyncMock, db_conn) -> AuthService: session_querier=session_queries.AsyncQuerier(db_conn), device_querier=device_queries.AsyncQuerier(db_conn), face_embedding_service=mock_face_embedding, + refresh_token_querier=refresh_token_queries.AsyncQuerier(db_conn), ) @@ -162,10 +164,10 @@ async def test_relogin_on_same_device_replaces_not_duplicates_real_db( async def test_revoke_device_cascades_delete_session_real_db( db_conn, ) -> None: - """Regression test verifying user_sessions_device_id_fkey is genuinely + r"""Regression test verifying user_sessions_device_id_fkey is genuinely ON DELETE CASCADE: deleting a device row via the real revoke_device query must also delete its session row, with no separate DELETE - needed. Previously verified once by hand via psql \\d user_sessions; + needed. Previously verified once by hand via psql \d user_sessions; this makes it automatic.""" email = f"test-revoke-{uuid.uuid4()}@multai.com" @@ -288,6 +290,7 @@ async def _login_task(task_idx: int) -> None: session_querier=session_queries.AsyncQuerier(conn), device_querier=device_queries.AsyncQuerier(conn), face_embedding_service=auth_service.face_embedding_service, + refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), ) req = MobileLoginRequest( email=email, @@ -360,6 +363,7 @@ async def _login() -> None: session_querier=session_queries.AsyncQuerier(conn), device_querier=device_queries.AsyncQuerier(conn), face_embedding_service=MagicMock(), + refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), ) req = MobileLoginRequest( email=email, password=password, @@ -379,6 +383,7 @@ async def _block() -> None: session_querier=session_queries.AsyncQuerier(conn), device_querier=device_queries.AsyncQuerier(conn), face_embedding_service=MagicMock(), + refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), ) await svc.block_user(redis=_FakeRedis(), user_id=user_id) await conn.commit() From eeb243093f8bcf007b35214312fe8b9225a51a23 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 18/25] test: add refresh_token_querier fixture and wire into auth_service --- tests/unit/test_auth_email_otp.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tests/unit/test_auth_email_otp.py b/tests/unit/test_auth_email_otp.py index b1850e8..7e6acd9 100644 --- a/tests/unit/test_auth_email_otp.py +++ b/tests/unit/test_auth_email_otp.py @@ -26,18 +26,24 @@ def mock_face_embedding_service() -> AsyncMock: def mock_redis() -> AsyncMock: return AsyncMock() +@pytest.fixture +def mock_refresh_token_querier() -> AsyncMock: + return AsyncMock() + @pytest.fixture def auth_service( mock_user_querier: AsyncMock, mock_device_querier: AsyncMock, mock_session_querier: AsyncMock, mock_face_embedding_service: AsyncMock, + mock_refresh_token_querier: AsyncMock, ) -> AuthService: return AuthService( user_querier=mock_user_querier, device_querier=mock_device_querier, session_querier=mock_session_querier, face_embedding_service=mock_face_embedding_service, + refresh_token_querier=mock_refresh_token_querier, ) @pytest.mark.asyncio From 95483e1f2f5ea5b610c0cb0d7dc8efeb1a315e63 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 19/25] test: update fixtures for opaque refresh tokens, wire refresh_token_querier --- tests/unit/test_mobile_auth_email_logging.py | 5 +++-- tests/unit/test_mobile_auth_intent_validation.py | 15 +++++++++++++-- 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/tests/unit/test_mobile_auth_email_logging.py b/tests/unit/test_mobile_auth_email_logging.py index e1bb4eb..d868f32 100644 --- a/tests/unit/test_mobile_auth_email_logging.py +++ b/tests/unit/test_mobile_auth_email_logging.py @@ -8,6 +8,7 @@ import logging import uuid from datetime import datetime, timezone +from unittest.mock import MagicMock import pytest @@ -133,6 +134,7 @@ def test_mobile_register_logs_without_plaintext_email( device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(session), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=MagicMock(), ) req = MobileRegisterRequest( @@ -148,8 +150,7 @@ async def _noop_cache_session_for_auth(**_: object) -> None: monkeypatch.setattr(SessionService, "cache_session_for_auth", _noop_cache_session_for_auth) monkeypatch.setattr(users_module, "create_acces_mobile_token", lambda _: "access") - monkeypatch.setattr(users_module, "create_refresh_mobile_token", lambda _: "refresh") - monkeypatch.setattr(users_module, "Get_expiry_time", lambda: 3600) + monkeypatch.setattr(users_module, "create_raw_refresh_token", lambda: "refresh") asyncio.run(service.mobile_register(FakeRedis(), req)) diff --git a/tests/unit/test_mobile_auth_intent_validation.py b/tests/unit/test_mobile_auth_intent_validation.py index a53ae74..1428ac2 100644 --- a/tests/unit/test_mobile_auth_intent_validation.py +++ b/tests/unit/test_mobile_auth_intent_validation.py @@ -11,6 +11,7 @@ import logging import uuid from datetime import datetime, timezone +from unittest.mock import AsyncMock import pytest from fastapi import HTTPException @@ -217,8 +218,7 @@ async def _noop_cache_session_for_auth(**_: object) -> None: monkeypatch.setattr(SessionService, "cache_session_for_auth", _noop_cache_session_for_auth) monkeypatch.setattr(users_module, "create_acces_mobile_token", lambda _: "access") - monkeypatch.setattr(users_module, "create_refresh_mobile_token", lambda _: "refresh") - monkeypatch.setattr(users_module, "Get_expiry_time", lambda: 3600) + monkeypatch.setattr(users_module, "create_raw_refresh_token", lambda: "refresh") def test_login_with_unknown_email_is_rejected() -> None: @@ -229,6 +229,7 @@ def test_login_with_unknown_email_is_rejected() -> None: device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -253,6 +254,7 @@ def test_register_with_existing_email_is_rejected() -> None: device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileRegisterRequest( @@ -279,6 +281,7 @@ def test_login_with_correct_credentials_succeeds( device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -307,6 +310,7 @@ def test_register_with_new_email_succeeds( device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileRegisterRequest( @@ -333,6 +337,7 @@ def test_register_then_login_same_device_succeeds( device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) _patch_token_helpers(monkeypatch) @@ -386,6 +391,7 @@ def test_login_with_wrong_password_fails() -> None: device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -415,6 +421,7 @@ def test_login_logs_correctly( device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -460,6 +467,7 @@ async def _raise_integrity_error(*args: Any, **kwargs: Any) -> Any: device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) verify_req = RegisterVerifyRequest( @@ -504,6 +512,7 @@ def test_session_device_id_matches_surrogate_pk_not_physical_id( device_querier=device_querier, session_querier=session_querier, face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) physical_id = uuid.uuid4() @@ -543,6 +552,7 @@ def test_relogin_on_existing_device_succeeds_even_at_session_cap( device_querier=device_querier, session_querier=session_querier, face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) _patch_token_helpers(monkeypatch) @@ -599,6 +609,7 @@ def test_same_physical_device_id_reuses_device_row( device_querier=device_querier, session_querier=session_querier, face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) _patch_token_helpers(monkeypatch) From 849baf0973526bc3edabce898b7584bfec240a87 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 20/25] test: add refresh_token_querier to AuthService construction --- tests/unit/test_mobile_auth_rate_limiting.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/unit/test_mobile_auth_rate_limiting.py b/tests/unit/test_mobile_auth_rate_limiting.py index 19c8f0c..e35d32c 100644 --- a/tests/unit/test_mobile_auth_rate_limiting.py +++ b/tests/unit/test_mobile_auth_rate_limiting.py @@ -1,6 +1,7 @@ import asyncio import uuid from typing import Any +from unittest.mock import MagicMock import pytest from fastapi import HTTPException @@ -65,6 +66,7 @@ def test_rate_limiting_triggered_after_max_attempts() -> None: device_querier=FakeDeviceQuerier(), session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=MagicMock(), ) # Stub session creation to avoid database / redis dependencies From 6dd1b3a22ccd6f1fe53ec7a1daa1d50be4d37f25 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 21/25] test(auth): add coverage for rate-limit fail-open, block/delete locking, and refresh-token grace-window behavior --- tests/unit/test_auth_service.py | 544 +++++++++++++++++++++++++++++++- 1 file changed, 535 insertions(+), 9 deletions(-) diff --git a/tests/unit/test_auth_service.py b/tests/unit/test_auth_service.py index 6a78ae6..d4b5ca5 100644 --- a/tests/unit/test_auth_service.py +++ b/tests/unit/test_auth_service.py @@ -151,6 +151,17 @@ def face_service() -> AsyncMock: svc.compute_average_embedding = AsyncMock(return_value=[0.1] * 512) return svc +@pytest.fixture +def refresh_token_querier() -> AsyncMock: + from db.generated import refresh_token as refresh_token_queries + q = MagicMock(spec=refresh_token_queries.AsyncQuerier) + q.get_refresh_token_by_hash_for_update = AsyncMock(return_value=None) + q.get_refresh_token_by_jti = AsyncMock(return_value=None) + q.create_refresh_token = AsyncMock() + q.revoke_refresh_token = AsyncMock() + q.revoke_all_user_refresh_tokens = AsyncMock() + q.mark_refresh_token_used = AsyncMock() + return q @pytest.fixture def redis() -> AsyncMock: @@ -169,12 +180,14 @@ def auth_service( device_querier: AsyncMock, session_querier: AsyncMock, face_service: AsyncMock, + refresh_token_querier: AsyncMock, ) -> AuthService: return AuthService( user_querier=user_querier, device_querier=device_querier, session_querier=session_querier, face_embedding_service=face_service, + refresh_token_querier=refresh_token_querier, ) @@ -190,6 +203,7 @@ async def test_new_user_is_created( auth_service: AuthService, user_querier: AsyncMock, redis: AsyncMock, + refresh_token_querier: AsyncMock, ) -> None: new_user = _make_user() user_querier.get_user_by_email.return_value = None @@ -207,6 +221,7 @@ async def test_pending_status_returned_on_register( auth_service: AuthService, user_querier: AsyncMock, redis: AsyncMock, + refresh_token_querier: AsyncMock, ) -> None: user_querier.get_user_by_email.return_value = None @@ -413,19 +428,49 @@ async def test_valid_refresh_returns_new_tokens( auth_service: AuthService, user_querier: AsyncMock, session_querier: AsyncMock, + refresh_token_querier: AsyncMock, redis: AsyncMock, ) -> None: - from app.core.securite import create_refresh_mobile_token + from app.core.securite import ( + create_raw_refresh_token, + hash_refresh_token, + decrypt_refresh_cache_payload, + ) session = _make_session() session_querier.get_session_by_id.return_value = session user_querier.get_user_by_id.return_value = _make_user(user_id=session.user_id) - refresh_token = create_refresh_mobile_token(str(session.id)) - result = await auth_service.refresh_token(redis, refresh_token) + raw_token = create_raw_refresh_token() + token_hash = hash_refresh_token(raw_token) + + row = MagicMock() + row.id = uuid.uuid4() + row.used = False + row.used_at = None + row.family_id = uuid.uuid4() + row.session_id = session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + refresh_token_querier.mark_refresh_token_used.return_value = row + + result = await auth_service.refresh_token(redis, raw_token) assert result.access_token assert result.refresh_token + refresh_token_querier.mark_refresh_token_used.assert_called_once_with(id=row.id) + refresh_token_querier.create_refresh_token.assert_called_once() + + # Verify the grace-window cache was written under the expected key, + # encrypted (not plaintext), and that it decrypts back to the response. + redis.set.assert_called_once() + call_args = redis.set.call_args + cache_key = call_args.args[0] if call_args.args else call_args.kwargs.get("key") + cache_value = call_args.args[1] if len(call_args.args) > 1 else call_args.kwargs.get("value") + + assert cache_key == f"refresh_retry:{token_hash}" + assert "access_token" not in cache_value # plaintext JSON would contain this literal key; encrypted payload must not + decrypted = decrypt_refresh_cache_payload(cache_value) + assert result.access_token in decrypted @pytest.mark.asyncio async def test_expired_session_raises_401( @@ -433,19 +478,28 @@ async def test_expired_session_raises_401( auth_service: AuthService, user_querier: AsyncMock, session_querier: AsyncMock, + refresh_token_querier: AsyncMock, redis: AsyncMock, ) -> None: - from app.core.securite import create_refresh_mobile_token + from app.core.securite import create_raw_refresh_token past_session = _make_session( expires_at=datetime.now(timezone.utc) - timedelta(days=1) ) session_querier.get_session_by_id.return_value = past_session - refresh_token = create_refresh_mobile_token(str(past_session.id)) + raw_token = create_raw_refresh_token() + row = MagicMock() + row.id = uuid.uuid4() + row.used = False + row.used_at = None + row.family_id = uuid.uuid4() + row.session_id = past_session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + refresh_token_querier.mark_refresh_token_used.return_value = row with pytest.raises(HTTPException) as exc_info: - await auth_service.refresh_token(redis, refresh_token) + await auth_service.refresh_token(redis, raw_token) assert exc_info.value.status_code == 401 @pytest.mark.asyncio @@ -454,9 +508,10 @@ async def test_blocked_user_on_refresh_raises_403( auth_service: AuthService, user_querier: AsyncMock, session_querier: AsyncMock, + refresh_token_querier: AsyncMock, redis: AsyncMock, ) -> None: - from app.core.securite import create_refresh_mobile_token + from app.core.securite import create_raw_refresh_token session = _make_session() session_querier.get_session_by_id.return_value = session @@ -464,18 +519,30 @@ async def test_blocked_user_on_refresh_raises_403( user_id=session.user_id, blocked=True ) - refresh_token = create_refresh_mobile_token(str(session.id)) + raw_token = create_raw_refresh_token() + row = MagicMock() + row.id = uuid.uuid4() + row.used = False + row.used_at = None + row.family_id = uuid.uuid4() + row.session_id = session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + refresh_token_querier.mark_refresh_token_used.return_value = row with pytest.raises(HTTPException) as exc_info: - await auth_service.refresh_token(redis, refresh_token) + await auth_service.refresh_token(redis, raw_token) assert exc_info.value.status_code == 403 @pytest.mark.asyncio async def test_invalid_refresh_token_raises_401( self, auth_service: AuthService, + refresh_token_querier: AsyncMock, redis: AsyncMock, ) -> None: + # Ensure the querier returns None so the token is rejected + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = None + with pytest.raises(HTTPException) as exc_info: await auth_service.refresh_token(redis, "completely.invalid.token") assert exc_info.value.status_code == 401 @@ -576,3 +643,462 @@ async def test_missing_user_at_lock_time_raises_401( await auth_service.mobile_login(redis, _make_login_request()) assert exc_info.value.status_code == 401 + +# =========================================================================== +# 7. check_rate_limit fail-open behavior +# =========================================================================== + + +class TestCheckRateLimitFailOpen: + @pytest.mark.asyncio + async def test_redis_outage_does_not_block_login( + self, + auth_service: AuthService, + user_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """A Redis failure during rate-limit checking must not prevent + login from proceeding — it should fail open, not crash the request.""" + user = _make_user(password="Correctpass1!") + user_querier.get_user_by_email.return_value = user + user_querier.get_user_by_id_for_update.return_value = user + + redis.incr = AsyncMock(side_effect=ConnectionError("redis unreachable")) + + result = await auth_service.mobile_login( + redis, _make_login_request(password="Correctpass1!") + ) + + assert result.access_token # login succeeded despite Redis being down + + @pytest.mark.asyncio + async def test_real_rate_limit_rejection_still_raises( + self, + auth_service: AuthService, + redis: AsyncMock, + ) -> None: + """Confirm the fail-open except clause doesn't accidentally swallow + the actual 429 rejection — only infra failures should be caught.""" + redis.incr = AsyncMock(return_value=999) # way over any reasonable limit + + with pytest.raises(HTTPException) as exc_info: + await auth_service.check_rate_limit(redis, "rate:test:key", max_requests=5, window_seconds=60) + assert exc_info.value.status_code == 429 + + +# =========================================================================== +# 8. block_user — lock ordering and session purge +# =========================================================================== + + +class TestBlockUser: + @pytest.mark.asyncio + async def test_takes_lock_before_mutating( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """block_user must call get_user_by_id_for_update (the lock) before + set_user_blocked — this is what serializes it against mobile_login.""" + target = _make_user() + user_querier.get_user_by_id_for_update.return_value = target + user_querier.set_user_blocked.return_value = _make_user( + user_id=target.id, blocked=True + ) + + call_order = [] + user_querier.get_user_by_id_for_update.side_effect = ( + lambda *a, **kw: call_order.append("lock") or target + ) + user_querier.set_user_blocked.side_effect = ( + lambda *a, **kw: call_order.append("mutate") or _make_user(user_id=target.id, blocked=True) + ) + + await auth_service.block_user(redis=redis, user_id=target.id) + + assert call_order == ["lock", "mutate"] + + @pytest.mark.asyncio + async def test_purges_all_sessions_and_invalidates_cache( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + target = _make_user() + user_querier.get_user_by_id_for_update.return_value = target + user_querier.set_user_blocked.return_value = _make_user( + user_id=target.id, blocked=True + ) + + session_ids = [uuid.uuid4(), uuid.uuid4()] + + async def _sessions(*, user_id): + for sid in session_ids: + s = MagicMock() + s.id = sid + yield s + + session_querier.list_sessions_by_user = MagicMock(side_effect=_sessions) + + await auth_service.block_user(redis=redis, user_id=target.id) + + session_querier.delete_all_user_sessions.assert_called_once_with(user_id=target.id) + assert redis.delete.call_count == len(session_ids) + + @pytest.mark.asyncio + async def test_missing_user_raises_404( + self, + auth_service: AuthService, + user_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + user_querier.get_user_by_id_for_update.return_value = None + + with pytest.raises(HTTPException) as exc_info: + await auth_service.block_user(redis=redis, user_id=uuid.uuid4()) + assert exc_info.value.status_code == 404 + + +# =========================================================================== +# 9. unblock_user +# =========================================================================== + + +class TestUnblockUser: + @pytest.mark.asyncio + async def test_unblocks_successfully( + self, + auth_service: AuthService, + user_querier: AsyncMock, + ) -> None: + target_id = uuid.uuid4() + user_querier.set_user_blocked.return_value = _make_user( + user_id=target_id, blocked=False + ) + + result = await auth_service.unblock_user(user_id=target_id) + + assert result.blocked is False + user_querier.set_user_blocked.assert_called_once_with(blocked=False, id=target_id) + + @pytest.mark.asyncio + async def test_missing_user_raises_404( + self, + auth_service: AuthService, + user_querier: AsyncMock, + ) -> None: + user_querier.set_user_blocked.return_value = None + + with pytest.raises(HTTPException) as exc_info: + await auth_service.unblock_user(user_id=uuid.uuid4()) + assert exc_info.value.status_code == 404 + + +# =========================================================================== +# 10. delete_user — lock ordering and session purge (same shape as block_user) +# =========================================================================== + + +class TestDeleteUser: + @pytest.mark.asyncio + async def test_takes_lock_before_deleting( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + target = _make_user() + + call_order = [] + user_querier.get_user_by_id_for_update.side_effect = ( + lambda *a, **kw: call_order.append("lock") or target + ) + user_querier.delete_user.side_effect = ( + lambda *a, **kw: call_order.append("delete") + ) + + await auth_service.delete_user(redis=redis, user_id=target.id) + + assert call_order == ["lock", "delete"] + + @pytest.mark.asyncio + async def test_purges_all_sessions_and_invalidates_cache( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + target = _make_user() + user_querier.get_user_by_id_for_update.return_value = target + + session_ids = [uuid.uuid4(), uuid.uuid4(), uuid.uuid4()] + + async def _sessions(*, user_id): + for sid in session_ids: + s = MagicMock() + s.id = sid + yield s + + session_querier.list_sessions_by_user = MagicMock(side_effect=_sessions) + + await auth_service.delete_user(redis=redis, user_id=target.id) + + session_querier.delete_all_user_sessions.assert_called_once_with(user_id=target.id) + assert redis.delete.call_count == len(session_ids) + + @pytest.mark.asyncio + async def test_missing_user_raises_404( + self, + auth_service: AuthService, + user_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + user_querier.get_user_by_id_for_update.return_value = None + + with pytest.raises(HTTPException) as exc_info: + await auth_service.delete_user(redis=redis, user_id=uuid.uuid4()) + assert exc_info.value.status_code == 404 + + +# =========================================================================== +# 11. Refresh token — grace window, §4.1a blocked re-check, encryption +# =========================================================================== + + +class TestRefreshTokenGraceWindow: + @pytest.mark.asyncio + async def test_used_token_within_grace_replays_cached_response( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """A used-but-within-grace token with a valid cached response should + replay it (idempotent retry), not treat it as new or as theft.""" + from app.core.securite import ( + create_raw_refresh_token, + encrypt_refresh_cache_payload, + ) + from app.schema.response.mobile.auth import MobileAuthResponse + + session = _make_session() + session_querier.get_session_by_id.return_value = session + user_querier.get_user_by_id.return_value = _make_user( + user_id=session.user_id, blocked=False + ) + + raw_token = create_raw_refresh_token() + cached_response = MobileAuthResponse( + access_token="cached-access-token", + refresh_token="cached-refresh-token", + session_id=str(session.id), + expires_in=900, + user_id=session.user_id, + ) + redis.get.return_value = encrypt_refresh_cache_payload( + cached_response.model_dump_json() + ) + + row = MagicMock() + row.id = uuid.uuid4() + row.used = True + row.used_at = datetime.now(timezone.utc) - timedelta(seconds=5) # within grace + row.family_id = uuid.uuid4() + row.session_id = session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + + result = await auth_service.refresh_token(redis, raw_token) + + assert result.access_token == "cached-access-token" + # must NOT have re-rotated — no new token row created for a replay + refresh_token_querier.create_refresh_token.assert_not_called() + + @pytest.mark.asyncio + async def test_used_token_within_grace_but_blocked_user_raises_403( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """§4.1a fix: even with a valid cached replay, a user blocked since + the original rotation must be rejected, not silently replayed.""" + from app.core.securite import ( + create_raw_refresh_token, + encrypt_refresh_cache_payload, + ) + from app.schema.response.mobile.auth import MobileAuthResponse + + session = _make_session() + session_querier.get_session_by_id.return_value = session + user_querier.get_user_by_id.return_value = _make_user( + user_id=session.user_id, blocked=True # blocked since original rotation + ) + + raw_token = create_raw_refresh_token() + cached_response = MobileAuthResponse( + access_token="cached-access-token", + refresh_token="cached-refresh-token", + session_id=str(session.id), + expires_in=900, + user_id=session.user_id, + ) + redis.get.return_value = encrypt_refresh_cache_payload( + cached_response.model_dump_json() + ) + + row = MagicMock() + row.id = uuid.uuid4() + row.used = True + row.used_at = datetime.now(timezone.utc) - timedelta(seconds=5) + row.family_id = uuid.uuid4() + row.session_id = session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + + with pytest.raises(HTTPException) as exc_info: + await auth_service.refresh_token(redis, raw_token) + assert exc_info.value.status_code == 403 + + @pytest.mark.asyncio + async def test_used_token_within_grace_but_cache_miss_raises_401( + self, + auth_service: AuthService, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """Regression test for the fall-through bug: a used token within + grace but with NO cached replay (Redis eviction/failure) must be + rejected outright, never silently re-rotated into new tokens.""" + from app.core.securite import create_raw_refresh_token + + raw_token = create_raw_refresh_token() + redis.get.return_value = None # cache miss + + row = MagicMock() + row.id = uuid.uuid4() + row.used = True + row.used_at = datetime.now(timezone.utc) - timedelta(seconds=5) # within grace + row.family_id = uuid.uuid4() + row.session_id = uuid.uuid4() + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + + with pytest.raises(HTTPException) as exc_info: + await auth_service.refresh_token(redis, raw_token) + assert exc_info.value.status_code == 401 + refresh_token_querier.create_refresh_token.assert_not_called() + + @pytest.mark.asyncio + async def test_used_token_outside_grace_revokes_session( + self, + auth_service: AuthService, + session_querier: AsyncMock, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """Reuse outside the grace window is treated as theft — the entire + session must be revoked.""" + from app.core.securite import create_raw_refresh_token + + raw_token = create_raw_refresh_token() + session = _make_session() + session_querier.get_session_by_id.return_value = session + + row = MagicMock() + row.id = uuid.uuid4() + row.used = True + row.used_at = datetime.now(timezone.utc) - timedelta(seconds=120) # well outside grace + row.family_id = uuid.uuid4() + row.session_id = session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + + with pytest.raises(HTTPException) as exc_info: + await auth_service.refresh_token(redis, raw_token) + + assert exc_info.value.status_code == 401 + session_querier.delete_session_by_id.assert_called_once_with( + id=session.id, user_id=session.user_id + ) + redis.delete.assert_called_once() + + @pytest.mark.asyncio + async def test_corrupted_cache_value_treated_as_cache_miss( + self, + auth_service: AuthService, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """A cached value that fails to decrypt (tampering, corruption, + wrong key) must be rejected, never trusted or crash the request.""" + from app.core.securite import create_raw_refresh_token + + raw_token = create_raw_refresh_token() + redis.get.return_value = "not-valid-encrypted-base64-data!!" + + row = MagicMock() + row.id = uuid.uuid4() + row.used = True + row.used_at = datetime.now(timezone.utc) - timedelta(seconds=5) + row.family_id = uuid.uuid4() + row.session_id = uuid.uuid4() + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + + with pytest.raises(HTTPException) as exc_info: + await auth_service.refresh_token(redis, raw_token) + assert exc_info.value.status_code == 401 + + @pytest.mark.asyncio + async def test_new_rotation_caches_encrypted_not_plaintext( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """The replay cache must never contain the raw token/response as + plaintext JSON — this is the fix for the Redis-plaintext-secret gap.""" + from app.core.securite import ( + create_raw_refresh_token, + hash_refresh_token, + decrypt_refresh_cache_payload, + ) + + session = _make_session() + session_querier.get_session_by_id.return_value = session + user_querier.get_user_by_id.return_value = _make_user(user_id=session.user_id) + + raw_token = create_raw_refresh_token() + token_hash = hash_refresh_token(raw_token) + + row = MagicMock() + row.id = uuid.uuid4() + row.used = False + row.used_at = None + row.family_id = uuid.uuid4() + row.session_id = session.id + refresh_token_querier.get_refresh_token_by_hash_for_update.return_value = row + refresh_token_querier.mark_refresh_token_used.return_value = row + + result = await auth_service.refresh_token(redis, raw_token) + + redis.set.assert_called_once() + call_args = redis.set.call_args + cache_key = call_args.args[0] if call_args.args else call_args.kwargs.get("key") + cache_value = call_args.args[1] if len(call_args.args) > 1 else call_args.kwargs.get("value") + + assert cache_key == f"refresh_retry:{token_hash}" + # plaintext JSON would contain this literal substring; encrypted payload must not + assert "access_token" not in cache_value + assert result.access_token not in cache_value + + decrypted = decrypt_refresh_cache_payload(cache_value) + assert result.access_token in decrypted From f8fdb500cf72612ef7d573735e0311fac2b97a56 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 22/25] fix(auth): close session/token security gaps in refresh, block, and rate-limit flows --- app/service/users.py | 204 +++++++++++++++++++++++++++++++------------ 1 file changed, 149 insertions(+), 55 deletions(-) diff --git a/app/service/users.py b/app/service/users.py index 6d3768f..9df39ef 100644 --- a/app/service/users.py +++ b/app/service/users.py @@ -11,9 +11,10 @@ hash_password, verify_password, create_acces_mobile_token, - create_refresh_mobile_token, - decode_refresh_mobile_token, - Get_expiry_time, + create_raw_refresh_token, + hash_refresh_token, + encrypt_refresh_cache_payload, + decrypt_refresh_cache_payload, ) from app.core.config import settings from app.infra.redis import RedisClient @@ -31,7 +32,8 @@ from db.generated import user as user_queries from db.generated import devices as device_queries from db.generated import session as session_queries -from db.generated.models import User, UserDevice +from db.generated.models import User, UserDevice, RefreshToken +from db.generated import refresh_token as refresh_token_queries from app.core.logger import logger from app.service.face_embedding import FaceImagePayload, FaceEmbeddingService from app.schema.internal.single_face_match import ClosestUserMatch @@ -42,19 +44,23 @@ class AuthService: user_querier: user_queries.AsyncQuerier device_querier: device_queries.AsyncQuerier session_querier: session_queries.AsyncQuerier + refresh_token_querier: refresh_token_queries.AsyncQuerier SESSION_LIMIT = settings.MOBILE_SESSION_LIMIT REDIS_SESSION_TTL = settings.MOBILE_SESSION_TTL_SECONDS + REFRESH_GRACE_SECONDS = settings.MOBILE_REFRESH_TOKEN_REUSE_GRACE_SECONDS def __init__( self, user_querier: user_queries.AsyncQuerier, device_querier: device_queries.AsyncQuerier, session_querier: session_queries.AsyncQuerier, + refresh_token_querier: refresh_token_queries.AsyncQuerier, face_embedding_service: FaceEmbeddingService, ): self.user_querier = user_querier self.device_querier = device_querier self.session_querier = session_querier + self.refresh_token_querier = refresh_token_querier self.face_embedding_service = face_embedding_service async def _ensure_device_for_login( @@ -310,8 +316,14 @@ async def _create_mobile_session( ) access_token = create_acces_mobile_token(str(session.id)) - refresh_token = create_refresh_mobile_token(str(session.id)) - expiry = Get_expiry_time() + + raw_refresh_token = create_raw_refresh_token() + await self.refresh_token_querier.create_refresh_token( + session_id=session.id, + family_id=uuid.uuid4(), + token_hash=hash_refresh_token(raw_refresh_token), + ) + expiry = settings.MOBILE_ACCESS_TOKEN_TTL_SECONDS logger.info("session_created session_id=%s user_id=%s", session.id, user_id) await SessionService.cache_session_for_auth( @@ -327,28 +339,90 @@ async def _create_mobile_session( return MobileAuthResponse( access_token=access_token, - refresh_token=refresh_token, + refresh_token=raw_refresh_token, session_id=str(session.id), expires_in=expiry, user_id=user_id, is_new_user=is_new_user, ) + + async def _handle_used_refresh_token( + self, + redis: RedisClient, + row: RefreshToken, + token_hash: str, + ) -> MobileAuthResponse: + """A `used=True` refresh token was presented. Returns a replayed + response if this is a benign grace-window retry with a cached + result, or raises if it's outside the grace window (theft) or + inside the grace window with no cached replay available (a used + token with nothing to replay is never treated as valid). + """ + within_grace = ( + row.used_at is not None + and (datetime.now(timezone.utc) - row.used_at) + <= timedelta(seconds=AuthService.REFRESH_GRACE_SECONDS) + ) + + if within_grace: + cache_key = f"refresh_retry:{token_hash}" + cached = await redis.get(cache_key) + if cached: + try: + decrypted = decrypt_refresh_cache_payload(cached) + except Exception: + # tampered, corrupted, or wrong key — treat exactly + # like a cache miss, never trust an undecryptable value + raise AppException.unauthorized("Invalid refresh token") + session_for_check = await self.session_querier.get_session_by_id(id=row.session_id) + if not session_for_check: + raise AppException.unauthorized("Session not found") + user_for_check = await self.user_querier.get_user_by_id(id=session_for_check.user_id) + if not user_for_check or user_for_check.blocked: + raise AppException.forbidden("User is blocked") + return MobileAuthResponse.model_validate_json(decrypted) + raise AppException.unauthorized("Invalid refresh token") + + logger.warning( + "refresh_token_reuse_detected family_id=%s session_id=%s", + row.family_id, row.session_id, + ) + session_for_revoke = await self.session_querier.get_session_by_id(id=row.session_id) + if session_for_revoke: + await self.session_querier.delete_session_by_id( + id=row.session_id, user_id=session_for_revoke.user_id + ) + await SessionService.delete_session_cache(redis, row.session_id) + raise AppException.unauthorized( + "Refresh token reuse detected; session revoked" + ) + async def refresh_token( self, redis: RedisClient, refresh_token: str, ) -> MobileAuthResponse: - payload = decode_refresh_mobile_token(refresh_token) - session_id = payload.get("session_id") + token_hash = hash_refresh_token(refresh_token) - if not session_id: + row = await self.refresh_token_querier.get_refresh_token_by_hash_for_update( + token_hash=token_hash + ) + if not row: raise AppException.unauthorized("Invalid refresh token") - session = await self.session_querier.get_session_by_id(id=uuid.UUID(session_id)) + if row.used: + return await self._handle_used_refresh_token(redis, row, token_hash) + + claimed = await self.refresh_token_querier.mark_refresh_token_used(id=row.id) + if claimed is None: + # Not expected to be reachable — the row lock above already + # serializes concurrent access to this row. + raise AppException.unauthorized("Invalid refresh token") + row = claimed + session = await self.session_querier.get_session_by_id(id=row.session_id) if not session: raise AppException.unauthorized("Session not found") - if session.expires_at < datetime.now(timezone.utc): raise AppException.unauthorized("Session expired") @@ -358,18 +432,30 @@ async def refresh_token( if user.blocked: raise AppException.forbidden("User is blocked") - new_access_token = create_acces_mobile_token(session_id) - new_refresh_token = create_refresh_mobile_token(session_id) - expiry = Get_expiry_time() + new_access_token = create_acces_mobile_token(str(session.id)) + new_raw_refresh_token = create_raw_refresh_token() + await self.refresh_token_querier.create_refresh_token( + session_id=session.id, + family_id=row.family_id, + token_hash=hash_refresh_token(new_raw_refresh_token), + ) - return MobileAuthResponse( + response = MobileAuthResponse( access_token=new_access_token, - refresh_token=new_refresh_token, - session_id=session_id, - expires_in=expiry, + refresh_token=new_raw_refresh_token, + session_id=str(session.id), + expires_in=settings.MOBILE_ACCESS_TOKEN_TTL_SECONDS, user_id=session.user_id, ) + await redis.set( + f"refresh_retry:{token_hash}", + encrypt_refresh_cache_payload(response.model_dump_json()), + expire=AuthService.REFRESH_GRACE_SECONDS, + ) + + return response + async def logout( self, redis: RedisClient, @@ -377,8 +463,8 @@ async def logout( session_id: str, ) -> dict[str, str]: sid = uuid.UUID(session_id) - await SessionService.delete_session_cache(redis, sid) await self.session_querier.delete_session_by_id(id=sid, user_id=uuid.UUID(user_id)) + await SessionService.delete_session_cache(redis, sid) return {"message": "Logged out successfully"} @@ -421,18 +507,15 @@ async def add_embbed_user( return user - async def validate_session( - self, - redis: RedisClient, - session_id: str, - ) -> bool: + async def validate_session(self, redis: RedisClient, session_id: str) -> bool: session = await self.session_querier.get_session_by_id(id=uuid.UUID(session_id)) - if not session: return False - if session.expires_at < datetime.now(timezone.utc): return False + user = await self.user_querier.get_user_by_id(id=session.user_id) + if not user or user.blocked: + return False return True async def get_user_by_id(self, user_id: uuid.UUID) -> User | None: @@ -566,35 +649,18 @@ async def delete_avatar_bytes(self, *, avatar_key: str) -> None: except Exception as exc: logger.warning("Failed to clean up orphaned avatar %s: %s", avatar_key, exc) - async def delete_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User: - try: - existing = await self.user_querier.get_user_by_id(id=user_id) - if not existing: - raise AppException.not_found("User not found") - - sessions = self.session_querier.list_sessions_by_user(user_id=user_id) - async for s in sessions: - await SessionService.delete_session_cache(redis=redis, session_id=s.id) - await self.session_querier.delete_all_user_sessions(user_id=user_id) - - await self.user_querier.delete_user(id=user_id) - - return existing - except Exception as exc: - logger.error("Failed to delete user: %s", exc) - raise DBException.handle(exc) - async def block_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User: try: + locked = await self.user_querier.get_user_by_id_for_update(id=user_id) + if not locked: + raise AppException.not_found("User not found") user = await self.user_querier.set_user_blocked(blocked=True, id=user_id) if not user: - raise AppException.not_found("User not found") - - sessions = self.session_querier.list_sessions_by_user(user_id=user_id) - async for s in sessions: - await SessionService.delete_session_cache(redis, s.id) + raise AppException.internal_error("Failed to block user") + session_ids = [s.id async for s in self.session_querier.list_sessions_by_user(user_id=user_id)] await self.session_querier.delete_all_user_sessions(user_id=user_id) - + for sid in session_ids: + await SessionService.delete_session_cache(redis, sid) return user except Exception as exc: logger.error("Failed to block user: %s", exc) @@ -610,6 +676,27 @@ async def unblock_user(self, *, user_id: uuid.UUID) -> User: logger.error("Failed to unblock user: %s", exc) raise DBException.handle(exc) + async def delete_user(self, *, redis: RedisClient, user_id: uuid.UUID) -> User: + try: + existing = await self.user_querier.get_user_by_id_for_update(id=user_id) + if not existing: + raise AppException.not_found("User not found") + + session_ids = [ + s.id + async for s in self.session_querier.list_sessions_by_user(user_id=user_id) + ] + await self.session_querier.delete_all_user_sessions(user_id=user_id) + await self.user_querier.delete_user(id=user_id) + + for sid in session_ids: + await SessionService.delete_session_cache(redis=redis, session_id=sid) + + return existing + except Exception as exc: + logger.error("Failed to delete user: %s", exc) + raise DBException.handle(exc) + async def find_closest_user(self, *, embedding_literal: str) -> ClosestUserMatch | None: row = await self.user_querier.find_closest_user_by_embedding( dollar_1=embedding_literal, @@ -625,10 +712,17 @@ async def check_rate_limit( max_requests: int, window_seconds: int, ) -> None: - """Enforce rate limiting using Redis INCR + EXPIRE.""" - current_count = await redis.incr(key) - if current_count == 1: - await redis.expire(key, window_seconds) + """Enforce rate limiting using Redis INCR + EXPIRE. Fails open if Redis is unavailable.""" + try: + current_count = await redis.incr(key) + if current_count == 1: + await redis.expire(key, window_seconds) + except HTTPException: + raise + except Exception: + logger.warning("check_rate_limit: redis unavailable, failing open for key=%s", key) + return + if current_count > max_requests: raise AppException.too_many_requests( "Too many requests. Please try again later.", From 6a1b02d8482317a46fbd0bb2a61b5e5e5caf7e8a Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:19:17 +0100 Subject: [PATCH 23/25] docs(env): add generation instructions for jwt_secret and encryption_key --- .env.example | 21 ++++++++++++++++++--- 1 file changed, 18 insertions(+), 3 deletions(-) diff --git a/.env.example b/.env.example index 22ec4e4..e01f5de 100644 --- a/.env.example +++ b/.env.example @@ -34,16 +34,31 @@ REDIS_PASSWORD= # ========================= PGADMIN_PORT=5050 -jwt_secret=super_secret_jwt_key +# Secret used to sign mobile access tokens and staff JWTs (HS256). +# Must be high-entropy, minimum 32 bytes to satisfy PyJWT's recommended +# HMAC key length for SHA-256 (RFC 7518 §3.2). +jwt_secret= jwt_algorithm=HS256 -encryption_key=super_secret_encryption_key + +# AES-256-GCM key used to encrypt the refresh-token grace-window replay +# cache before it's stored in Redis (see app/core/securite.py: +# encrypt_refresh_cache_payload / decrypt_refresh_cache_payload). +# Must be a base64-encoded 32-byte (256-bit) key. +encryption_key= + totp_issuer=MultiAI GOOGLE_CLIENT_ID= GOOGLE_CLIENT_SECRET= GOOGLE_REDIRECT_URI=http://127.0.0.1:8000/staff/drive/callback GOOGLE_OAUTH_SCOPES=https://www.googleapis.com/auth/drive.readonly openid email profile -FACE_ENCRYPTION_KEY=hkbribvfirirbvivbibvib + +# Key for the (currently dormant/commented-out) EmbeddingCrypto class in +# app/core/securite.py. Same format requirement as encryption_key above — +# base64-encoded 32-byte key — if this class is ever re-enabled, a weak +# placeholder value here will fail the same way an under-length +# encryption_key did. +FACE_ENCRYPTION_KEY= # CORS Configuration CORS_ORIGINS=["http://localhost:3000", "http://localhost:5173", "http://127.0.0.1:3000", "http://127.0.0.1:5173"] From fb453a3d5eb3f1cec219f7399b9746e85781536f Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:21:57 +0100 Subject: [PATCH 24/25] feat(db): add refresh_token queries and generated querier --- db/generated/refresh_token.py | 107 ++++++++++++++++++++++++++++++++++ db/queries/refresh_token.sql | 26 +++++++++ 2 files changed, 133 insertions(+) create mode 100644 db/generated/refresh_token.py create mode 100644 db/queries/refresh_token.sql diff --git a/db/generated/refresh_token.py b/db/generated/refresh_token.py new file mode 100644 index 0000000..a7ed63c --- /dev/null +++ b/db/generated/refresh_token.py @@ -0,0 +1,107 @@ +# Code generated by sqlc. DO NOT EDIT. +# versions: +# sqlc v1.31.1 +# source: refresh_token.sql +from typing import Optional +import uuid + +import sqlalchemy +import sqlalchemy.ext.asyncio + +from db.generated import models + + +CREATE_REFRESH_TOKEN = """-- name: create_refresh_token \\:one +INSERT INTO refresh_tokens ( + session_id, + family_id, + token_hash +) VALUES ( + :p1, :p2, :p3 +) +RETURNING id, session_id, family_id, token_hash, used, created_at, used_at +""" + + +GET_REFRESH_TOKEN_BY_HASH = """-- name: get_refresh_token_by_hash \\:one +SELECT id, session_id, family_id, token_hash, used, created_at, used_at +FROM refresh_tokens +WHERE token_hash = :p1 +""" + + +GET_REFRESH_TOKEN_BY_HASH_FOR_UPDATE = """-- name: get_refresh_token_by_hash_for_update \\:one +SELECT id, session_id, family_id, token_hash, used, created_at, used_at +FROM refresh_tokens +WHERE token_hash = :p1 +FOR UPDATE +""" + + +MARK_REFRESH_TOKEN_USED = """-- name: mark_refresh_token_used \\:one +UPDATE refresh_tokens +SET used = TRUE, used_at = NOW() +WHERE id = :p1 AND used = FALSE +RETURNING id, session_id, family_id, token_hash, used, created_at, used_at +""" + + +class AsyncQuerier: + def __init__(self, conn: sqlalchemy.ext.asyncio.AsyncConnection): + self._conn = conn + + async def create_refresh_token(self, *, session_id: uuid.UUID, family_id: uuid.UUID, token_hash: str) -> Optional[models.RefreshToken]: + row = (await self._conn.execute(sqlalchemy.text(CREATE_REFRESH_TOKEN), {"p1": session_id, "p2": family_id, "p3": token_hash})).first() + if row is None: + return None + return models.RefreshToken( + id=row[0], + session_id=row[1], + family_id=row[2], + token_hash=row[3], + used=row[4], + created_at=row[5], + used_at=row[6], + ) + + async def get_refresh_token_by_hash(self, *, token_hash: str) -> Optional[models.RefreshToken]: + row = (await self._conn.execute(sqlalchemy.text(GET_REFRESH_TOKEN_BY_HASH), {"p1": token_hash})).first() + if row is None: + return None + return models.RefreshToken( + id=row[0], + session_id=row[1], + family_id=row[2], + token_hash=row[3], + used=row[4], + created_at=row[5], + used_at=row[6], + ) + + async def get_refresh_token_by_hash_for_update(self, *, token_hash: str) -> Optional[models.RefreshToken]: + row = (await self._conn.execute(sqlalchemy.text(GET_REFRESH_TOKEN_BY_HASH_FOR_UPDATE), {"p1": token_hash})).first() + if row is None: + return None + return models.RefreshToken( + id=row[0], + session_id=row[1], + family_id=row[2], + token_hash=row[3], + used=row[4], + created_at=row[5], + used_at=row[6], + ) + + async def mark_refresh_token_used(self, *, id: uuid.UUID) -> Optional[models.RefreshToken]: + row = (await self._conn.execute(sqlalchemy.text(MARK_REFRESH_TOKEN_USED), {"p1": id})).first() + if row is None: + return None + return models.RefreshToken( + id=row[0], + session_id=row[1], + family_id=row[2], + token_hash=row[3], + used=row[4], + created_at=row[5], + used_at=row[6], + ) diff --git a/db/queries/refresh_token.sql b/db/queries/refresh_token.sql new file mode 100644 index 0000000..222b759 --- /dev/null +++ b/db/queries/refresh_token.sql @@ -0,0 +1,26 @@ +-- name: create_refresh_token :one +INSERT INTO refresh_tokens ( + session_id, + family_id, + token_hash +) VALUES ( + $1, $2, $3 +) +RETURNING *; + +-- name: get_refresh_token_by_hash :one +SELECT * +FROM refresh_tokens +WHERE token_hash = $1; + +-- name: mark_refresh_token_used :one +UPDATE refresh_tokens +SET used = TRUE, used_at = NOW() +WHERE id = $1 AND used = FALSE +RETURNING *; + +-- name: get_refresh_token_by_hash_for_update :one +SELECT id, session_id, family_id, token_hash, used, created_at, used_at +FROM refresh_tokens +WHERE token_hash = $1 +FOR UPDATE; \ No newline at end of file From f7a8374df6aa97a2f2303f19f7bf54b9a090f40a Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Sun, 26 Jul 2026 04:52:58 +0100 Subject: [PATCH 25/25] fix: switched encryption key to valide base 64 key --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b1543f8..a55a58f 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -51,7 +51,7 @@ jobs: MINIO_ROOT_PASSWORD: dummy MINIO_HOST: localhost jwt_secret: test_secret - encryption_key: test_encryption_key + encryption_key: MPCSXH0IYfkp8JTpUNH0vUVyDlUeP6OKI8kz5iK54mw= FACE_ENCRYPTION_KEY: test_face_encryption_key FIREBASE_CREDENTIALS_PATH: dummy.json services: