feat(auth): 重构认证系统,引入 session 会话管理

- 新增 user_sessions 数据库表及 Alembic 迁移
- 实现 SessionService 会话生命周期管理(创建/刷新/撤销/清理)
- 认证流程改用 refresh token cookie + access token 双令牌模式
- 前端实现自动静默刷新、跨标签页同步及设备指纹
- 用户设置页新增会话管理和密码修改功能
- 管理员用户管理新增强制登出和会话查看
- 密码策略增强,支持强度校验和泄露检测
- OAuth 登录流程适配新会话机制
- 新增完整的单元测试和 API 测试覆盖

Closes #232

Co-authored-by: LewisPen <LewisPen@nyadoo.com>
This commit is contained in:
fawney19
2026-03-17 16:34:09 +08:00
parent d480aa11f3
commit c4bb6b8161
58 changed files with 4654 additions and 508 deletions

View File

@@ -0,0 +1,144 @@
from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from src.api.auth.routes import AuthLoginAdapter
from src.core.enums import UserRole
from src.services.auth.service import AuthenticatedUserSnapshot
def _build_login_context(db: MagicMock) -> SimpleNamespace:
return SimpleNamespace(
db=db,
request=SimpleNamespace(state=SimpleNamespace(), headers={}),
ensure_json_body=lambda: {
"email": "user@example.com",
"password": "password123",
"auth_type": "local",
},
)
@pytest.mark.asyncio
async def test_auth_login_adapter_commits_login_success_metadata_with_session() -> None:
adapter = AuthLoginAdapter()
db = MagicMock()
db_user = SimpleNamespace(id="user-1", email="user@example.com", last_login_at=None)
db.query.return_value.filter.return_value.first.return_value = db_user
context = _build_login_context(db)
snapshot = AuthenticatedUserSnapshot(
user_id="user-1",
email="user@example.com",
username="tester",
role=UserRole.USER,
created_at=datetime.now(timezone.utc),
)
with (
patch(
"src.api.auth.routes.IPRateLimiter.check_limit",
new=AsyncMock(return_value=(True, 19, 0)),
),
patch(
"src.api.auth.routes.AuthService.authenticate_user_threadsafe",
new=AsyncMock(return_value=snapshot),
),
patch(
"src.api.auth.routes.get_client_ip",
return_value="127.0.0.1",
),
patch(
"src.api.auth.routes.get_user_agent",
return_value="pytest-agent",
),
patch(
"src.api.auth.routes.AuditService.log_login_attempt",
) as mock_log,
patch(
"src.api.auth.routes.UserCacheService.invalidate_user_cache",
new=AsyncMock(),
) as mock_invalidate,
):
def _issue_tokens(**_kwargs: object) -> tuple[str, str, str]:
assert mock_log.call_count == 0
return ("session-1", "access-1", "refresh-1")
with patch(
"src.api.auth.routes._issue_session_bound_tokens",
side_effect=_issue_tokens,
):
result = await adapter.handle(context)
assert result["access_token"] == "access-1"
assert result["_refresh_token"] == "refresh-1"
assert db_user.last_login_at is not None
db.commit.assert_called_once()
assert context.request.state.tx_committed_by_route is True
mock_log.assert_called_once_with(
db=db,
email="user@example.com",
success=True,
ip_address="127.0.0.1",
user_agent="pytest-agent",
user_id="user-1",
)
mock_invalidate.assert_awaited_once_with("user-1", "user@example.com")
@pytest.mark.asyncio
async def test_auth_login_adapter_skips_success_audit_when_session_creation_fails() -> None:
adapter = AuthLoginAdapter()
db = MagicMock()
db_user = SimpleNamespace(id="user-1", email="user@example.com", last_login_at=None)
db.query.return_value.filter.return_value.first.return_value = db_user
context = _build_login_context(db)
snapshot = AuthenticatedUserSnapshot(
user_id="user-1",
email="user@example.com",
username="tester",
role=UserRole.USER,
created_at=datetime.now(timezone.utc),
)
with (
patch(
"src.api.auth.routes.IPRateLimiter.check_limit",
new=AsyncMock(return_value=(True, 19, 0)),
),
patch(
"src.api.auth.routes.AuthService.authenticate_user_threadsafe",
new=AsyncMock(return_value=snapshot),
),
patch(
"src.api.auth.routes.get_client_ip",
return_value="127.0.0.1",
),
patch(
"src.api.auth.routes.get_user_agent",
return_value="pytest-agent",
),
patch(
"src.api.auth.routes.AuditService.log_login_attempt",
) as mock_log,
patch(
"src.api.auth.routes.UserCacheService.invalidate_user_cache",
new=AsyncMock(),
) as mock_invalidate,
patch(
"src.api.auth.routes._issue_session_bound_tokens",
side_effect=HTTPException(status_code=400, detail="缺少或无效的设备标识"),
),
):
with pytest.raises(HTTPException, match="缺少或无效的设备标识"):
await adapter.handle(context)
db.commit.assert_not_called()
mock_log.assert_not_called()
mock_invalidate.assert_not_awaited()
assert not hasattr(context.request.state, "tx_committed_by_route")

