From ce0dc12dc290f7a5e74496b8452b396ad171d669 Mon Sep 17 00:00:00 2001 From: Tyjfre-j Date: Mon, 20 Jul 2026 13:09:23 +0100 Subject: [PATCH 1/5] 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 2/5] 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 3/5] 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 4/5] 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 5/5] 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, *,