mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
144
tests/unit/test_auth_login_adapter.py
Normal file
144
tests/unit/test_auth_login_adapter.py
Normal 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")
|
||||
84
tests/unit/test_auth_refresh_adapter.py
Normal file
84
tests/unit/test_auth_refresh_adapter.py
Normal 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
|
||||
43
tests/unit/test_config_auth_cookies.py
Normal file
43
tests/unit/test_config_auth_cookies.py
Normal 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()
|
||||
)
|
||||
@@ -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"}
|
||||
)
|
||||
|
||||
408
tests/unit/test_session_service.py
Normal file
408
tests/unit/test_session_service.py
Normal 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()
|
||||
54
tests/unit/test_user_me_password.py
Normal file
54
tests/unit/test_user_me_password.py
Normal 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()
|
||||
Reference in New Issue
Block a user