View File

@@ -0,0 +1,84 @@
from __future__ import annotations
from datetime import datetime, timezone
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from src.api.auth.routes import AuthRefreshAdapter
from src.config import config
@pytest.mark.asyncio
async def test_auth_refresh_adapter_skips_rotation_for_grace_window_token() -> None:
adapter = AuthRefreshAdapter()
created_at = datetime.now(timezone.utc)
user = SimpleNamespace(
id="user-1",
role=SimpleNamespace(value="user"),
created_at=created_at,
is_active=True,
is_deleted=False,
)
session = SimpleNamespace(id="session-1", client_device_id="device-1")
db = MagicMock()
db.query.return_value.filter.return_value.first.return_value = user
request = SimpleNamespace(
headers={"content-length": "0", "user-agent": "pytest"},
cookies={config.auth_refresh_cookie_name: "refresh-old"},
query_params={},
state=SimpleNamespace(),
)
context = SimpleNamespace(db=db, request=request)
with (
patch(
"src.api.auth.routes.AuthService.verify_token",
new=AsyncMock(
return_value={
"user_id": "user-1",
"session_id": "session-1",
"created_at": created_at.isoformat(),
}
),
),
patch(
"src.api.auth.routes.AuthService.token_identity_matches_user",
return_value=True,
),
patch(
"src.api.auth.routes.SessionService.extract_client_device_id",
return_value="device-1",
),
patch(
"src.api.auth.routes.SessionService.validate_refresh_session",
return_value=(session, True),
),
patch("src.api.auth.routes.SessionService.assert_session_device_matches"),
patch(
"src.api.auth.routes.AuthService.create_access_token",
return_value="access-new",
),
patch("src.api.auth.routes.AuthService.create_refresh_token") as mock_create_refresh,
patch("src.api.auth.routes.SessionService.rotate_refresh_token") as mock_rotate,
patch(
"src.api.auth.routes.get_client_ip",
return_value="127.0.0.1",
),
patch(
"src.api.auth.routes.get_user_agent",
return_value="pytest-agent",
),
):
result = await adapter.handle(context)
assert result == {
"access_token": "access-new",
"token_type": "bearer",
"expires_in": 86400,
}
mock_create_refresh.assert_not_called()
mock_rotate.assert_not_called()
db.commit.assert_called_once()
assert context.request.state.tx_committed_by_route is True

View File

