from __future__ import annotations from datetime import datetime, timedelta, timezone from types import SimpleNamespace import pytest from fastapi import HTTPException from sqlalchemy import create_engine from sqlalchemy.orm import Session, sessionmaker from src.core.enums import AuthSource, UserRole from src.models.database import Base, User, UserSession from src.services.auth.session_service import ( _DEVICE_ID_PATTERN, MAX_SESSIONS_PER_USER, TERMINAL_SESSION_RETENTION_DAYS, SessionClientContext, SessionService, ) def _make_session( *, token: str = "refresh-token-value", expires_delta: timedelta | None = None, revoked: bool = False, ) -> UserSession: expires_at = datetime.now(timezone.utc) + (expires_delta or timedelta(days=7)) session = UserSession( user_id="user-1", client_device_id="device-1", refresh_token_hash="", expires_at=expires_at, ) session.set_refresh_token(token) if revoked: session.revoked_at = datetime.now(timezone.utc) return session def _make_db_session() -> Session: engine = create_engine("sqlite:///:memory:") Base.metadata.create_all(engine, tables=[User.__table__, UserSession.__table__]) session_factory = sessionmaker(bind=engine) db = session_factory() db.info["test_engine"] = engine return db def _close_db_session(db: Session) -> None: engine = db.info.pop("test_engine", None) try: db.close() finally: if engine is not None: engine.dispose() def _make_user(db: Session, *, user_id: str = "user-1") -> User: user = User( id=user_id, email=f"{user_id}@example.com", email_verified=True, username=user_id, role=UserRole.USER, auth_source=AuthSource.LOCAL, is_active=True, is_deleted=False, ) db.add(user) db.commit() return user def _make_client_context(device_id: str) -> SessionClientContext: return SessionClientContext( client_device_id=device_id, device_label=f"Device {device_id}", device_type="desktop", browser_name="Chrome", browser_version="136.0", os_name="macOS", os_version="15", device_model=None, client_hints={}, ip_address="127.0.0.1", user_agent="pytest-agent", ) def test_build_client_context_prefers_client_hints() -> None: context = SessionService.build_client_context( client_device_id="device-1", client_ip="127.0.0.1", user_agent=( "Mozilla/5.0 (Linux; Android 15; Pixel 9) AppleWebKit/537.36 " "(KHTML, like Gecko) Chrome/134.0.0.0 Mobile Safari/537.36" ), headers={ "sec-ch-ua-platform": '"Android"', "sec-ch-ua-platform-version": '"15.0.0"', "sec-ch-ua-model": '"Pixel 9"', "sec-ch-ua-mobile": "?1", }, ) assert context.client_device_id == "device-1" assert context.device_type == "mobile" assert context.os_name == "Android" assert context.os_version == "15.0.0" assert context.device_model == "Pixel 9" assert context.device_label == "Pixel 9" # ── Refresh token hash roundtrip ── def test_user_session_refresh_token_hash_roundtrip() -> None: session = _make_session(token="refresh-token-value") is_valid, is_prev = session.verify_refresh_token("refresh-token-value") assert is_valid is True assert is_prev is False is_valid, is_prev = session.verify_refresh_token("different-token") assert is_valid is False assert is_prev is False def test_verify_refresh_token_grace_window() -> None: """轮换后短时间内旧 token 仍应通过验证(grace window)。""" session = _make_session(token="old-token") # 模拟轮换 session.set_refresh_token("new-token") # 新 token 正常通过 is_valid, is_prev = session.verify_refresh_token("new-token") assert is_valid is True assert is_prev is False # 旧 token 在宽限窗口内通过 is_valid, is_prev = session.verify_refresh_token("old-token") assert is_valid is True assert is_prev is True # 完全无关的 token 不通过 is_valid, is_prev = session.verify_refresh_token("random-token") assert is_valid is False assert is_prev is False def test_verify_refresh_token_grace_window_expired() -> None: """宽限窗口超时后,旧 token 不再通过。""" session = _make_session(token="old-token") session.set_refresh_token("new-token") # 手动把 rotated_at 设为 30 秒前,超过 REFRESH_GRACE_SECONDS session.rotated_at = datetime.now(timezone.utc) - timedelta(seconds=30) is_valid, is_prev = session.verify_refresh_token("old-token") assert is_valid is False assert is_prev is False # ── Session expiry ── def test_user_session_expired_property_handles_aware_datetime() -> None: session = _make_session(expires_delta=timedelta(minutes=-1)) assert session.is_expired is True def test_user_session_not_expired() -> None: session = _make_session(expires_delta=timedelta(days=7)) assert session.is_expired is False def test_user_session_revoked_property() -> None: session = _make_session(revoked=True) assert session.is_revoked is True session2 = _make_session(revoked=False) assert session2.is_revoked is False # ── Device ID validation ── def test_device_id_pattern_accepts_valid_ids() -> None: assert _DEVICE_ID_PATTERN.match("abc-123_DEF") assert _DEVICE_ID_PATTERN.match("a" * 128) assert _DEVICE_ID_PATTERN.match("simple-uuid-v4-like-id") def test_device_id_pattern_rejects_invalid_ids() -> None: assert not _DEVICE_ID_PATTERN.match("") assert not _DEVICE_ID_PATTERN.match("a" * 129) assert not _DEVICE_ID_PATTERN.match("has spaces") assert not _DEVICE_ID_PATTERN.match("has") assert not _DEVICE_ID_PATTERN.match("日本語") def test_normalize_device_id_strips_and_validates() -> None: assert SessionService._normalize_device_id(" valid-id ") == "valid-id" assert SessionService._normalize_device_id("a" * 200) is not None # truncated to 128 assert SessionService._normalize_device_id("