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("