@@ -0,0 +1,43 @@
from __future__ import annotations
import pytest
from src.config.settings import Config
def test_production_defaults_refresh_cookie_to_cross_site(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("ENVIRONMENT", "production")
monkeypatch.delenv("AUTH_REFRESH_COOKIE_SAMESITE", raising=False)
monkeypatch.delenv("AUTH_REFRESH_COOKIE_SECURE", raising=False)
cfg = Config()
assert cfg.auth_refresh_cookie_samesite == "none"
assert cfg.auth_refresh_cookie_secure is True
def test_validate_security_config_rejects_insecure_none_cookie(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("ENVIRONMENT", "production")
monkeypatch.setenv("AUTH_REFRESH_COOKIE_SAMESITE", "none")
monkeypatch.setenv("AUTH_REFRESH_COOKIE_SECURE", "false")
cfg = Config()
assert (
"AUTH_REFRESH_COOKIE_SECURE must be true when AUTH_REFRESH_COOKIE_SAMESITE=none."
in cfg.validate_security_config()
)
def test_validate_security_config_rejects_invalid_refresh_cookie_samesite(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setenv("AUTH_REFRESH_COOKIE_SAMESITE", "invalid-value")
cfg = Config()
assert "AUTH_REFRESH_COOKIE_SAMESITE must be one of: lax, strict, none." in (
cfg.validate_security_config()
)

View File

@@ -1,6 +1,8 @@
import pytest
from pydantic import ValidationError
from src.core.validators import PasswordPolicyLevel, PasswordValidator
from src.models.api import LoginRequest
class TestPasswordPolicy:
@@ -52,3 +54,28 @@ class TestPasswordPolicy:
assert valid is True
assert error is None
def test_rejects_password_longer_than_72_bytes(self) -> None:
valid, error = PasswordValidator.validate("a" * 80)
assert valid is False
assert error == "密码长度不能超过72字节"
def test_rejects_multibyte_password_longer_than_72_bytes(self) -> None:
valid, error = PasswordValidator.validate("" * 25)
assert valid is False
assert error == "密码长度不能超过72字节"
def test_login_request_preserves_leading_and_trailing_spaces(self) -> None:
request = LoginRequest.model_validate(
{"email": "tester", "password": " pass word ", "auth_type": "local"}
)
assert request.password == " pass word "
def test_login_request_rejects_password_longer_than_72_bytes(self) -> None:
with pytest.raises(ValidationError, match="密码长度不能超过72字节"):
LoginRequest.model_validate(
{"email": "tester", "password": "a" * 80, "auth_type": "local"}
)

View File

@@ -0,0 +1,408 @@
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)
return session_factory()
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<script>xss</script>")
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("<script>") is None
assert SessionService._normalize_device_id(" ") is None
def test_extract_client_device_id_requires_valid_value() -> None:
request = SimpleNamespace(headers={}, query_params={})
with pytest.raises(HTTPException, match="缺少或无效的设备标识"):
SessionService.extract_client_device_id(request) # type: ignore[arg-type]
def test_assert_session_device_matches_rejects_mismatch() -> None:
session = _make_session()
session.client_device_id = "device-expected"
with pytest.raises(HTTPException, match="设备标识与登录会话不匹配"):
SessionService.assert_session_device_matches(session, "device-actual")
# ── set_refresh_token preserves previous hash ──
def test_set_refresh_token_stores_previous_hash() -> None:
session = _make_session(token="token-1")
original_hash = session.refresh_token_hash
assert session.prev_refresh_token_hash is None or session.prev_refresh_token_hash == ""
session.set_refresh_token("token-2")
assert session.prev_refresh_token_hash == original_hash
assert session.rotated_at is not None
assert session.refresh_token_hash != original_hash
def test_get_session_for_user_can_lock_for_update() -> None:
expected = object()
class DummyQuery:
def __init__(self) -> None:
self.locked = False
def filter(self, *_args: object, **_kwargs: object) -> "DummyQuery":
return self
def with_for_update(self) -> "DummyQuery":
self.locked = True
return self
def first(self) -> object:
return expected
query = DummyQuery()
db = SimpleNamespace(query=lambda _model: query)
result = SessionService.get_session_for_user(
db, # type: ignore[arg-type]
user_id="user-1",
session_id="session-1",
lock_for_update=True,
)
assert result is expected
assert query.locked is True
def test_create_session_session_limit_ignores_expired_sessions() -> None:
db = _make_db_session()
try:
user = _make_user(db)
now = datetime.now(timezone.utc)
for idx in range(MAX_SESSIONS_PER_USER - 1):
session = UserSession(
id=f"active-{idx}",
user_id=user.id,
client_device_id=f"active-device-{idx}",
refresh_token_hash="",
expires_at=now + timedelta(days=7),
last_seen_at=now + timedelta(minutes=idx),
)
session.set_refresh_token(f"active-token-{idx}")
db.add(session)
expired = UserSession(
id="expired-1",
user_id=user.id,
client_device_id="expired-device",
refresh_token_hash="",
expires_at=now - timedelta(minutes=5),
last_seen_at=now - timedelta(days=1),
)
expired.set_refresh_token("expired-token")
db.add(expired)
db.commit()
created = SessionService.create_session(
db,
user=user,
session_id="new-session",
refresh_token="new-refresh-token",
expires_at=now + timedelta(days=7),
client=_make_client_context("new-device"),
)
db.commit()
assert created.id == "new-session"
active_sessions = SessionService.list_user_sessions(db, user_id=user.id)
assert len(active_sessions) == MAX_SESSIONS_PER_USER
assert all(session.revoke_reason != "session_limit_exceeded" for session in active_sessions)
finally:
db.close()
def test_revoke_all_user_sessions_skips_expired_sessions() -> None:
db = _make_db_session()
try:
user = _make_user(db)
now = datetime.now(timezone.utc)
active = UserSession(
id="active-1",
user_id=user.id,
client_device_id="active-device",
refresh_token_hash="",
expires_at=now + timedelta(days=7),
last_seen_at=now,
)
active.set_refresh_token("active-token")
expired = UserSession(
id="expired-1",
user_id=user.id,
client_device_id="expired-device",
refresh_token_hash="",
expires_at=now - timedelta(minutes=1),
last_seen_at=now - timedelta(days=1),
)
expired.set_refresh_token("expired-token")
db.add_all([active, expired])
db.commit()
revoked_count = SessionService.revoke_all_user_sessions(
db,
user_id=user.id,
reason="security_review",
)
db.flush()
db.refresh(active)
db.refresh(expired)
assert revoked_count == 1
assert active.revoked_at is not None
assert expired.revoked_at is None
finally:
db.close()
def test_list_user_sessions_prunes_old_terminal_sessions() -> None:
db = _make_db_session()
try:
user = _make_user(db)
now = datetime.now(timezone.utc)
active = UserSession(
id="active-1",
user_id=user.id,
client_device_id="active-device",
refresh_token_hash="",
expires_at=now + timedelta(days=7),
last_seen_at=now,
)
active.set_refresh_token("active-token")
old_expired = UserSession(
id="expired-old",
user_id=user.id,
client_device_id="expired-old-device",
refresh_token_hash="",
expires_at=now - timedelta(days=TERMINAL_SESSION_RETENTION_DAYS + 5),
last_seen_at=now - timedelta(days=40),
)
old_expired.set_refresh_token("expired-old-token")
old_revoked = UserSession(
id="revoked-old",
user_id=user.id,
client_device_id="revoked-old-device",
refresh_token_hash="",
expires_at=now + timedelta(days=7),
last_seen_at=now - timedelta(days=35),
)
old_revoked.set_refresh_token("revoked-old-token")
old_revoked.revoked_at = now - timedelta(days=TERMINAL_SESSION_RETENTION_DAYS + 1)
old_revoked.revoke_reason = "manual_revoke"
recent_expired = UserSession(
id="expired-recent",
user_id=user.id,
client_device_id="expired-recent-device",
refresh_token_hash="",
expires_at=now - timedelta(days=1),
last_seen_at=now - timedelta(days=2),
)
recent_expired.set_refresh_token("expired-recent-token")
db.add_all([active, old_expired, old_revoked, recent_expired])
db.commit()
sessions = SessionService.list_user_sessions(db, user_id=user.id)
remaining_ids = {session.id for session in db.query(UserSession).all()}
assert [session.id for session in sessions] == ["active-1"]
assert "expired-old" not in remaining_ids
assert "revoked-old" not in remaining_ids
assert "expired-recent" in remaining_ids
finally:
db.close()

View File

@@ -0,0 +1,54 @@
from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from src.api.user_me.routes import _change_password_sync
from src.core.exceptions import InvalidRequestException
from src.core.validators import PasswordPolicyLevel
from src.models.database import User
def test_verify_password_returns_false_for_password_over_72_bytes() -> None:
user = User(email=None, email_verified=False, username="tester")
user.set_password("abc12345")
assert user.verify_password("a" * 80) is False
def test_change_password_rejects_same_as_current_password(
monkeypatch: pytest.MonkeyPatch,
) -> None:
user = SimpleNamespace(
id="user-1",
email="user@example.com",
password_hash="hashed",
auth_source=SimpleNamespace(value="local"),
verify_password=MagicMock(side_effect=lambda password: password == "Abcd1234!"),
set_password=MagicMock(),
updated_at=None,
)
db = MagicMock()
db.query.return_value.filter.return_value.first.return_value = user
@contextmanager
def _fake_get_db_context() -> MagicMock:
yield db
monkeypatch.setattr("src.api.user_me.routes.get_db_context", _fake_get_db_context)
monkeypatch.setattr(
"src.api.user_me.routes.SystemConfigService.get_password_policy_level",
lambda _db: PasswordPolicyLevel.STRONG.value,
)
with pytest.raises(InvalidRequestException, match="新密码不能与当前密码相同"):
_change_password_sync(
"user-1",
SimpleNamespace(old_password="Abcd1234!", new_password="Abcd1234!"),
)
user.set_password.assert_not_called()