diff --git a/.env.example b/.env.example index 22ec4e4b..e01f5de1 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"] diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b1543f8b..a55a58f9 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: diff --git a/app/container.py b/app/container.py index 86e0d46a..c845258c 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/core/config.py b/app/core/config.py index b04e3ef9..dea2f69a 100644 --- a/app/core/config.py +++ b/app/core/config.py @@ -39,6 +39,13 @@ class Settings(BaseSettings): MOBILE_SESSION_LIMIT: int = 3 MOBILE_SESSION_TTL_SECONDS: int = 180 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_SESSION_ABSOLUTE_DAYS: int = 30 + # Mobile auth validation defaults MOBILE_AUTH_PASSWORD_MIN_LEN: int = 8 MOBILE_AUTH_PASSWORD_MAX_LEN: int = 128 diff --git a/app/core/image_validation.py b/app/core/image_validation.py index 7e427c77..c4680400 100644 --- a/app/core/image_validation.py +++ b/app/core/image_validation.py @@ -29,6 +29,7 @@ def sanitise_filename(raw: str | None, extension: str) -> str: if not raw: return f"{prefix}.{extension}" name = re.sub(r'[\\/:*?"<>|\x00-\x1f]', "_", raw) + name = name.replace("..", "_") name = name.lstrip(".")[:128] return f"{prefix}_{name}" diff --git a/app/core/securite.py b/app/core/securite.py index 77342a36..b9bc82a0 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: diff --git a/app/deps/client_ip.py b/app/deps/client_ip.py new file mode 100644 index 00000000..9bae50e0 --- /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 diff --git a/app/deps/rate_limit.py b/app/deps/rate_limit.py index a21e2051..b146adf0 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") diff --git a/app/deps/token_auth.py b/app/deps/token_auth.py index c6c2ff59..a14de6c4 100644 --- a/app/deps/token_auth.py +++ b/app/deps/token_auth.py @@ -1,34 +1,35 @@ -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Annotated import uuid from fastapi import Depends, HTTPException from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer -from pydantic import BaseModel from app.container import get_container, Container from app.core.config import settings from app.core.securite import decode_access_mobile_token from app.infra.redis import RedisClient +from app.schema.response.mobile.auth import MobileUserSchema from app.service.session import MobileSessionCache, SessionService security = HTTPBearer() -class MobileUserSchema(BaseModel): - user_id: uuid.UUID - email: str - session_id: uuid.UUID - - async def get_current_mobile_user( credentials: Annotated[HTTPAuthorizationCredentials, Depends(security)], container: Annotated[Container, Depends(get_container)], ) -> MobileUserSchema: """ Dependency to get the current logged-in mobile user. - Fast path: Redis cache (0 DB queries). + Fast path: Redis cache hit. Usually 0 DB queries; occasionally 1 cheap, + throttled UPDATE to last_active (see SESSION_ACTIVITY_THROTTLE_SECONDS) — + this is not a full round trip through the slow path, just a single + indexed write on the request's existing connection. Slow path: Postgres fallback (2 DB queries) with cache re-population. + + idle_expires_at slides forward on each throttled activity refresh, up to + MOBILE_SESSION_DAYS from now, but is capped so it never exceeds + absolute_expires_at — the hard ceiling set at login that never moves. """ token = credentials.credentials payload = decode_access_mobile_token(token) @@ -45,10 +46,33 @@ async def get_current_mobile_user( redis, session_id ) if cached is not None: - if cached.expires_at < datetime.now(timezone.utc): + now = datetime.now(timezone.utc) + if cached.idle_expires_at < now or cached.absolute_expires_at < now: raise HTTPException(status_code=401, detail="Session expired") if cached.blocked: raise HTTPException(status_code=403, detail="User is blocked") + + if (now - cached.last_active).total_seconds() > settings.SESSION_ACTIVITY_THROTTLE_SECONDS: + new_idle_expires_at = min( + now + timedelta(days=settings.MOBILE_SESSION_DAYS), + cached.absolute_expires_at, + ) + await container.session_service.session_querier.update_session_activity( + id=cached.session_id, + idle_expires_at=new_idle_expires_at, + ) + await SessionService.cache_session_for_auth( + redis=redis, + session_id=cached.session_id, + user_id=cached.user_id, + email=cached.email, + idle_expires_at=new_idle_expires_at, + absolute_expires_at=cached.absolute_expires_at, + blocked=cached.blocked, + ttl=settings.MOBILE_SESSION_TTL_SECONDS, + last_active=now, + ) + return MobileUserSchema( user_id=cached.user_id, email=cached.email, @@ -60,8 +84,8 @@ async def get_current_mobile_user( if not session: raise HTTPException(status_code=401, detail="Session not found") - exp_ts = payload.get("exp") - if exp_ts and session.expires_at.timestamp() < exp_ts: + now = datetime.now(timezone.utc) + if session.idle_expires_at < now or session.absolute_expires_at < now: raise HTTPException(status_code=401, detail="Session expired") user = await container.auth_service.user_querier.get_user_by_id(id=session.user_id) @@ -70,15 +94,19 @@ async def get_current_mobile_user( if user.blocked: raise HTTPException(status_code=403, detail="User is blocked") - # Re-populate cache so next request hits Redis + # Re-populate cache so next request hits Redis. The session row was just + # fetched fresh from Postgres, so its last_active is already accurate — + # no extra write needed here, only cache population. await SessionService.cache_session_for_auth( redis=redis, session_id=session.id, user_id=session.user_id, email=user.email or "", - expires_at=session.expires_at, + idle_expires_at=session.idle_expires_at, + absolute_expires_at=session.absolute_expires_at, blocked=user.blocked, ttl=settings.MOBILE_SESSION_TTL_SECONDS, + last_active=session.last_active, ) return MobileUserSchema( diff --git a/app/router/mobile/auth.py b/app/router/mobile/auth.py index b0aefb36..ce444c00 100644 --- a/app/router/mobile/auth.py +++ b/app/router/mobile/auth.py @@ -7,8 +7,8 @@ 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 app.deps.rate_limit import RateLimiter @@ -25,27 +25,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 +42,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 +53,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 +69,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 +79,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), @@ -206,7 +196,8 @@ async def get_me( session_id=sessions_objs.id, device_id=sessions_objs.device_id, last_active=sessions_objs.last_active, - expires_at=sessions_objs.expires_at, + idle_expires_at=sessions_objs.idle_expires_at, + absolute_expires_at=sessions_objs.absolute_expires_at, ) return MeResponse( diff --git a/app/schema/request/mobile/auth.py b/app/schema/request/mobile/auth.py index 23a13717..dea9933a 100644 --- a/app/schema/request/mobile/auth.py +++ b/app/schema/request/mobile/auth.py @@ -20,7 +20,7 @@ class MobileAuthBaseRequest(BaseModel): min_length=1, max_length=settings.MOBILE_AUTH_DEVICE_TYPE_MAX_LEN, ) - device_id: UUID + physical_device_id: UUID @field_validator("email", mode="before") @classmethod diff --git a/app/schema/response/mobile/auth.py b/app/schema/response/mobile/auth.py index d00d26a8..67bf1398 100644 --- a/app/schema/response/mobile/auth.py +++ b/app/schema/response/mobile/auth.py @@ -13,7 +13,13 @@ class SessionSchema(BaseModel): session_id: uuid.UUID device_id: uuid.UUID last_active: datetime - expires_at: datetime + idle_expires_at: datetime + absolute_expires_at: datetime + +class MobileUserSchema(BaseModel): + user_id: uuid.UUID + email: str + session_id: uuid.UUID class UserSchema(BaseModel): id: uuid.UUID @@ -26,7 +32,6 @@ class MeResponse(BaseModel): devices: List[DeviceSchema] sessions: Optional[SessionSchema] - class RegisterPendingResponse(BaseModel): message: str status: str diff --git a/app/service/device.py b/app/service/device.py index 6bf69f17..b5840ea7 100644 --- a/app/service/device.py +++ b/app/service/device.py @@ -1,5 +1,4 @@ from db.generated import devices as device_queries -from app.core.securite import create_totp_secret import uuid from app.core.exceptions import DBException,AppException, DBExceptionImpl from db.generated.models import UserDevice @@ -11,30 +10,6 @@ class DeviceService: def init(self: "DeviceService", device_querier: device_queries.AsyncQuerier) -> None: self.device_querier = device_querier - async def create_device( - self: "DeviceService", - user_id: uuid.UUID, - device_name: str, - device_type: str, - id: uuid.UUID | None = None, - ) -> UserDevice | None: - try : - DeviceCount = await self.count_devices(user_id=user_id) - if DeviceCount >=3: - raise AppException.bad_request("You can only have 3 devices") - return await self.device_querier.create_device( - arg=device_queries.CreateDeviceParams( - column_1=id, - user_id=user_id, - device_name=device_name, - device_type=device_type, - totp_secret=create_totp_secret(), - ) - - ) - except Exception as e : - raise DBException.handle(e) - async def activate_device( self: "DeviceService", device_id: uuid.UUID, diff --git a/app/service/session.py b/app/service/session.py index 403af567..97b980f6 100644 --- a/app/service/session.py +++ b/app/service/session.py @@ -6,25 +6,27 @@ 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): session_id: uuid.UUID user_id: uuid.UUID email: str - expires_at: datetime + idle_expires_at: datetime + absolute_expires_at: datetime blocked: bool + last_active: datetime 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( @@ -32,19 +34,28 @@ async def cache_session_for_auth( session_id: uuid.UUID, user_id: uuid.UUID, email: str, - expires_at: datetime, + idle_expires_at: datetime, + absolute_expires_at: datetime, blocked: bool, ttl: int, + last_active: datetime, ) -> None: key = RedisKey.MobileSessionCache.value.format(session_id=session_id) payload = MobileSessionCache( session_id=session_id, user_id=user_id, email=email, - expires_at=expires_at, + idle_expires_at=idle_expires_at, + absolute_expires_at=absolute_expires_at, blocked=blocked, + 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( @@ -52,7 +63,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) @@ -63,32 +80,27 @@ 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: - try: - await SessionService.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) diff --git a/app/service/users.py b/app/service/users.py index 603d1c34..3864757b 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( @@ -62,26 +68,28 @@ async def _ensure_device_for_login( user_id: uuid.UUID, req: MobileAuthBaseRequest, ) -> UserDevice: - existing_device = await self.device_querier.get_device_by_id_any(id=req.device_id) + existing_device = await self.device_querier.get_device_by_physical_id( + user_id=user_id, + physical_device_id=req.physical_device_id + ) if existing_device: - if existing_device.user_id != user_id: - raise AppException.forbidden("Device already registered to another user") if existing_device.is_invalid_token: raise AppException.forbidden( "Device push token is invalid. Update the token before logging in." ) if not existing_device.is_active: - await self.device_querier.activate_device(id=req.device_id, user_id=user_id) + await self.device_querier.activate_device(id=existing_device.id, user_id=user_id) return existing_device device = await self.device_querier.create_device( arg=device_queries.CreateDeviceParams( - column_1=req.device_id, + column_1=None, user_id=user_id, device_name=req.device_name, device_type=req.device_type, totp_secret=None, + physical_device_id=req.physical_device_id ) ) if not device: @@ -122,10 +130,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, ) @@ -274,34 +290,41 @@ async def _create_mobile_session( ) -> MobileAuthResponse: user_id: uuid.UUID = user.id - session_count = await self.session_querier.count_user_sessions(user_id=user_id) - if session_count and session_count >= AuthService.SESSION_LIMIT: - logger.warning( - "session_limit_reached user_id=%s limit=%s", - user_id, - AuthService.SESSION_LIMIT, - ) - raise AppException.forbidden("Maximum session limit reached") + await self.session_querier.lock_user_sessions(user_id=str(user_id)) - device_id = req.device_id - expires_at = datetime.now(timezone.utc) + timedelta( - days=settings.MOBILE_SESSION_DAYS - ) + device = await self._ensure_device_for_login(user_id, req) - await self._ensure_device_for_login(user_id, req) + now = datetime.now(timezone.utc) + idle_expires_at = now + timedelta(days=settings.MOBILE_SESSION_DAYS) + absolute_expires_at = now + timedelta(days=settings.MOBILE_SESSION_ABSOLUTE_DAYS) session = await self.session_querier.upsert_session( user_id=user_id, - device_id=device_id, - expires_at=expires_at, + device_id=device.id, + idle_expires_at=idle_expires_at, + absolute_expires_at=absolute_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() + + 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( @@ -309,37 +332,101 @@ async def _create_mobile_session( session_id=session.id, user_id=user_id, email=user.email or "", - expires_at=session.expires_at, + idle_expires_at=session.idle_expires_at, + absolute_expires_at=session.absolute_expires_at, blocked=user.blocked, ttl=AuthService.REDIS_SESSION_TTL, + last_active=session.last_active, ) 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): + now = datetime.now(timezone.utc) + if session.idle_expires_at < now or session.absolute_expires_at < now: raise AppException.unauthorized("Session expired") user = await self.user_querier.get_user_by_id(id=session.user_id) @@ -348,18 +435,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, @@ -367,8 +466,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"} @@ -411,17 +510,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): + now = datetime.now(timezone.utc) + if session.idle_expires_at < now or session.absolute_expires_at < now: + 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 @@ -556,35 +653,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) @@ -600,6 +680,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, @@ -615,10 +716,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.", diff --git a/db/generated/devices.py b/db/generated/devices.py index 7da5fb91..ffb04426 100644 --- a/db/generated/devices.py +++ b/db/generated/devices.py @@ -33,11 +33,12 @@ user_id, device_name, device_type, - totp_secret + totp_secret, + physical_device_id ) VALUES ( - COALESCE(:p1, uuid_generate_v4()), :p2, :p3, :p4, :p5 + COALESCE(:p1, uuid_generate_v4()), :p2, :p3, :p4, :p5, :p6 ) -RETURNING id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token +RETURNING id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id """ @@ -48,6 +49,7 @@ class CreateDeviceParams: device_name: Optional[str] device_type: Optional[str] totp_secret: Optional[str] + physical_device_id: uuid.UUID DEACTIVATE_DEVICE = """-- name: deactivate_device \\:exec @@ -67,21 +69,33 @@ class CreateDeviceParams: """ +GET_ANY_DEVICE_BY_PHYSICAL_ID = """-- name: get_any_device_by_physical_id \\:many +SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id FROM user_devices +WHERE physical_device_id = :p1 +""" + + GET_DEVICE_BY_ID = """-- name: get_device_by_id \\:one -SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token from user_devices +SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id from user_devices WHERE id = :p1 AND user_id = :p2 """ GET_DEVICE_BY_ID_ANY = """-- name: get_device_by_id_any \\:one -SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token from user_devices +SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id from user_devices WHERE id = :p1 """ +GET_DEVICE_BY_PHYSICAL_ID = """-- name: get_device_by_physical_id \\:one +SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id FROM user_devices +WHERE user_id = :p1 AND physical_device_id = :p2 +""" + + LIST_USER_DEVICES = """-- name: list_user_devices \\:many -SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token +SELECT id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id FROM user_devices WHERE user_id = :p1 ORDER BY last_active DESC @@ -119,7 +133,7 @@ class CreateDeviceParams: is_invalid_token = FALSE WHERE id = :p1 AND user_id = :p3 -RETURNING id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token +RETURNING id, user_id, device_name, device_type, totp_secret, is_2fa_enabled, last_active, created_at, push_token, is_active, is_invalid_token, physical_device_id """ @@ -143,6 +157,7 @@ async def create_device(self, arg: CreateDeviceParams) -> Optional[models.UserDe "p3": arg.device_name, "p4": arg.device_type, "p5": arg.totp_secret, + "p6": arg.physical_device_id, })).first() if row is None: return None @@ -158,6 +173,7 @@ async def create_device(self, arg: CreateDeviceParams) -> Optional[models.UserDe push_token=row[8], is_active=row[9], is_invalid_token=row[10], + physical_device_id=row[11], ) async def deactivate_device(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None: @@ -166,6 +182,24 @@ async def deactivate_device(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None: async def enable_device2_fa(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(ENABLE_DEVICE2_FA), {"p1": id, "p2": user_id}) + async def get_any_device_by_physical_id(self, *, physical_device_id: uuid.UUID) -> AsyncIterator[models.UserDevice]: + result = await self._conn.stream(sqlalchemy.text(GET_ANY_DEVICE_BY_PHYSICAL_ID), {"p1": physical_device_id}) + async for row in result: + yield models.UserDevice( + id=row[0], + user_id=row[1], + device_name=row[2], + device_type=row[3], + totp_secret=row[4], + is_2fa_enabled=row[5], + last_active=row[6], + created_at=row[7], + push_token=row[8], + is_active=row[9], + is_invalid_token=row[10], + physical_device_id=row[11], + ) + async def get_device_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> Optional[models.UserDevice]: row = (await self._conn.execute(sqlalchemy.text(GET_DEVICE_BY_ID), {"p1": id, "p2": user_id})).first() if row is None: @@ -182,6 +216,7 @@ async def get_device_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> Option push_token=row[8], is_active=row[9], is_invalid_token=row[10], + physical_device_id=row[11], ) async def get_device_by_id_any(self, *, id: uuid.UUID) -> Optional[models.UserDevice]: @@ -200,6 +235,26 @@ async def get_device_by_id_any(self, *, id: uuid.UUID) -> Optional[models.UserDe push_token=row[8], is_active=row[9], is_invalid_token=row[10], + physical_device_id=row[11], + ) + + async def get_device_by_physical_id(self, *, user_id: uuid.UUID, physical_device_id: uuid.UUID) -> Optional[models.UserDevice]: + row = (await self._conn.execute(sqlalchemy.text(GET_DEVICE_BY_PHYSICAL_ID), {"p1": user_id, "p2": physical_device_id})).first() + if row is None: + return None + return models.UserDevice( + id=row[0], + user_id=row[1], + device_name=row[2], + device_type=row[3], + totp_secret=row[4], + is_2fa_enabled=row[5], + last_active=row[6], + created_at=row[7], + push_token=row[8], + is_active=row[9], + is_invalid_token=row[10], + physical_device_id=row[11], ) async def list_user_devices(self, *, user_id: uuid.UUID) -> AsyncIterator[models.UserDevice]: @@ -217,6 +272,7 @@ async def list_user_devices(self, *, user_id: uuid.UUID) -> AsyncIterator[models push_token=row[8], is_active=row[9], is_invalid_token=row[10], + physical_device_id=row[11], ) async def mark_device_token_invalid(self, *, push_token: Optional[str]) -> None: @@ -244,4 +300,5 @@ async def update_device_push_token(self, *, id: uuid.UUID, push_token: Optional[ push_token=row[8], is_active=row[9], is_invalid_token=row[10], + physical_device_id=row[11], ) diff --git a/db/generated/models.py b/db/generated/models.py index 9908c54d..21bac799 100644 --- a/db/generated/models.py +++ b/db/generated/models.py @@ -146,6 +146,17 @@ class ProcessingJob: completed_at: Optional[datetime.datetime] +@dataclasses.dataclass() +class RefreshToken: + id: uuid.UUID + session_id: uuid.UUID + family_id: uuid.UUID + token_hash: str + used: bool + created_at: datetime.datetime + used_at: Optional[datetime.datetime] + + @dataclasses.dataclass() class StaffDriveConnection: id: uuid.UUID @@ -261,6 +272,7 @@ class UserDevice: push_token: Optional[str] is_active: bool is_invalid_token: bool + physical_device_id: uuid.UUID @dataclasses.dataclass() @@ -279,4 +291,5 @@ class UserSession: device_id: uuid.UUID created_at: datetime.datetime last_active: datetime.datetime - expires_at: datetime.datetime + idle_expires_at: datetime.datetime + absolute_expires_at: datetime.datetime diff --git a/db/generated/refresh_token.py b/db/generated/refresh_token.py new file mode 100644 index 00000000..a7ed63c1 --- /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/generated/session.py b/db/generated/session.py index fb9528a1..5bb8b045 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 @@ -24,12 +24,6 @@ """ -DELETE_EXPIRED_SESSIONS = """-- name: delete_expired_sessions \\:exec -DELETE FROM user_sessions -WHERE expires_at < NOW() -""" - - DELETE_SESSION_BY_DEVICE = """-- name: delete_session_by_device \\:exec DELETE FROM user_sessions WHERE device_id = :p1 @@ -43,30 +37,55 @@ """ +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 +SELECT id, user_id, device_id, created_at, last_active, idle_expires_at, absolute_expires_at FROM user_sessions WHERE device_id = :p1 AND user_id = :p2 """ GET_SESSION_BY_ID = """-- name: get_session_by_id \\:one -SELECT id, user_id, device_id, created_at, last_active, expires_at +SELECT id, user_id, device_id, created_at, last_active, idle_expires_at, absolute_expires_at FROM user_sessions WHERE id = :p1 """ LIST_SESSIONS_BY_USER = """-- name: list_sessions_by_user \\:many -SELECT id, user_id, device_id, created_at, last_active, expires_at +SELECT id, user_id, device_id, created_at, last_active, idle_expires_at, absolute_expires_at FROM user_sessions WHERE user_id = :p1 """ +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() +SET last_active = NOW(), + idle_expires_at = :p2 WHERE id = :p1 """ @@ -75,20 +94,22 @@ INSERT INTO user_sessions ( user_id, device_id, - expires_at + idle_expires_at, + absolute_expires_at ) VALUES ( - :p1, :p2, :p3 + :p1, :p2, :p3, :p4 ) ON CONFLICT (user_id, device_id) DO UPDATE SET last_active = NOW(), - expires_at = EXCLUDED.expires_at + idle_expires_at = EXCLUDED.idle_expires_at RETURNING id, user_id, device_id, last_active, - expires_at, + idle_expires_at, + absolute_expires_at, created_at """ @@ -99,7 +120,8 @@ class UpsertSessionRow: user_id: uuid.UUID device_id: uuid.UUID last_active: datetime.datetime - expires_at: datetime.datetime + idle_expires_at: datetime.datetime + absolute_expires_at: datetime.datetime created_at: datetime.datetime @@ -116,15 +138,17 @@ async def count_user_sessions(self, *, user_id: uuid.UUID) -> Optional[int]: async def delete_all_user_sessions(self, *, user_id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(DELETE_ALL_USER_SESSIONS), {"p1": user_id}) - async def delete_expired_sessions(self) -> None: - await self._conn.execute(sqlalchemy.text(DELETE_EXPIRED_SESSIONS)) - async def delete_session_by_device(self, *, device_id: uuid.UUID, user_id: uuid.UUID) -> None: await self._conn.execute(sqlalchemy.text(DELETE_SESSION_BY_DEVICE), {"p1": device_id, "p2": user_id}) 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: @@ -135,7 +159,8 @@ async def get_session_by_device_for_user(self, *, device_id: uuid.UUID, user_id: device_id=row[2], created_at=row[3], last_active=row[4], - expires_at=row[5], + idle_expires_at=row[5], + absolute_expires_at=row[6], ) async def get_session_by_id(self, *, id: uuid.UUID) -> Optional[models.UserSession]: @@ -148,7 +173,8 @@ async def get_session_by_id(self, *, id: uuid.UUID) -> Optional[models.UserSessi device_id=row[2], created_at=row[3], last_active=row[4], - expires_at=row[5], + idle_expires_at=row[5], + absolute_expires_at=row[6], ) async def list_sessions_by_user(self, *, user_id: uuid.UUID) -> AsyncIterator[models.UserSession]: @@ -160,14 +186,23 @@ async def list_sessions_by_user(self, *, user_id: uuid.UUID) -> AsyncIterator[mo device_id=row[2], created_at=row[3], last_active=row[4], - expires_at=row[5], + idle_expires_at=row[5], + absolute_expires_at=row[6], ) - async def update_session_activity(self, *, id: uuid.UUID) -> None: - await self._conn.execute(sqlalchemy.text(UPDATE_SESSION_ACTIVITY), {"p1": id}) + 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, idle_expires_at: datetime.datetime) -> None: + await self._conn.execute(sqlalchemy.text(UPDATE_SESSION_ACTIVITY), {"p1": id, "p2": idle_expires_at}) - async def upsert_session(self, *, user_id: uuid.UUID, device_id: uuid.UUID, expires_at: datetime.datetime) -> Optional[UpsertSessionRow]: - row = (await self._conn.execute(sqlalchemy.text(UPSERT_SESSION), {"p1": user_id, "p2": device_id, "p3": expires_at})).first() + async def upsert_session(self, *, user_id: uuid.UUID, device_id: uuid.UUID, idle_expires_at: datetime.datetime, absolute_expires_at: datetime.datetime) -> Optional[UpsertSessionRow]: + row = (await self._conn.execute(sqlalchemy.text(UPSERT_SESSION), { + "p1": user_id, + "p2": device_id, + "p3": idle_expires_at, + "p4": absolute_expires_at, + })).first() if row is None: return None return UpsertSessionRow( @@ -175,6 +210,7 @@ async def upsert_session(self, *, user_id: uuid.UUID, device_id: uuid.UUID, expi user_id=row[1], device_id=row[2], last_active=row[3], - expires_at=row[4], - created_at=row[5], + idle_expires_at=row[4], + absolute_expires_at=row[5], + created_at=row[6], ) diff --git a/db/queries/devices.sql b/db/queries/devices.sql index 8c91a255..a5d71817 100644 --- a/db/queries/devices.sql +++ b/db/queries/devices.sql @@ -4,9 +4,10 @@ INSERT INTO user_devices ( user_id, device_name, device_type, - totp_secret + totp_secret, + physical_device_id ) VALUES ( - COALESCE($1, uuid_generate_v4()), $2, $3, $4, $5 + COALESCE($1, uuid_generate_v4()), $2, $3, $4, $5, $6 ) RETURNING *; @@ -76,3 +77,11 @@ SET is_invalid_token = TRUE, is_active = FALSE WHERE push_token = $1; + +-- name: GetDeviceByPhysicalId :one +SELECT * FROM user_devices +WHERE user_id = $1 AND physical_device_id = $2; + +-- name: GetAnyDeviceByPhysicalId :many +SELECT * FROM user_devices +WHERE physical_device_id = $1; \ No newline at end of file diff --git a/db/queries/refresh_token.sql b/db/queries/refresh_token.sql new file mode 100644 index 00000000..222b7592 --- /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 diff --git a/db/queries/session.sql b/db/queries/session.sql index 83d7a2b9..d8d6cb58 100644 --- a/db/queries/session.sql +++ b/db/queries/session.sql @@ -2,20 +2,22 @@ INSERT INTO user_sessions ( user_id, device_id, - expires_at + idle_expires_at, + absolute_expires_at ) VALUES ( - $1, $2, $3 + $1, $2, $3, $4 ) ON CONFLICT (user_id, device_id) DO UPDATE SET last_active = NOW(), - expires_at = EXCLUDED.expires_at + idle_expires_at = EXCLUDED.idle_expires_at RETURNING id, user_id, device_id, last_active, - expires_at, + idle_expires_at, + absolute_expires_at, created_at; -- name: GetSessionByDeviceForUser :one @@ -35,7 +37,8 @@ WHERE user_id = $1; -- name: UpdateSessionActivity :exec UPDATE user_sessions -SET last_active = NOW() +SET last_active = NOW(), + idle_expires_at = $2 WHERE id = $1; -- name: DeleteSessionByDevice :exec @@ -51,9 +54,25 @@ WHERE id = $1 AND user_id = $2; DELETE FROM user_sessions WHERE user_id = $1; --- name: DeleteExpiredSessions :exec -DELETE FROM user_sessions -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 diff --git a/migrations/sql/down/add_physical_device_id.sql b/migrations/sql/down/add_physical_device_id.sql new file mode 100644 index 00000000..d0ef0a39 --- /dev/null +++ b/migrations/sql/down/add_physical_device_id.sql @@ -0,0 +1,4 @@ +DROP INDEX IF EXISTS idx_user_devices_user_physical; + +ALTER TABLE user_devices + DROP COLUMN IF EXISTS physical_device_id; \ No newline at end of file diff --git a/migrations/sql/down/add_session_idle_absolute_expiry.sql b/migrations/sql/down/add_session_idle_absolute_expiry.sql new file mode 100644 index 00000000..ae8cb295 --- /dev/null +++ b/migrations/sql/down/add_session_idle_absolute_expiry.sql @@ -0,0 +1,5 @@ +ALTER TABLE user_sessions + DROP COLUMN absolute_expires_at; + +ALTER TABLE user_sessions + RENAME COLUMN idle_expires_at TO expires_at; \ No newline at end of file diff --git a/migrations/sql/down/create_refresh_tokens.sql b/migrations/sql/down/create_refresh_tokens.sql new file mode 100644 index 00000000..a72a1751 --- /dev/null +++ b/migrations/sql/down/create_refresh_tokens.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS refresh_tokens; \ No newline at end of file diff --git a/migrations/sql/up/add_physical_device_id.sql b/migrations/sql/up/add_physical_device_id.sql new file mode 100644 index 00000000..934cb1b3 --- /dev/null +++ b/migrations/sql/up/add_physical_device_id.sql @@ -0,0 +1,10 @@ +ALTER TABLE user_devices + ADD COLUMN physical_device_id UUID; + +UPDATE user_devices SET physical_device_id = id WHERE physical_device_id IS NULL; + +ALTER TABLE user_devices + ALTER COLUMN physical_device_id SET NOT NULL; + +CREATE UNIQUE INDEX idx_user_devices_user_physical + ON user_devices (user_id, physical_device_id); \ No newline at end of file diff --git a/migrations/sql/up/add_session_idle_absolute_expiry.sql b/migrations/sql/up/add_session_idle_absolute_expiry.sql new file mode 100644 index 00000000..98c51a25 --- /dev/null +++ b/migrations/sql/up/add_session_idle_absolute_expiry.sql @@ -0,0 +1,8 @@ +ALTER TABLE user_sessions + RENAME COLUMN expires_at TO idle_expires_at; + +ALTER TABLE user_sessions + ADD COLUMN absolute_expires_at timestamp with time zone NOT NULL DEFAULT (now() + interval '30 days'); + +ALTER TABLE user_sessions + ALTER COLUMN absolute_expires_at DROP DEFAULT; \ No newline at end of file diff --git a/migrations/sql/up/create_refresh_tokens.sql b/migrations/sql/up/create_refresh_tokens.sql new file mode 100644 index 00000000..a48d7caf --- /dev/null +++ b/migrations/sql/up/create_refresh_tokens.sql @@ -0,0 +1,13 @@ +CREATE TABLE IF NOT EXISTS refresh_tokens ( + id UUID PRIMARY KEY DEFAULT uuid_generate_v4(), + session_id UUID NOT NULL REFERENCES user_sessions(id) ON DELETE CASCADE, + family_id UUID NOT NULL, + token_hash TEXT NOT NULL, + used BOOLEAN NOT NULL DEFAULT FALSE, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + used_at TIMESTAMPTZ +); + +CREATE UNIQUE INDEX idx_refresh_tokens_hash ON refresh_tokens (token_hash); +CREATE INDEX idx_refresh_tokens_family ON refresh_tokens (family_id); +CREATE INDEX idx_refresh_tokens_session ON refresh_tokens (session_id); \ No newline at end of file diff --git a/migrations/versions/2c8a676ecccf_add_physical_device_id.py b/migrations/versions/2c8a676ecccf_add_physical_device_id.py new file mode 100644 index 00000000..72976f67 --- /dev/null +++ b/migrations/versions/2c8a676ecccf_add_physical_device_id.py @@ -0,0 +1,25 @@ +"""add_physical_device_id + +Revision ID: 2c8a676ecccf +Revises: f0fa13623f6c +Create Date: 2026-07-17 02:15:43.587673 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_down, run_sql_up + + +# revision identifiers, used by Alembic. +revision: str = '2c8a676ecccf' +down_revision: Union[str, Sequence[str], None] = 'f0fa13623f6c' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + run_sql_up("add_physical_device_id") + + +def downgrade() -> None: + run_sql_down("add_physical_device_id") \ No newline at end of file diff --git a/migrations/versions/46bceaf84bd8_add_session_idle_absolute_expiry.py b/migrations/versions/46bceaf84bd8_add_session_idle_absolute_expiry.py new file mode 100644 index 00000000..1c60a74c --- /dev/null +++ b/migrations/versions/46bceaf84bd8_add_session_idle_absolute_expiry.py @@ -0,0 +1,27 @@ +"""add_session_idle_absolute_expiry + +Revision ID: 46bceaf84bd8 +Revises: e49065cb125a +Create Date: 2026-07-27 13:03:49.830264 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_up + + +# revision identifiers, used by Alembic. +revision: str = '46bceaf84bd8' +down_revision: Union[str, Sequence[str], None] = 'e49065cb125a' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + run_sql_up("add_session_idle_absolute_expiry") + pass + + +def downgrade() -> None: + run_sql_up("add_session_idle_absolute_expiry") + pass diff --git a/migrations/versions/e49065cb125a_create_refresh_tokens.py b/migrations/versions/e49065cb125a_create_refresh_tokens.py new file mode 100644 index 00000000..4e9928e5 --- /dev/null +++ b/migrations/versions/e49065cb125a_create_refresh_tokens.py @@ -0,0 +1,25 @@ +"""create_refresh_tokens + +Revision ID: e49065cb125a +Revises: 2c8a676ecccf +Create Date: 2026-07-17 02:21:18.789969 + +""" +from typing import Sequence, Union + +from migrations.helper import run_sql_down, run_sql_up + + +# revision identifiers, used by Alembic. +revision: str = 'e49065cb125a' +down_revision: Union[str, Sequence[str], None] = '2c8a676ecccf' +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + run_sql_up("create_refresh_tokens") + + +def downgrade() -> None: + run_sql_down("create_refresh_tokens") diff --git a/tests/e2e/test_mobile_auth_intent_e2e.py b/tests/e2e/test_mobile_auth_intent_e2e.py index 9bdaa76f..f37ba8d9 100644 --- a/tests/e2e/test_mobile_auth_intent_e2e.py +++ b/tests/e2e/test_mobile_auth_intent_e2e.py @@ -29,7 +29,7 @@ def test_login_with_unknown_email_fails(self) -> None: "password": "anypassword", "device_name": "TestDevice", "device_type": "android", - "device_id": str(uuid.uuid4()), + "physical_device_id": str(uuid.uuid4()), } response = requests.post( f"{self.base_url}/user/auth/login", @@ -50,7 +50,7 @@ def test_register_with_existing_email_fails(self) -> None: "password": "ValidPass@123", "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } response1 = requests.post( f"{self.base_url}/user/auth/register", @@ -74,7 +74,7 @@ def test_register_with_existing_email_fails(self) -> None: "otp": otp, "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } verify_response = requests.post( f"{self.base_url}/user/auth/register/verify", @@ -104,7 +104,7 @@ def test_register_then_login_succeeds(self) -> None: "password": password, "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } register_response = requests.post( f"{self.base_url}/user/auth/register", @@ -124,7 +124,7 @@ def test_register_then_login_succeeds(self) -> None: "otp": otp, "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } verify_response = requests.post( f"{self.base_url}/user/auth/register/verify", @@ -141,7 +141,7 @@ def test_register_then_login_succeeds(self) -> None: "password": password, "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } login_response = requests.post( f"{self.base_url}/user/auth/login", @@ -168,7 +168,7 @@ def test_login_with_wrong_password_fails(self) -> None: "password": password, "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } register_response = requests.post( f"{self.base_url}/user/auth/register", @@ -187,7 +187,7 @@ def test_login_with_wrong_password_fails(self) -> None: "otp": otp, "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } verify_response = requests.post( f"{self.base_url}/user/auth/register/verify", @@ -202,7 +202,7 @@ def test_login_with_wrong_password_fails(self) -> None: "password": "WrongPass@123", "device_name": "TestDevice", "device_type": "android", - "device_id": device_id, + "physical_device_id": device_id, } response = requests.post( f"{self.base_url}/user/auth/login", @@ -218,7 +218,7 @@ def test_register_requires_password(self) -> None: "email": "user@example.com", "device_name": "TestDevice", "device_type": "android", - "device_id": str(uuid.uuid4()), + "physical_device_id": str(uuid.uuid4()), # Missing password } response = requests.post( @@ -234,7 +234,7 @@ def test_login_requires_device_type(self) -> None: "email": "user@example.com", "password": "ValidPass@123", "device_name": "TestDevice", - "device_id": str(uuid.uuid4()), + "physical_device_id": str(uuid.uuid4()), # Missing device_type } response = requests.post( diff --git a/tests/e2e/test_mobile_auth_request_validation_e2e.py b/tests/e2e/test_mobile_auth_request_validation_e2e.py index 96c7f0e2..6852bcc3 100644 --- a/tests/e2e/test_mobile_auth_request_validation_e2e.py +++ b/tests/e2e/test_mobile_auth_request_validation_e2e.py @@ -25,7 +25,7 @@ def _valid_payload() -> dict[str, object]: "password": "ValidPass@123", "device_name": "Pixel 8", "device_type": "android", - "device_id": str(uuid.uuid4()), + "physical_device_id": str(uuid.uuid4()), } diff --git a/tests/integration/test_enrollment_flow.py b/tests/integration/test_enrollment_flow.py index 836f1ff3..1968946b 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), ) diff --git a/tests/integration/test_session_device_management.py b/tests/integration/test_session_device_management.py new file mode 100644 index 00000000..f65ca98e --- /dev/null +++ b/tests/integration/test_session_device_management.py @@ -0,0 +1,426 @@ +""" +Integration tests for session & device management. + +These tests use a real PostgreSQL database (not fakes) specifically because +the behaviors under test depend on real SQL semantics that a fake cannot +verify: the UpsertSession ON CONFLICT clause actually matching the live +UNIQUE(user_id, device_id) constraint, and the user_sessions.device_id FK +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 + +import pytest +from sqlalchemy import text + +from app.core.securite import hash_password +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 + +pytestmark = pytest.mark.integration + + +# =========================================================================== +# Fixtures +# =========================================================================== + + +@pytest.fixture +async def db_conn(): + from app.core.config import settings + 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 with engine.connect() as conn: + yield conn + await engine.dispose() + + +@pytest.fixture +def mock_face_embedding() -> AsyncMock: + from app.service.face_embedding import FaceEmbeddingService + + svc = MagicMock(spec=FaceEmbeddingService) + return svc + + +@pytest.fixture +def auth_service(mock_face_embedding: AsyncMock, db_conn) -> AuthService: + return AuthService( + user_querier=user_queries.AsyncQuerier(db_conn), + 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), + ) + + +class _FakeRedis: + """Minimal redis stand-in — no rate-limit backing store or cache + assertions needed for these tests, just enough for the login flow + to complete without touching a real Redis instance.""" + + def __init__(self) -> None: + self._store: dict[str, int] = {} + + async def incr(self, key: str) -> int: + self._store[key] = self._store.get(key, 0) + 1 + return self._store[key] + + async def expire(self, key: str, seconds: int) -> None: + pass + + async def ttl(self, key: str) -> int: + return -1 + + async def set(self, key: str, value: str, expire: int) -> None: + return None + + async def get(self, key: str) -> str | None: + return None + + async def delete(self, key: str) -> None: + return None + + +# =========================================================================== +# Tests +# =========================================================================== + + +@pytest.mark.asyncio +async def test_relogin_on_same_device_replaces_not_duplicates_real_db( + auth_service: AuthService, + db_conn, +) -> None: + """Regression test against the real UpsertSession query: logging in twice + from the same (user, physical_device_id) must produce exactly one + session row with a stable id, not two rows — verifying the live + UNIQUE(user_id, device_id) constraint and ON CONFLICT clause actually + match, which a fake-based unit test cannot verify.""" + password = "ValidPass@123" + email = f"test-session-{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 + + try: + req = MobileLoginRequest( + email=email, + password=password, + device_name="Integration Test Device", + device_type="android", + physical_device_id=physical_device_id, + ) + + result1 = await auth_service.mobile_login(_FakeRedis(), req) + result2 = await auth_service.mobile_login(_FakeRedis(), req) + + assert result1.session_id == result2.session_id + + row = ( + await db_conn.execute( + text("SELECT COUNT(*) FROM user_sessions WHERE user_id = :uid"), + {"uid": user_id}, + ) + ).scalar() + assert row == 1 + + device_row = ( + await db_conn.execute( + text("SELECT COUNT(*) FROM user_devices WHERE user_id = :uid"), + {"uid": user_id}, + ) + ).scalar() + assert device_row == 1 + finally: + 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() + + +@pytest.mark.asyncio +async def test_revoke_device_cascades_delete_session_real_db( + db_conn, +) -> None: + 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; + this makes it automatic.""" + email = f"test-revoke-{uuid.uuid4()}@multai.com" + + user_querier = user_queries.AsyncQuerier(db_conn) + device_querier = device_queries.AsyncQuerier(db_conn) + session_querier = session_queries.AsyncQuerier(db_conn) + + user = await user_querier.create_user(email=email, hashed_password="hash") + assert user is not None + user_id = user.id + + try: + device = await device_querier.create_device( + arg=device_queries.CreateDeviceParams( + column_1=None, + user_id=user_id, + device_name="Integration Test Device", + device_type="android", + totp_secret=None, + physical_device_id=uuid.uuid4(), + ) + ) + assert device is not None + + session = await session_querier.upsert_session( + user_id=user_id, + device_id=device.id, + idle_expires_at=datetime.now(timezone.utc) + timedelta(days=7), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(days=30), + ) + assert session is not None + session_id = session.id + + # Confirm the session actually exists before we revoke. + pre = await session_querier.get_session_by_id(id=session_id) + assert pre is not None + + await device_querier.revoke_device(id=device.id, user_id=user_id) + + post = await session_querier.get_session_by_id(id=session_id) + assert post is None + + device_still_there = ( + await db_conn.execute( + text("SELECT COUNT(*) FROM user_devices WHERE id = :did"), + {"did": device.id}, + ) + ).scalar() + assert device_still_there == 0 + finally: + 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() + +@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, + idle_expires_at=datetime.now(timezone.utc) + timedelta(days=7), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(days=30), + ) + + # 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, + refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), + ) + 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() + +@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(), + refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), + ) + 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(), + refresh_token_querier=refresh_token_queries.AsyncQuerier(conn), + ) + 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/security/test_auth_security.py b/tests/security/test_auth_security.py index d010e704..7cb83405 100644 --- a/tests/security/test_auth_security.py +++ b/tests/security/test_auth_security.py @@ -89,9 +89,11 @@ async def test_blocked_user_access(client): session_id=session_id, user_id=user_id, email="blocked@test.com", - expires_at=datetime.now(timezone.utc) + timedelta(hours=1), + idle_expires_at=datetime.now(timezone.utc) + timedelta(days=7), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(days=30), blocked=True, ttl=3600, + last_active=datetime.now(timezone.utc) ) payload = { @@ -127,9 +129,11 @@ async def test_rate_limiting(client): session_id=session_id, user_id=user_id, email="rate@test.com", - expires_at=datetime.now(timezone.utc) + timedelta(hours=1), + idle_expires_at=datetime.now(timezone.utc) + timedelta(days=7), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(days=30), blocked=False, ttl=3600, + last_active=datetime.now(timezone.utc) ) payload = { @@ -158,3 +162,88 @@ async def test_rate_limiting(client): await redis._client.delete("rate_limit:/user/photos:testclient") except Exception: pass + +async def test_fast_path_rejects_idle_expired_cached_session(client): + """Cached session past idle timeout but before absolute → 401.""" + async with engine.begin() as conn: + uq = user_queries.AsyncQuerier(conn) + user = await uq.create_user( + email=f"idle-expired-{uuid.uuid4()}@test.com", + hashed_password="hash" + ) + user_id = user.id + session_id = uuid.uuid4() + + redis = RedisClient.get_instance() + now = datetime.now(timezone.utc) + await SessionService.cache_session_for_auth( + redis=redis, + session_id=session_id, + user_id=user_id, + email="idle@test.com", + idle_expires_at=now - timedelta(hours=1), # expired + absolute_expires_at=now + timedelta(days=30), # still valid + blocked=False, + ttl=3600, + last_active=now - timedelta(hours=2), + ) + + payload = { + "session_id": str(session_id), + "exp": datetime.now(timezone.utc) + timedelta(hours=1), + } + token = jwt.encode(payload, settings.jwt_secret, algorithm="HS256") + + try: + response = await client.get( + "/user/photos", + headers={"Authorization": f"Bearer {token}"} + ) + assert response.status_code == 401 + assert "expired" in response.json()["detail"].lower() + finally: + async with engine.begin() as conn: + await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) + + +async def test_fast_path_rejects_absolute_expired_cached_session(client): + """Cached session past absolute timeout but before idle → 401.""" + async with engine.begin() as conn: + uq = user_queries.AsyncQuerier(conn) + user = await uq.create_user( + email=f"abs-expired-{uuid.uuid4()}@test.com", + hashed_password="hash" + ) + user_id = user.id + session_id = uuid.uuid4() + + redis = RedisClient.get_instance() + now = datetime.now(timezone.utc) + await SessionService.cache_session_for_auth( + redis=redis, + session_id=session_id, + user_id=user_id, + email="abs@test.com", + idle_expires_at=now + timedelta(days=7), # still valid + absolute_expires_at=now - timedelta(hours=1), # expired + blocked=False, + ttl=3600, + last_active=now - timedelta(hours=2), + ) + + payload = { + "session_id": str(session_id), + "exp": datetime.now(timezone.utc) + timedelta(hours=1), + } + token = jwt.encode(payload, settings.jwt_secret, algorithm="HS256") + + try: + response = await client.get( + "/user/photos", + headers={"Authorization": f"Bearer {token}"} + ) + assert response.status_code == 401 + assert "expired" in response.json()["detail"].lower() + finally: + async with engine.begin() as conn: + await conn.execute(text("DELETE FROM users WHERE id = :uid"), {"uid": user_id}) diff --git a/tests/unit/test_auth_email_otp.py b/tests/unit/test_auth_email_otp.py index 2dcfe269..7e6acd95 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 @@ -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 @@ -48,28 +54,22 @@ async def test_mobile_register_sends_otp( mock_redis: AsyncMock, mock_user_querier: AsyncMock, ) -> None: - # Arrange req = MobileRegisterRequest( email="test@example.com", password="Password1!", device_name="iPhone", device_type="iOS", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) - mock_user_querier.get_user_by_email.return_value = None # User does not exist - mock_redis.incr.return_value = 1 # Rate limit check passes + mock_user_querier.get_user_by_email.return_value = None + mock_redis.incr.return_value = 1 - # Act res = await auth_service.mobile_register(redis=mock_redis, req=req) - # Assert assert res.status == "pending_verification" assert res.email == "test@example.com" - - # Verify redis was called to save pending user and OTP assert mock_redis.set.call_count == 2 - # Verify NATS publish was called mock_publish.assert_called_once() args, _ = mock_publish.call_args assert args[0] == "email.send_otp" @@ -85,20 +85,19 @@ async def test_verify_mobile_register_success( mock_device_querier: AsyncMock, mock_session_querier: AsyncMock, ) -> None: - # Arrange device_id = uuid.uuid4() req = RegisterVerifyRequest( email="test@example.com", password="Password1!", device_name="iPhone", device_type="iOS", - device_id=device_id, + physical_device_id=device_id, otp="123456" ) mock_redis.get.side_effect = [ - "123456", # First call gets OTP - json.dumps({"hashed_password": "hashed_pass"}) # Second call gets pending user + "123456", + json.dumps({"hashed_password": "hashed_pass"}) ] mock_user = AsyncMock() @@ -107,64 +106,60 @@ 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 - mock_device_querier.get_device_by_id.return_value = None - mock_device_querier.get_device_by_id_any.return_value = None - # Act + # FIX 1: Mock the new lookup method to return None (device not found) + mock_device_querier.get_device_by_physical_id.return_value = None + # FIX 2: Mock create_device to return a truthy device + mock_device_querier.create_device.return_value = AsyncMock() + with patch("app.service.users.SessionService.cache_session_for_auth", new_callable=AsyncMock): res = await auth_service.verify_mobile_register(redis=mock_redis, req=req) - # Assert assert res.is_new_user is True assert res.user_id == mock_user.id - - # Verify user was created mock_user_querier.create_user.assert_called_once_with(email="test@example.com", hashed_password="hashed_pass") - - # Verify redis cleanup assert mock_redis.delete.call_count == 2 - @pytest.mark.asyncio async def test_mobile_register_resend_otp_success( auth_service: AuthService, mock_redis: AsyncMock, ) -> None: - # Arrange email = "test@example.com" mock_redis.get.return_value = '{"hashed_password": "fake"}' mock_redis.incr.return_value = 1 - # Act with patch("app.service.users.NatsClient.publish", new_callable=AsyncMock) as mock_publish: res = await auth_service.mobile_register_resend_otp(redis=mock_redis, email=email) - # Assert assert res.status == "pending_verification" assert res.message == "New OTP sent to email" assert res.email == email - mock_redis.get.assert_called_with(f"pending_user:{email}") mock_redis.set.assert_called_with(f"otp:{email}", ANY, expire=600) mock_publish.assert_called_once() - @pytest.mark.asyncio async def test_mobile_register_resend_otp_not_found( auth_service: AuthService, mock_redis: AsyncMock, ) -> None: from fastapi import HTTPException - # Arrange email = "test@example.com" mock_redis.incr.return_value = 1 - mock_redis.get.return_value = None # No pending user + mock_redis.get.return_value = None - # Act & Assert with pytest.raises(HTTPException) as exc_info: await auth_service.mobile_register_resend_otp(redis=mock_redis, email=email) diff --git a/tests/unit/test_auth_service.py b/tests/unit/test_auth_service.py index 91b306f2..622d074e 100644 --- a/tests/unit/test_auth_service.py +++ b/tests/unit/test_auth_service.py @@ -17,6 +17,9 @@ from app.core.securite import hash_password from app.schema.request.mobile.auth import MobileLoginRequest, MobileRegisterRequest +async def _empty_async_iter(): + return + yield # pragma: no cover # --------------------------------------------------------------------------- # Factories @@ -45,13 +48,15 @@ def _make_session( *, session_id: uuid.UUID | None = None, user_id: uuid.UUID | None = None, - expires_at: datetime | None = None, + idle_expires_at: datetime | None = None, + absolute_expires_at: datetime | None = None, ) -> MagicMock: s = MagicMock() s.id = session_id or uuid.uuid4() s.user_id = user_id or uuid.uuid4() s.device_id = uuid.uuid4() - s.expires_at = expires_at or datetime.now(timezone.utc) + timedelta(days=30) + s.idle_expires_at = idle_expires_at or datetime.now(timezone.utc) + timedelta(days=7) + s.absolute_expires_at = absolute_expires_at or datetime.now(timezone.utc) + timedelta(days=30) s.last_active = datetime.now(timezone.utc) return s @@ -73,7 +78,7 @@ def _make_login_request( return MobileLoginRequest( email=email, password=password, - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), # was: device_id device_name="iPhone 15", device_type="ios", ) @@ -87,12 +92,11 @@ def _make_register_request( return MobileRegisterRequest( email=email, password=password, - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), # was: device_id device_name="iPhone 15", device_type="ios", ) - # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- @@ -103,6 +107,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) @@ -116,6 +121,7 @@ def device_querier() -> AsyncMock: q = MagicMock(spec=device_queries.AsyncQuerier) q.get_device_by_id = AsyncMock(return_value=None) q.get_device_by_id_any = AsyncMock(return_value=None) + q.get_device_by_physical_id = AsyncMock(return_value=None) # new q.create_device = AsyncMock(return_value=_make_device()) q.activate_device = AsyncMock() return q @@ -125,7 +131,16 @@ 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() q.upsert_session = AsyncMock(return_value=_make_session()) q.get_session_by_id = AsyncMock() return q @@ -138,6 +153,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: @@ -156,12 +182,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, ) @@ -177,6 +205,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 @@ -194,6 +223,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 @@ -234,6 +264,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!") @@ -251,6 +282,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( @@ -280,7 +312,28 @@ async def test_blocked_user_raises_403( class TestSessionLimit: @pytest.mark.asyncio - async def test_exceeding_session_limit_raises_403( + async def test_at_cap_evicts_oldest_and_succeeds( + self, auth_service, user_querier, session_querier, redis, + ) -> 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() + + async def _evict(*, user_id, id, session_limit): + assert session_limit == AuthService.SESSION_LIMIT + yield evicted_id + + 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() + + @pytest.mark.asyncio + async def test_within_session_limit_succeeds( self, auth_service: AuthService, user_querier: AsyncMock, @@ -289,28 +342,48 @@ async def test_exceeding_session_limit_raises_403( ) -> None: user = _make_user() user_querier.get_user_by_email.return_value = user - # Return a count >= SESSION_LIMIT - session_querier.count_user_sessions.return_value = AuthService.SESSION_LIMIT + 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()) + assert result.access_token + session_querier.delete_session_by_id.assert_not_called() - with pytest.raises(HTTPException) as exc_info: - await auth_service.mobile_login(redis, _make_login_request()) - assert exc_info.value.status_code == 403 - assert "session limit" in exc_info.value.detail.lower() @pytest.mark.asyncio - async def test_within_session_limit_succeeds( + 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 - session_querier.count_user_sessions.return_value = AuthService.SESSION_LIMIT - 1 + user_querier.get_user_by_id_for_update.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 + + # =========================================================================== @@ -357,19 +430,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( @@ -377,19 +480,29 @@ 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) + idle_expires_at=datetime.now(timezone.utc) - timedelta(days=1), + absolute_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 @@ -398,9 +511,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 @@ -408,22 +522,104 @@ 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 + @pytest.mark.asyncio + async def test_refresh_rejects_idle_expired_session( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """Idle timeout expired but absolute still valid → refresh must reject.""" + from app.core.securite import create_raw_refresh_token + + now = datetime.now(timezone.utc) + session = _make_session( + idle_expires_at=now - timedelta(hours=1), # expired + absolute_expires_at=now + timedelta(days=30), # valid + ) + 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() + 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, raw_token) + assert exc_info.value.status_code == 401 + assert "expired" in exc_info.value.detail.lower() + + @pytest.mark.asyncio + async def test_refresh_rejects_absolute_expired_session( + self, + auth_service: AuthService, + user_querier: AsyncMock, + session_querier: AsyncMock, + refresh_token_querier: AsyncMock, + redis: AsyncMock, + ) -> None: + """Absolute timeout expired but idle still valid → refresh must reject.""" + from app.core.securite import create_raw_refresh_token + + now = datetime.now(timezone.utc) + session = _make_session( + idle_expires_at=now + timedelta(days=7), # valid + absolute_expires_at=now - timedelta(hours=1), # expired + ) + 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() + 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, raw_token) + assert exc_info.value.status_code == 401 + assert "expired" in exc_info.value.detail.lower() + # =========================================================================== # 6. find_closest_user @@ -459,3 +655,560 @@ 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()) + + 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 + +# =========================================================================== +# 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 + +class TestValidateSession: + @pytest.mark.asyncio + async def test_validate_session_false_when_idle_expired( + self, + auth_service: AuthService, + session_querier: AsyncMock, + user_querier: AsyncMock, + ) -> None: + now = datetime.now(timezone.utc) + session = _make_session( + idle_expires_at=now - timedelta(hours=1), + absolute_expires_at=now + timedelta(days=30), + ) + session_querier.get_session_by_id.return_value = session + user_querier.get_user_by_id.return_value = _make_user(user_id=session.user_id) + + result = await auth_service.validate_session(AsyncMock(), str(session.id)) + assert result is False + + @pytest.mark.asyncio + async def test_validate_session_false_when_absolute_expired( + self, + auth_service: AuthService, + session_querier: AsyncMock, + user_querier: AsyncMock, + ) -> None: + now = datetime.now(timezone.utc) + session = _make_session( + idle_expires_at=now + timedelta(days=7), + absolute_expires_at=now - timedelta(hours=1), + ) + session_querier.get_session_by_id.return_value = session + user_querier.get_user_by_id.return_value = _make_user(user_id=session.user_id) + + result = await auth_service.validate_session(AsyncMock(), str(session.id)) + assert result is False diff --git a/tests/unit/test_enroll_security.py b/tests/unit/test_enroll_security.py index 4f31584d..d944cb03 100644 --- a/tests/unit/test_enroll_security.py +++ b/tests/unit/test_enroll_security.py @@ -1,8 +1,5 @@ """ -Unit tests for the enrollment security helpers introduced in fix/enroll. - -These helpers only exist on the fix/enroll branch. On other branches, -all tests are automatically skipped via pytest.importorskip. +Unit tests for image validation helpers in app.core.image_validation. Run with: uv run pytest tests/unit/test_enroll_security.py -v """ @@ -15,27 +12,16 @@ from fastapi import UploadFile from fastapi.exceptions import HTTPException -# Guard: skip gracefully if the helpers don't exist on this branch -_router_mod = pytest.importorskip( - "app.router.mobile.enrollement", - reason="fix/enroll branch required for enrollment security helpers", +from app.core.constant import MAX_IMAGE_SIZE, MIN_IMAGE_DIM, MAX_IMAGE_DIM # noqa: E402 +from app.core.image_validation import ( + sanitise_filename, + validate_dimensions, + precheck_upload_headers, + read_limited, ) -_sanitise_filename = getattr(_router_mod, "_sanitise_filename", None) -_validate_dimensions = getattr(_router_mod, "_validate_dimensions", None) -_precheck_upload_headers = getattr(_router_mod, "_precheck_upload_headers", None) -_enrollment_lock_key = getattr(_router_mod, "_enrollment_lock_key", None) -read_limited = getattr(_router_mod, "read_limited", None) - -# If any helper is missing, skip the whole module -if any(fn is None for fn in [_sanitise_filename, _validate_dimensions, - _precheck_upload_headers, _enrollment_lock_key, read_limited]): - pytest.skip( - "Enrollment security helpers not available on this branch", - allow_module_level=True, - ) - -from app.core.constant import MAX_IMAGE_SIZE, MIN_IMAGE_DIM, MAX_IMAGE_DIM # noqa: E402 +# _enrollment_lock_key does not exist in the refactored module — skip those tests +_enrollment_lock_key = None # --------------------------------------------------------------------------- @@ -74,133 +60,134 @@ def _make_upload_file( # =========================================================================== -# 1. _sanitise_filename +# 1. sanitise_filename # =========================================================================== class TestSanitiseFilename: def test_normal_name_gets_uuid_prefix(self) -> None: - result = _sanitise_filename("portrait.jpg", "jpg") + result = sanitise_filename("portrait.jpg", "jpg") parts = result.split("_", 1) assert len(parts) == 2 uuid.UUID(parts[0]) # raises if not a valid UUID assert parts[1] == "portrait.jpg" def test_path_traversal_is_neutralised(self) -> None: - result = _sanitise_filename("../../../etc/passwd.jpg", "jpg") + result = sanitise_filename("../../../etc/passwd.jpg", "jpg") assert ".." not in result assert "/" not in result + assert "\\" not in result def test_null_bytes_are_replaced(self) -> None: - assert "\x00" not in _sanitise_filename("face\x00evil.jpg", "jpg") + assert "\x00" not in sanitise_filename("face\x00evil.jpg", "jpg") def test_control_characters_are_replaced(self) -> None: - assert "\x1f" not in _sanitise_filename("face\x1fmalicious.jpg", "jpg") + assert "\x1f" not in sanitise_filename("face\x1fmalicious.jpg", "jpg") def test_windows_reserved_chars_are_replaced(self) -> None: - for char in r'\/:*?"<>|': - assert char not in _sanitise_filename(f"face{char}name.jpg", "jpg"), \ + for char in r'\\/:*?"<>|': + assert char not in sanitise_filename(f"face{char}name.jpg", "jpg"), \ f"char {char!r} must be replaced" def test_none_filename_returns_uuid_only(self) -> None: - result = _sanitise_filename(None, "png") + result = sanitise_filename(None, "png") base, ext = result.rsplit(".", 1) uuid.UUID(base) assert ext == "png" def test_empty_filename_returns_uuid_only(self) -> None: - result = _sanitise_filename("", "jpg") + result = sanitise_filename("", "jpg") base, ext = result.rsplit(".", 1) uuid.UUID(base) assert ext == "jpg" def test_long_filename_is_truncated(self) -> None: - result = _sanitise_filename("a" * 200, "jpg") + result = sanitise_filename("a" * 200, "jpg") name_part = result.split("_", 1)[1] assert len(name_part) <= 128 def test_leading_dots_stripped(self) -> None: - result = _sanitise_filename("...hidden.jpg", "jpg") + result = sanitise_filename("...hidden.jpg", "jpg") name_part = result.split("_", 1)[1] assert not name_part.startswith(".") # =========================================================================== -# 2. _validate_dimensions +# 2. validate_dimensions # =========================================================================== class TestValidateDimensions: def test_valid_image_passes(self) -> None: - _validate_dimensions(_make_jpeg_bytes(200, 200)) # must not raise + validate_dimensions(_make_jpeg_bytes(200, 200)) # must not raise def test_too_small_width_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - _validate_dimensions(_make_jpeg_bytes(MIN_IMAGE_DIM - 1, 200)) + validate_dimensions(_make_jpeg_bytes(MIN_IMAGE_DIM - 1, 200)) assert exc_info.value.status_code == 400 assert "too small" in exc_info.value.detail.lower() def test_too_small_height_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - _validate_dimensions(_make_jpeg_bytes(200, MIN_IMAGE_DIM - 1)) + validate_dimensions(_make_jpeg_bytes(200, MIN_IMAGE_DIM - 1)) assert exc_info.value.status_code == 400 def test_too_large_width_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - _validate_dimensions(_make_jpeg_bytes(MAX_IMAGE_DIM + 1, 200)) + validate_dimensions(_make_jpeg_bytes(MAX_IMAGE_DIM + 1, 200)) assert exc_info.value.status_code == 400 assert "too large" in exc_info.value.detail.lower() def test_corrupt_bytes_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - _validate_dimensions(b"this-is-not-an-image") + validate_dimensions(b"this-is-not-an-image") assert exc_info.value.status_code == 400 def test_boundary_min_dimension_passes(self) -> None: - _validate_dimensions(_make_jpeg_bytes(MIN_IMAGE_DIM, MIN_IMAGE_DIM)) + validate_dimensions(_make_jpeg_bytes(MIN_IMAGE_DIM, MIN_IMAGE_DIM)) def test_boundary_max_dimension_passes(self) -> None: - _validate_dimensions(_make_jpeg_bytes(MAX_IMAGE_DIM, MAX_IMAGE_DIM)) + validate_dimensions(_make_jpeg_bytes(MAX_IMAGE_DIM, MAX_IMAGE_DIM)) # =========================================================================== -# 3. _precheck_upload_headers +# 3. precheck_upload_headers # =========================================================================== class TestPrecheckUploadHeaders: def test_valid_jpeg_header_passes(self) -> None: - _precheck_upload_headers(_make_upload_file(b"", content_type="image/jpeg")) + precheck_upload_headers(_make_upload_file(b"", content_type="image/jpeg")) def test_valid_png_header_passes(self) -> None: - _precheck_upload_headers(_make_upload_file(b"", content_type="image/png")) + precheck_upload_headers(_make_upload_file(b"", content_type="image/png")) def test_missing_content_type_raises_400(self) -> None: f = _make_upload_file(b"", content_type=None) with pytest.raises(HTTPException) as exc_info: - _precheck_upload_headers(f) + precheck_upload_headers(f) assert exc_info.value.status_code == 400 def test_unsupported_content_type_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - _precheck_upload_headers(_make_upload_file(b"", content_type="application/pdf")) + precheck_upload_headers(_make_upload_file(b"", content_type="application/pdf")) assert exc_info.value.status_code == 400 def test_content_type_with_charset_param_accepted(self) -> None: # "image/jpeg; charset=utf-8" should normalise to "image/jpeg" - _precheck_upload_headers( + precheck_upload_headers( _make_upload_file(b"", content_type="image/jpeg; charset=utf-8") ) def test_oversized_content_length_raises_400(self) -> None: with pytest.raises(HTTPException) as exc_info: - _precheck_upload_headers( + precheck_upload_headers( _make_upload_file(b"", content_type="image/jpeg", content_length=MAX_IMAGE_SIZE + 1) ) assert exc_info.value.status_code == 400 def test_valid_content_length_passes(self) -> None: - _precheck_upload_headers( + precheck_upload_headers( _make_upload_file(b"", content_type="image/jpeg", content_length=1024) ) @@ -208,7 +195,7 @@ def test_invalid_content_length_string_raises_400(self) -> None: f = _make_upload_file(b"", content_type="image/jpeg") f.headers = {"content-length": "not_a_number"} with pytest.raises(HTTPException) as exc_info: - _precheck_upload_headers(f) + precheck_upload_headers(f) assert exc_info.value.status_code == 400 @@ -267,23 +254,12 @@ async def mock_read(n: int = -1) -> bytes: # =========================================================================== -# 5. _enrollment_lock_key +# 5. _enrollment_lock_key — REMOVED +# This helper does not exist in app.core.image_validation. +# If it exists elsewhere, add a separate test file for it. # =========================================================================== -class TestEnrollmentLockKey: - def test_key_contains_user_id(self) -> None: - user_id = uuid.uuid4() - assert str(user_id) in _enrollment_lock_key(user_id) - - def test_different_users_have_different_keys(self) -> None: - assert _enrollment_lock_key(uuid.uuid4()) != _enrollment_lock_key(uuid.uuid4()) - - def test_key_format(self) -> None: - user_id = uuid.uuid4() - assert _enrollment_lock_key(user_id) == f"enroll:in_progress:{user_id}" - - # =========================================================================== # 6. Magic byte sniffing (unit-level verification) # =========================================================================== diff --git a/tests/unit/test_mobile_auth_email_logging.py b/tests/unit/test_mobile_auth_email_logging.py index 8180cfa7..d868f324 100644 --- a/tests/unit/test_mobile_auth_email_logging.py +++ b/tests/unit/test_mobile_auth_email_logging.py @@ -1,13 +1,14 @@ - # Test doubles intentionally implement only the AuthService methods exercised here. # They do not subclass the generated queriers, so mypy would otherwise flag each # constructor injection as an arg-type mismatch. # mypy: disable-error-code=arg-type import asyncio +from collections.abc import AsyncIterator import logging import uuid from datetime import datetime, timezone +from unittest.mock import MagicMock import pytest @@ -26,8 +27,10 @@ def __init__(self, email: str) -> None: class FakeDevice: - is_invalid_token = False - is_active = True + def __init__(self) -> None: + self.id = uuid.uuid4() + self.is_invalid_token = False + self.is_active = True class FakeSession: @@ -50,6 +53,11 @@ async def create_user(self, *, email: str, hashed_password: str) -> FakeUser: class FakeDeviceQuerier: + async def get_device_by_physical_id( + self, *, user_id: uuid.UUID, physical_device_id: uuid.UUID + ) -> FakeDevice | None: + return None + async def get_device_by_id_any(self, id: uuid.UUID) -> FakeDevice | None: return None @@ -64,8 +72,20 @@ 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 + ) -> FakeSession | None: + return None async def upsert_session( self, @@ -95,6 +115,7 @@ async def ttl(self, key: str) -> int: async def set(self, key: str, value: str, expire: int) -> None: return None + class FakeFaceEmbeddingService: pass @@ -113,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( @@ -120,7 +142,7 @@ def test_mobile_register_logs_without_plaintext_email( password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) async def _noop_cache_session_for_auth(**_: object) -> None: @@ -128,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)) @@ -137,4 +158,3 @@ async def _noop_cache_session_for_auth(**_: object) -> None: assert req.email not in caplog.text assert "user@example.com" not in caplog.text assert "mobile_register attempt" in caplog.text - diff --git a/tests/unit/test_mobile_auth_intent_validation.py b/tests/unit/test_mobile_auth_intent_validation.py index 574b5aa8..2311b7c9 100644 --- a/tests/unit/test_mobile_auth_intent_validation.py +++ b/tests/unit/test_mobile_auth_intent_validation.py @@ -1,3 +1,4 @@ +from collections.abc import AsyncIterator from typing import Any # Test doubles intentionally implement only the AuthService methods exercised here. @@ -6,16 +7,18 @@ # mypy: disable-error-code=arg-type import asyncio +import json import logging import uuid from datetime import datetime, timezone +from unittest.mock import AsyncMock import pytest from fastapi import HTTPException import app.service.users as users_module from app.core.securite import hash_password -from app.schema.request.mobile.auth import MobileLoginRequest, MobileRegisterRequest +from app.schema.request.mobile.auth import MobileLoginRequest, MobileRegisterRequest, RegisterVerifyRequest from app.service.session import SessionService from app.service.users import AuthService @@ -30,14 +33,22 @@ def __init__(self, email: str, exists: bool = True, password: str = "ValidPass@1 class FakeDevice: - is_invalid_token = False - is_active = True + def __init__(self, physical_device_id: uuid.UUID, user_id: uuid.UUID) -> None: + self.id = uuid.uuid4() + self.physical_device_id = physical_device_id + self.user_id = user_id + self.is_invalid_token = False + self.is_active = True class FakeSession: - def __init__(self) -> None: + def __init__(self, user_id: uuid.UUID, device_id: uuid.UUID) -> None: self.id = uuid.uuid4() + self.user_id = user_id + self.device_id = device_id self.expires_at = datetime.now(timezone.utc) + self.last_active = datetime.now(timezone.utc) + self.created_at = datetime.now(timezone.utc) class FakeUserQuerier: @@ -55,8 +66,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 @@ -65,47 +82,117 @@ async def create_user(self, *, email: str, hashed_password: str) -> FakeUser: class FakeDeviceQuerier: + """Stateful fake — tracks devices keyed by (user_id, physical_device_id), + matching the real UNIQUE(user_id, physical_device_id) constraint.""" + + def __init__(self) -> None: + self._devices: dict[tuple[uuid.UUID, uuid.UUID], FakeDevice] = {} + + async def get_device_by_physical_id( + self, *, user_id: uuid.UUID, physical_device_id: uuid.UUID + ) -> FakeDevice | None: + return self._devices.get((user_id, physical_device_id)) + async def get_device_by_id_any(self, id: uuid.UUID) -> FakeDevice | None: + # Kept only because the generated querier exposes it; application code + # no longer calls this for auth decisions. return None - async def get_device_by_id(self, id: uuid.UUID) -> FakeDevice | None: + async def get_device_by_id(self, id: uuid.UUID, user_id: uuid.UUID) -> FakeDevice | None: + for device in self._devices.values(): + if device.id == id and device.user_id == user_id: + return device return None - async def create_device(self, arg: object) -> FakeDevice: - return FakeDevice() + async def create_device(self, arg: Any) -> FakeDevice: + device = FakeDevice(physical_device_id=arg.physical_device_id, user_id=arg.user_id) + self._devices[(arg.user_id, arg.physical_device_id)] = device + return device async def activate_device(self, id: uuid.UUID, user_id: uuid.UUID) -> None: return None class FakeSessionQuerier: - def __init__(self, session: FakeSession) -> None: - self._session = session + """Stateful fake — tracks sessions keyed by (user_id, device_id), matching + the real UNIQUE(user_id, device_id) constraint and upsert-on-conflict + behavior.""" - async def count_user_sessions(self, user_id: uuid.UUID) -> int: - return 0 + def __init__(self) -> None: + self._sessions: dict[tuple[uuid.UUID, uuid.UUID], FakeSession] = {} + + async def get_session_by_device_for_user( + self, *, device_id: uuid.UUID, user_id: uuid.UUID + ) -> FakeSession | None: + return self._sessions.get((user_id, device_id)) async def get_session_by_id(self, id: uuid.UUID) -> FakeSession | None: - return self._session + for session in self._sessions.values(): + if session.id == id: + return session + return None + + async def list_sessions_by_user(self, user_id: uuid.UUID) -> AsyncIterator[FakeSession]: + for (u, _d), session in self._sessions.items(): + if u == user_id: + yield session + + async def delete_session_by_id(self, *, id: uuid.UUID, user_id: uuid.UUID) -> None: + key_to_remove = None + for key, session in self._sessions.items(): + if session.id == id and session.user_id == user_id: + key_to_remove = key + break + 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, *, user_id: uuid.UUID, device_id: uuid.UUID, - expires_at: datetime, + idle_expires_at: datetime, + absolute_expires_at: datetime, ) -> FakeSession: - self._session.expires_at = expires_at - return self._session - + key = (user_id, device_id) + existing = self._sessions.get(key) + if existing: + existing.idle_expires_at = idle_expires_at + existing.last_active = datetime.now(timezone.utc) + return existing + session = FakeSession(user_id=user_id, device_id=device_id) + session.idle_expires_at = idle_expires_at + session.absolute_expires_at = absolute_expires_at + self._sessions[key] = session + return session class FakeRedis: def __init__(self) -> None: - self._store: dict[str, int] = {} + self._store: dict[str, str] = {} async def incr(self, key: str) -> int: - self._store[key] = self._store.get(key, 0) + 1 - return self._store[key] + current = int(self._store.get(key, "0")) + 1 + self._store[key] = str(current) + return current async def expire(self, key: str, seconds: int) -> None: pass @@ -114,7 +201,10 @@ async def ttl(self, key: str) -> int: return -1 async def set(self, key: str, value: str, expire: int) -> None: - return None + self._store[key] = value + + async def get(self, key: str) -> str | None: + return self._store.get(key) async def delete(self, key: str) -> None: self._store.pop(key, None) @@ -124,15 +214,24 @@ class FakeFaceEmbeddingService: pass +def _patch_token_helpers(monkeypatch: pytest.MonkeyPatch) -> None: + async def _noop_cache_session_for_auth(**_: object) -> None: + return 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_raw_refresh_token", lambda: "refresh") + + def test_login_with_unknown_email_is_rejected() -> None: """Test that login with unknown email fails.""" user = FakeUser(email="user@example.com", exists=True) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -140,7 +239,7 @@ def test_login_with_unknown_email_is_rejected() -> None: password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) with pytest.raises(HTTPException) as exc_info: @@ -152,12 +251,12 @@ def test_login_with_unknown_email_is_rejected() -> None: def test_register_with_existing_email_is_rejected() -> None: """Test that registration with existing email fails.""" user = FakeUser(email="user@example.com", exists=True) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileRegisterRequest( @@ -165,7 +264,7 @@ def test_register_with_existing_email_is_rejected() -> None: password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) with pytest.raises(HTTPException) as exc_info: @@ -179,12 +278,12 @@ def test_login_with_correct_credentials_succeeds( ) -> None: """Test that login with correct credentials succeeds.""" user = FakeUser(email="user@example.com", exists=True) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -192,16 +291,10 @@ def test_login_with_correct_credentials_succeeds( password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) - async def _noop_cache_session_for_auth(**_: object) -> None: - return 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) + _patch_token_helpers(monkeypatch) result = asyncio.run(service.mobile_login(FakeRedis(), req)) assert result.access_token == "access" @@ -214,12 +307,12 @@ def test_register_with_new_email_succeeds( ) -> None: """Test that registration with new email succeeds.""" user = FakeUser(email="user@example.com", exists=False) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileRegisterRequest( @@ -227,16 +320,10 @@ def test_register_with_new_email_succeeds( password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) - async def _noop_cache_session_for_auth(**_: object) -> None: - return 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) + _patch_token_helpers(monkeypatch) result = asyncio.run(service.mobile_register(FakeRedis(), req)) assert result.status == "pending_verification" @@ -247,73 +334,52 @@ def test_register_then_login_same_device_succeeds( ) -> None: """Test full flow: register then login with same device.""" user = FakeUser(email="newuser@example.com", exists=False) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) - async def _noop_cache_session_for_auth(**_: object) -> None: - return 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) + _patch_token_helpers(monkeypatch) - device_id = uuid.uuid4() + physical_device_id = uuid.uuid4() password = "ValidPass@123" - # Register register_req = MobileRegisterRequest( email="newuser@example.com", password=password, device_name="TestDevice", device_type="android", - device_id=device_id, + physical_device_id=physical_device_id, ) fake_redis = FakeRedis() result1 = asyncio.run(service.mobile_register(fake_redis, register_req)) assert result1.status == "pending_verification" - # Try to register again (should just resend OTP, not fail with 409 because user is not yet fully enrolled) - # Actually, we expect 200 with pending_verification again if they try to register while pending, OR it might - # just succeed. Wait, the actual flow drops it in Redis and returns pending. So it won't raise 409 unless - # the user is IN THE DB. Since FakeUserQuerier won't have it, it won't raise 409. - # We will just verify OTP instead. - - # Now verify to fully create user - from app.schema.request.mobile.auth import RegisterVerifyRequest verify_req = RegisterVerifyRequest( email="newuser@example.com", password=password, otp="123456", device_name="TestDevice", device_type="android", - device_id=device_id, + physical_device_id=physical_device_id, ) - # Fake Redis returning the raw data and otp - import json fake_redis._store["otp:newuser@example.com"] = "123456" - fake_redis._store["pending_user:newuser@example.com"] = json.dumps({"hashed_password": hash_password(password)}) - - # We have to stub get() on FakeRedis since it doesn't support it by default - async def fake_get(key: str) -> str | None: - return fake_redis._store.get(key) - fake_redis.get = fake_get # type: ignore + fake_redis._store["pending_user:newuser@example.com"] = json.dumps( + {"hashed_password": hash_password(password)} + ) verify_result = asyncio.run(service.verify_mobile_register(fake_redis, verify_req)) assert verify_result.is_new_user is True - # Now login login_req = MobileLoginRequest( email="newuser@example.com", password=password, device_name="TestDevice", device_type="android", - device_id=device_id, + physical_device_id=physical_device_id, ) result2 = asyncio.run(service.mobile_login(FakeRedis(), login_req)) assert result2.is_new_user is False @@ -322,12 +388,12 @@ async def fake_get(key: str) -> str | None: def test_login_with_wrong_password_fails() -> None: """Test that login with wrong password fails.""" user = FakeUser(email="user@example.com", exists=True) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -335,7 +401,7 @@ def test_login_with_wrong_password_fails() -> None: password="wrongpassword", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) with pytest.raises(HTTPException) as exc_info: @@ -352,12 +418,12 @@ def test_login_logs_correctly( caplog.set_level(logging.INFO, logger="multAI") user = FakeUser(email="user@example.com", exists=True) - session = FakeSession() service = AuthService( user_querier=FakeUserQuerier(user), device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) req = MobileLoginRequest( @@ -365,16 +431,10 @@ def test_login_logs_correctly( password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) - async def _noop_cache_session_for_auth(**_: object) -> None: - return 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) + _patch_token_helpers(monkeypatch) asyncio.run(service.mobile_login(FakeRedis(), req)) @@ -393,7 +453,6 @@ class FakeOrigException(Exception): constraint_name = "idx_users_email" user = FakeUser(email="user@example.com", exists=False) - session = FakeSession() async def _raise_integrity_error(*args: Any, **kwargs: Any) -> Any: raise IntegrityError( @@ -403,35 +462,30 @@ async def _raise_integrity_error(*args: Any, **kwargs: Any) -> Any: ) user_querier = FakeUserQuerier(user) - # Stub create_user to raise IntegrityError user_querier.create_user = _raise_integrity_error # type: ignore service = AuthService( user_querier=user_querier, device_querier=FakeDeviceQuerier(), - session_querier=FakeSessionQuerier(session), + session_querier=FakeSessionQuerier(), face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), ) - # Stub create_user to raise IntegrityError during VERIFY, because mobile_register - # no longer calls create_user directly! - from app.schema.request.mobile.auth import RegisterVerifyRequest verify_req = RegisterVerifyRequest( email="newuser@example.com", password="ValidPass@123", otp="123456", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) fake_redis = FakeRedis() - import json fake_redis._store["otp:newuser@example.com"] = "123456" - fake_redis._store["pending_user:newuser@example.com"] = json.dumps({"hashed_password": hash_password("ValidPass@123")}) - async def fake_get(key: str) -> str | None: - return fake_redis._store.get(key) - fake_redis.get = fake_get # type: ignore + fake_redis._store["pending_user:newuser@example.com"] = json.dumps( + {"hashed_password": hash_password("ValidPass@123")} + ) with pytest.raises(HTTPException) as exc_info: asyncio.run(service.verify_mobile_register(fake_redis, verify_req)) @@ -439,3 +493,140 @@ async def fake_get(key: str) -> str | None: assert exc_info.value.status_code == 409 assert "already in use" in exc_info.value.detail.lower() + +# =========================================================================== +# Regression tests — Phase 2 bug fixes +# =========================================================================== + + +def test_session_device_id_matches_surrogate_pk_not_physical_id( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression test for the FK bug: user_sessions.device_id must be the + device row's surrogate id, never the client-supplied physical_device_id + directly. Pre-migration these happened to be equal by construction; + post-migration they are unrelated UUIDs.""" + user = FakeUser(email="user@example.com", exists=True) + device_querier = FakeDeviceQuerier() + session_querier = FakeSessionQuerier() + service = AuthService( + user_querier=FakeUserQuerier(user), + device_querier=device_querier, + session_querier=session_querier, + face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), + ) + + physical_id = uuid.uuid4() + req = MobileLoginRequest( + email="user@example.com", + password="ValidPass@123", + device_name="Pixel 8", + device_type="android", + physical_device_id=physical_id, + ) + + _patch_token_helpers(monkeypatch) + + asyncio.run(service.mobile_login(FakeRedis(), req)) + + assert len(device_querier._devices) == 1 + device = next(iter(device_querier._devices.values())) + assert len(session_querier._sessions) == 1 + session = next(iter(session_querier._sessions.values())) + + assert session.device_id == device.id + assert session.device_id != physical_id + + +def test_relogin_on_existing_device_succeeds_even_at_session_cap( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression test for the session-cap-vs-replace bug: a user at the + session cap must still be able to re-login on a device they already + have an active session on (replace), while a genuinely new device + should evict the oldest and succeed (Phase 4 behavior).""" + user = FakeUser(email="user@example.com", exists=True) + device_querier = FakeDeviceQuerier() + session_querier = FakeSessionQuerier() + service = AuthService( + user_querier=FakeUserQuerier(user), + device_querier=device_querier, + session_querier=session_querier, + face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), + ) + + _patch_token_helpers(monkeypatch) + + # Fill up the cap with sessions on distinct devices. + for i in range(AuthService.SESSION_LIMIT): + other_req = MobileLoginRequest( + email="user@example.com", + password="ValidPass@123", + device_name=f"Device {i}", + device_type="android", + physical_device_id=uuid.uuid4(), + ) + asyncio.run(service.mobile_login(FakeRedis(), other_req)) + + assert len(session_querier._sessions) == AuthService.SESSION_LIMIT + + # A genuinely NEW device at the cap should evict the oldest and SUCCEED. + result = asyncio.run(service.mobile_login(FakeRedis(), MobileLoginRequest( + email="user@example.com", + password="ValidPass@123", + device_name="New device", + device_type="android", + physical_device_id=uuid.uuid4(), + ))) + assert result.access_token == "access" + # Count stays at cap — one evicted, one added. + assert len(session_querier._sessions) == AuthService.SESSION_LIMIT + + # Re-logging in on an EXISTING device (replace) must still succeed. + existing_physical_id = next(iter(device_querier._devices.values())).physical_device_id + repeat_req = MobileLoginRequest( + email="user@example.com", + password="ValidPass@123", + device_name="Device 0", + device_type="android", + physical_device_id=existing_physical_id, + ) + result = asyncio.run(service.mobile_login(FakeRedis(), repeat_req)) + assert result.access_token == "access" + # Session count must NOT have grown — this was a replace, not an addition. + assert len(session_querier._sessions) == AuthService.SESSION_LIMIT + +def test_same_physical_device_id_reuses_device_row( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Two logins with the same (user, physical_device_id) must reuse the + same device row, never create a duplicate.""" + user = FakeUser(email="user@example.com", exists=True) + device_querier = FakeDeviceQuerier() + session_querier = FakeSessionQuerier() + service = AuthService( + user_querier=FakeUserQuerier(user), + device_querier=device_querier, + session_querier=session_querier, + face_embedding_service=FakeFaceEmbeddingService(), + refresh_token_querier=AsyncMock(), + ) + + _patch_token_helpers(monkeypatch) + + physical_id = uuid.uuid4() + + for _ in range(3): + req = MobileLoginRequest( + email="user@example.com", + password="ValidPass@123", + device_name="Pixel 8", + device_type="android", + physical_device_id=physical_id, + ) + asyncio.run(service.mobile_login(FakeRedis(), req)) + + assert len(device_querier._devices) == 1 + assert len(session_querier._sessions) == 1 diff --git a/tests/unit/test_mobile_auth_rate_limiting.py b/tests/unit/test_mobile_auth_rate_limiting.py index e5f227ed..e35d32c7 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 @@ -35,11 +36,15 @@ def __init__(self) -> None: self.hashed_password = hash_password("ValidPass@123") self.blocked = False - class FakeUserQuerier: + def __init__(self) -> None: + self._user = FakeUser() + async def get_user_by_email(self, email: str) -> FakeUser: - return FakeUser() + return self._user + async def get_user_by_id_for_update(self, id: uuid.UUID) -> FakeUser: + return self._user class FakeDeviceQuerier: pass @@ -61,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 @@ -81,7 +87,7 @@ async def _dummy_create_session(*args: object, **kwargs: object) -> Any: password="ValidPass@123", device_name="Pixel 8", device_type="android", - device_id=uuid.uuid4(), + physical_device_id=uuid.uuid4(), ) # Call mobile_login 5 times (which is the default max limit in settings) diff --git a/tests/unit/test_mobile_auth_request_validation.py b/tests/unit/test_mobile_auth_request_validation.py index 56d3a82e..39b287ed 100644 --- a/tests/unit/test_mobile_auth_request_validation.py +++ b/tests/unit/test_mobile_auth_request_validation.py @@ -88,7 +88,7 @@ def _valid_payload() -> dict[str, object]: "password": "ValidPass@123", "device_name": "Pixel 8", "device_type": "android", - "device_id": str(uuid.uuid4()), + "physical_device_id": str(uuid.uuid4()), }