mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +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:
@@ -96,6 +96,69 @@ def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) ->
|
||||
assert batch_getter.call_args.args[1] == ["user-1", "user-2"]
|
||||
|
||||
|
||||
def test_list_user_sessions_route_returns_sessions(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_admin_users_app(db, monkeypatch)
|
||||
sessions = [
|
||||
{
|
||||
"id": "session-1",
|
||||
"device_label": "Chrome / macOS",
|
||||
"device_type": "desktop",
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
"is_current": False,
|
||||
}
|
||||
]
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes._list_user_sessions_sync",
|
||||
lambda user_id: (
|
||||
sessions,
|
||||
{"action": "list_user_sessions", "target_user_id": user_id},
|
||||
),
|
||||
)
|
||||
|
||||
response = client.get("/api/admin/users/user-1/sessions")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == sessions
|
||||
|
||||
|
||||
def test_revoke_user_session_route_returns_message(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_admin_users_app(db, monkeypatch)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes._revoke_user_session_sync",
|
||||
lambda user_id, session_id, admin_user_id: (
|
||||
{"message": f"{user_id}:{session_id}:revoked"},
|
||||
{"action": "revoke_user_session", "target_user_id": user_id, "session_id": session_id},
|
||||
),
|
||||
)
|
||||
|
||||
response = client.delete("/api/admin/users/user-1/sessions/session-1")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"message": "user-1:session-1:revoked"}
|
||||
|
||||
|
||||
def test_revoke_all_user_sessions_route_returns_count(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_admin_users_app(db, monkeypatch)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes._revoke_all_user_sessions_sync",
|
||||
lambda user_id, admin_user_id: (
|
||||
{"message": "done", "revoked_count": 2},
|
||||
{"action": "revoke_all_user_sessions", "target_user_id": user_id, "revoked_count": 2},
|
||||
),
|
||||
)
|
||||
|
||||
response = client.delete("/api/admin/users/user-1/sessions")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"message": "done", "revoked_count": 2}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_user_adapter_preserves_empty_restriction_lists(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
@@ -103,14 +166,12 @@ async def test_create_user_adapter_preserves_empty_restriction_lists(
|
||||
db = MagicMock()
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _create_user_sync(request: Any, role: Any) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
captured["allowed_providers"] = request.allowed_providers
|
||||
captured["allowed_api_formats"] = request.allowed_api_formats
|
||||
captured["allowed_models"] = request.allowed_models
|
||||
def _fake_create_user_sync(request: Any, role: Any) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
captured["request"] = request
|
||||
captured["role"] = role
|
||||
return {"id": "user-3"}, {}
|
||||
return {"id": "user-3"}, {"action": "create_user", "target_user_id": "user-3"}
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes._create_user_sync", _create_user_sync)
|
||||
monkeypatch.setattr("src.api.admin.users.routes._create_user_sync", _fake_create_user_sync)
|
||||
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
@@ -131,6 +192,6 @@ async def test_create_user_adapter_preserves_empty_restriction_lists(
|
||||
result = await AdminCreateUserAdapter().handle(context)
|
||||
|
||||
assert result == {"id": "user-3"}
|
||||
assert captured["allowed_providers"] == []
|
||||
assert captured["allowed_api_formats"] == []
|
||||
assert captured["allowed_models"] == []
|
||||
assert captured["request"].allowed_providers == []
|
||||
assert captured["request"].allowed_api_formats == []
|
||||
assert captured["request"].allowed_models == []
|
||||
|
||||
220
tests/api/test_auth_routes.py
Normal file
220
tests/api/test_auth_routes.py
Normal file
@@ -0,0 +1,220 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.api.auth.routes import _logout_with_refresh_cookie_fallback
|
||||
from src.api.auth.routes import router as auth_router
|
||||
from src.config import config
|
||||
from src.database import get_db
|
||||
|
||||
|
||||
def _build_auth_app(
|
||||
db: MagicMock,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
pipeline_result: Any = None,
|
||||
pipeline_exception: Exception | None = None,
|
||||
) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(auth_router)
|
||||
app.dependency_overrides[get_db] = lambda: db
|
||||
|
||||
async def _fake_pipeline_run(
|
||||
*, adapter: Any, http_request: object, db: MagicMock, mode: object
|
||||
) -> Any:
|
||||
_ = adapter, http_request, db, mode
|
||||
if pipeline_exception is not None:
|
||||
raise pipeline_exception
|
||||
return pipeline_result
|
||||
|
||||
monkeypatch.setattr("src.api.auth.routes.pipeline.run", _fake_pipeline_run)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_login_route_sets_refresh_cookie_and_hides_token(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_auth_app(
|
||||
db,
|
||||
monkeypatch,
|
||||
pipeline_result={
|
||||
"access_token": "access-1",
|
||||
"token_type": "bearer",
|
||||
"expires_in": 86400,
|
||||
"user_id": "user-1",
|
||||
"username": "tester",
|
||||
"role": "user",
|
||||
"_refresh_token": "refresh-1",
|
||||
},
|
||||
)
|
||||
|
||||
response = client.post("/api/auth/login", json={"email": "user@example.com", "password": "pw"})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {
|
||||
"access_token": "access-1",
|
||||
"token_type": "bearer",
|
||||
"expires_in": 86400,
|
||||
"user_id": "user-1",
|
||||
"username": "tester",
|
||||
"role": "user",
|
||||
}
|
||||
set_cookie = response.headers.get("set-cookie", "")
|
||||
assert config.auth_refresh_cookie_name in set_cookie
|
||||
assert "refresh-1" in set_cookie
|
||||
assert "HttpOnly" in set_cookie
|
||||
|
||||
|
||||
def test_refresh_route_sets_cookie_and_hides_token(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_auth_app(
|
||||
db,
|
||||
monkeypatch,
|
||||
pipeline_result={
|
||||
"access_token": "access-2",
|
||||
"token_type": "bearer",
|
||||
"expires_in": 86400,
|
||||
"_refresh_token": "refresh-2",
|
||||
},
|
||||
)
|
||||
|
||||
response = client.post("/api/auth/refresh")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {
|
||||
"access_token": "access-2",
|
||||
"token_type": "bearer",
|
||||
"expires_in": 86400,
|
||||
}
|
||||
set_cookie = response.headers.get("set-cookie", "")
|
||||
assert config.auth_refresh_cookie_name in set_cookie
|
||||
assert "refresh-2" in set_cookie
|
||||
assert "HttpOnly" in set_cookie
|
||||
|
||||
|
||||
def test_refresh_route_clears_cookie_on_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_auth_app(
|
||||
db,
|
||||
monkeypatch,
|
||||
pipeline_exception=HTTPException(status_code=401, detail="登录会话已失效,请重新登录"),
|
||||
)
|
||||
|
||||
response = client.post("/api/auth/refresh")
|
||||
|
||||
assert response.status_code == 401
|
||||
assert response.json()["error"]["message"] == "登录会话已失效,请重新登录"
|
||||
set_cookie = response.headers.get("set-cookie", "")
|
||||
assert config.auth_refresh_cookie_name in set_cookie
|
||||
assert "Max-Age=0" in set_cookie
|
||||
|
||||
|
||||
def test_logout_route_clears_cookie_on_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_auth_app(
|
||||
db,
|
||||
monkeypatch,
|
||||
pipeline_exception=HTTPException(status_code=401, detail="缺少认证令牌"),
|
||||
)
|
||||
|
||||
response = client.post("/api/auth/logout")
|
||||
|
||||
assert response.status_code == 401
|
||||
assert response.json()["error"]["message"] == "缺少认证令牌"
|
||||
set_cookie = response.headers.get("set-cookie", "")
|
||||
assert config.auth_refresh_cookie_name in set_cookie
|
||||
assert "Max-Age=0" in set_cookie
|
||||
|
||||
|
||||
def test_logout_route_uses_refresh_cookie_fallback_on_auth_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
client = _build_auth_app(
|
||||
db,
|
||||
monkeypatch,
|
||||
pipeline_exception=HTTPException(status_code=401, detail="Token已过期"),
|
||||
)
|
||||
|
||||
async def _fake_fallback(_request: object, _db: MagicMock) -> dict[str, Any]:
|
||||
return {"message": "登出成功", "success": True}
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes._logout_with_refresh_cookie_fallback",
|
||||
_fake_fallback,
|
||||
)
|
||||
|
||||
response = client.post(
|
||||
"/api/auth/logout",
|
||||
cookies={config.auth_refresh_cookie_name: "refresh-1"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"message": "登出成功", "success": True}
|
||||
set_cookie = response.headers.get("set-cookie", "")
|
||||
assert config.auth_refresh_cookie_name in set_cookie
|
||||
assert "Max-Age=0" in set_cookie
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logout_with_refresh_cookie_fallback_revokes_session(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
user = SimpleNamespace(id="user-1", email="user@example.com")
|
||||
session = SimpleNamespace(id="session-1", is_revoked=False, is_expired=False)
|
||||
db.query.return_value.filter.return_value.first.return_value = user
|
||||
request = SimpleNamespace(
|
||||
cookies={config.auth_refresh_cookie_name: "refresh-1"},
|
||||
headers={},
|
||||
query_params={},
|
||||
state=SimpleNamespace(),
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes.AuthService.verify_token",
|
||||
AsyncMock(
|
||||
return_value={
|
||||
"user_id": "user-1",
|
||||
"session_id": "session-1",
|
||||
}
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes.SessionService.extract_client_device_id",
|
||||
lambda _request: "device-1",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes.SessionService.get_session_for_user",
|
||||
lambda *_args, **_kwargs: session,
|
||||
)
|
||||
assert_session_device_matches = MagicMock()
|
||||
revoke_session = MagicMock()
|
||||
log_event = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes.SessionService.assert_session_device_matches",
|
||||
assert_session_device_matches,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes.SessionService.revoke_session",
|
||||
revoke_session,
|
||||
)
|
||||
monkeypatch.setattr("src.api.auth.routes.AuditService.log_event", log_event)
|
||||
monkeypatch.setattr("src.api.auth.routes.get_client_ip", lambda _request: "127.0.0.1")
|
||||
monkeypatch.setattr(
|
||||
"src.api.auth.routes.get_user_agent",
|
||||
lambda _request: "pytest-agent",
|
||||
)
|
||||
|
||||
result = await _logout_with_refresh_cookie_fallback(request, db)
|
||||
|
||||
assert result == {"message": "登出成功", "success": True}
|
||||
assert_session_device_matches.assert_called_once_with(session, "device-1")
|
||||
revoke_session.assert_called_once()
|
||||
log_event.assert_called_once()
|
||||
db.commit.assert_called_once()
|
||||
assert request.state.tx_committed_by_route is True
|
||||
@@ -8,6 +8,7 @@ API Pipeline 测试
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
@@ -227,6 +228,35 @@ class TestPipelineAuditLogging:
|
||||
assert call_kwargs["status_code"] == 500
|
||||
assert call_kwargs["error_message"] == "Internal error"
|
||||
|
||||
def test_record_audit_event_commits_when_route_already_committed(
|
||||
self, pipeline: ApiRequestPipeline
|
||||
) -> None:
|
||||
mock_context = MagicMock()
|
||||
mock_context.db = MagicMock()
|
||||
mock_context.user = MagicMock()
|
||||
mock_context.user.id = "user-123"
|
||||
mock_context.api_key = None
|
||||
mock_context.request_id = "req-123"
|
||||
mock_context.client_ip = "127.0.0.1"
|
||||
mock_context.user_agent = "test-agent"
|
||||
mock_context.request = MagicMock()
|
||||
mock_context.request.method = "POST"
|
||||
mock_context.request.url.path = "/api/auth/refresh"
|
||||
mock_context.request.state = SimpleNamespace(tx_committed_by_route=True)
|
||||
mock_context.start_time = 1000.0
|
||||
|
||||
mock_adapter = MagicMock()
|
||||
mock_adapter.name = "test-adapter"
|
||||
mock_adapter.audit_log_enabled = True
|
||||
mock_adapter.audit_success_event = None
|
||||
mock_adapter.audit_failure_event = None
|
||||
|
||||
with patch.object(pipeline.audit_service, "log_event") as mock_log:
|
||||
pipeline._record_audit_event(mock_context, mock_adapter, success=True, status_code=200)
|
||||
|
||||
mock_log.assert_called_once()
|
||||
mock_context.db.commit.assert_called_once()
|
||||
|
||||
def test_record_audit_event_no_db(self, pipeline: ApiRequestPipeline) -> None:
|
||||
"""测试没有数据库会话时跳过审计"""
|
||||
mock_context = MagicMock()
|
||||
@@ -755,6 +785,8 @@ class TestPipelineAdminAuth:
|
||||
async def test_authenticate_admin_success(self, pipeline: ApiRequestPipeline) -> None:
|
||||
"""测试管理员认证成功"""
|
||||
created_at = datetime.now(timezone.utc)
|
||||
mock_session = MagicMock()
|
||||
mock_session.id = "session-123"
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.id = "admin-123"
|
||||
@@ -765,28 +797,46 @@ class TestPipelineAdminAuth:
|
||||
mock_user.created_at = created_at
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"authorization": "Bearer valid-token"}
|
||||
mock_request.headers = {
|
||||
"authorization": "Bearer valid-token",
|
||||
"X-Client-Device-Id": "device-admin-123",
|
||||
}
|
||||
mock_request.state = MagicMock()
|
||||
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||
|
||||
with patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"user_id": "admin-123", "created_at": created_at.isoformat()},
|
||||
with (
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"user_id": "admin-123",
|
||||
"created_at": created_at.isoformat(),
|
||||
"session_id": "session-123",
|
||||
},
|
||||
),
|
||||
patch(
|
||||
"src.api.base.pipeline.SessionService.get_active_session",
|
||||
return_value=mock_session,
|
||||
),
|
||||
patch("src.api.base.pipeline.SessionService.touch_session"),
|
||||
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||
):
|
||||
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
||||
|
||||
assert user == mock_user
|
||||
assert management_token is None
|
||||
assert mock_request.state.user_id == "admin-123"
|
||||
assert mock_request.state.user_session_id == "session-123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_admin_lowercase_bearer(self, pipeline: ApiRequestPipeline) -> None:
|
||||
"""测试 bearer (小写) 前缀也能正确解析"""
|
||||
created_at = datetime.now(timezone.utc)
|
||||
mock_session = MagicMock()
|
||||
mock_session.id = "session-123"
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.id = "admin-123"
|
||||
@@ -797,24 +847,79 @@ class TestPipelineAdminAuth:
|
||||
mock_user.created_at = created_at
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"authorization": "bearer valid-token"}
|
||||
mock_request.headers = {
|
||||
"authorization": "bearer valid-token",
|
||||
"X-Client-Device-Id": "device-admin-123",
|
||||
}
|
||||
mock_request.state = MagicMock()
|
||||
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||
|
||||
with patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"user_id": "admin-123", "created_at": created_at.isoformat()},
|
||||
) as mock_verify:
|
||||
with (
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"user_id": "admin-123",
|
||||
"created_at": created_at.isoformat(),
|
||||
"session_id": "session-123",
|
||||
},
|
||||
) as mock_verify,
|
||||
patch(
|
||||
"src.api.base.pipeline.SessionService.get_active_session",
|
||||
return_value=mock_session,
|
||||
),
|
||||
patch("src.api.base.pipeline.SessionService.touch_session"),
|
||||
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||
):
|
||||
user, management_token = await pipeline._authenticate_admin(mock_request, mock_db)
|
||||
|
||||
mock_verify.assert_awaited_once_with("valid-token", token_type="access")
|
||||
assert user == mock_user
|
||||
assert management_token is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_admin_rejects_legacy_token_without_session_id(
|
||||
self, pipeline: ApiRequestPipeline
|
||||
) -> None:
|
||||
created_at = datetime.now(timezone.utc)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {
|
||||
"authorization": "Bearer valid-token",
|
||||
"X-Client-Device-Id": "device-admin-legacy",
|
||||
}
|
||||
mock_request.state = MagicMock()
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.id = "admin-123"
|
||||
mock_user.is_active = True
|
||||
mock_user.is_deleted = False
|
||||
mock_user.role = UserRole.ADMIN
|
||||
mock_user.email = "admin@example.com"
|
||||
mock_user.created_at = created_at
|
||||
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"user_id": "admin-123", "created_at": created_at.isoformat()},
|
||||
),
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"token_identity_matches_user",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException, match="登录会话已失效,请重新登录"):
|
||||
await pipeline._authenticate_admin(mock_request, mock_db)
|
||||
|
||||
|
||||
class TestPipelineUserAuth:
|
||||
"""测试普通用户 JWT 认证"""
|
||||
@@ -827,6 +932,8 @@ class TestPipelineUserAuth:
|
||||
async def test_authenticate_user_lowercase_bearer(self, pipeline: ApiRequestPipeline) -> None:
|
||||
"""测试 bearer (小写) 前缀也能正确解析"""
|
||||
created_at = datetime.now(timezone.utc)
|
||||
mock_session = MagicMock()
|
||||
mock_session.id = "session-456"
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.id = "user-123"
|
||||
@@ -836,20 +943,119 @@ class TestPipelineUserAuth:
|
||||
mock_user.created_at = created_at
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"authorization": "bearer valid-token"}
|
||||
mock_request.headers = {
|
||||
"authorization": "bearer valid-token",
|
||||
"X-Client-Device-Id": "device-user-456",
|
||||
}
|
||||
mock_request.state = MagicMock()
|
||||
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||
|
||||
with patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"user_id": "user-123", "created_at": created_at.isoformat()},
|
||||
) as mock_verify:
|
||||
with (
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"user_id": "user-123",
|
||||
"created_at": created_at.isoformat(),
|
||||
"session_id": "session-456",
|
||||
},
|
||||
) as mock_verify,
|
||||
patch(
|
||||
"src.api.base.pipeline.SessionService.get_active_session",
|
||||
return_value=mock_session,
|
||||
),
|
||||
patch("src.api.base.pipeline.SessionService.touch_session"),
|
||||
patch("src.api.base.pipeline.SessionService.assert_session_device_matches"),
|
||||
):
|
||||
user, management_token = await pipeline._authenticate_user(mock_request, mock_db)
|
||||
|
||||
mock_verify.assert_awaited_once_with("valid-token", token_type="access")
|
||||
assert user == mock_user
|
||||
assert management_token is None
|
||||
assert mock_request.state.user_session_id == "session-456"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_user_rejects_legacy_token_without_session_id(
|
||||
self, pipeline: ApiRequestPipeline
|
||||
) -> None:
|
||||
"""历史 JWT 无 session_id 时应拒绝。"""
|
||||
created_at = datetime.now(timezone.utc)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {
|
||||
"authorization": "Bearer valid-token",
|
||||
"X-Client-Device-Id": "device-user-legacy",
|
||||
}
|
||||
mock_request.state = MagicMock()
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.id = "user-123"
|
||||
mock_user.is_active = True
|
||||
mock_user.is_deleted = False
|
||||
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={"user_id": "user-123", "created_at": created_at.isoformat()},
|
||||
),
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"token_identity_matches_user",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException, match="登录会话已失效,请重新登录"):
|
||||
await pipeline._authenticate_user(mock_request, mock_db)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_user_requires_device_id(self, pipeline: ApiRequestPipeline) -> None:
|
||||
created_at = datetime.now(timezone.utc)
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.headers = {"authorization": "Bearer valid-token"}
|
||||
mock_request.query_params = {}
|
||||
mock_request.state = MagicMock()
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.id = "user-123"
|
||||
mock_user.is_active = True
|
||||
mock_user.is_deleted = False
|
||||
mock_user.created_at = created_at
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.id = "session-456"
|
||||
|
||||
mock_db = MagicMock()
|
||||
mock_db.query.return_value.filter.return_value.first.return_value = mock_user
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"verify_token",
|
||||
new_callable=AsyncMock,
|
||||
return_value={
|
||||
"user_id": "user-123",
|
||||
"created_at": created_at.isoformat(),
|
||||
"session_id": "session-456",
|
||||
},
|
||||
),
|
||||
patch(
|
||||
"src.api.base.pipeline.SessionService.get_active_session",
|
||||
return_value=mock_session,
|
||||
),
|
||||
patch.object(
|
||||
pipeline.auth_service,
|
||||
"token_identity_matches_user",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
with pytest.raises(HTTPException, match="缺少或无效的设备标识"):
|
||||
await pipeline._authenticate_user(mock_request, mock_db)
|
||||
|
||||
@@ -268,10 +268,10 @@ class TestUserAuthentication:
|
||||
assert isinstance(result, AuthenticatedUserSnapshot)
|
||||
assert result.user_id == "user-123"
|
||||
assert result.username == "tester"
|
||||
thread_db.commit.assert_called_once()
|
||||
thread_db.commit.assert_not_called()
|
||||
thread_db.close.assert_called_once()
|
||||
route_db.commit.assert_not_called()
|
||||
invalidate_cache.assert_awaited_once_with("user-123", "test@example.com")
|
||||
invalidate_cache.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_user_for_pipeline_threadsafe_prefetches_balance(self) -> None:
|
||||
|
||||
129
tests/services/test_oauth_service.py
Normal file
129
tests/services/test_oauth_service.py
Normal file
@@ -0,0 +1,129 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import pytest
|
||||
|
||||
from src.core.enums import AuthSource
|
||||
from src.services.auth.oauth.service import OAuthService
|
||||
from src.services.auth.oauth.state import OAuthStateData
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_bind_authorize_url_includes_client_device_id(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
user = SimpleNamespace(id="user-1", auth_source=AuthSource.LOCAL)
|
||||
provider = SimpleNamespace(
|
||||
get_authorization_url=MagicMock(return_value="https://provider.example/authorize")
|
||||
)
|
||||
config = SimpleNamespace()
|
||||
create_state = AsyncMock(return_value="state-1")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._require_module_active",
|
||||
lambda _db: None,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._get_provider_impl",
|
||||
lambda _provider_type: provider,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._get_enabled_provider_config",
|
||||
lambda _db, _provider_type: config,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.get_redis_client",
|
||||
AsyncMock(return_value=object()),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.create_oauth_state",
|
||||
create_state,
|
||||
)
|
||||
|
||||
url = await OAuthService.build_bind_authorize_url(
|
||||
db,
|
||||
user,
|
||||
"github",
|
||||
client_device_id="device-1",
|
||||
)
|
||||
|
||||
assert url == "https://provider.example/authorize"
|
||||
assert create_state.await_args.kwargs["client_device_id"] == "device-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_callback_allows_bind_state_without_device_id(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
provider = SimpleNamespace(
|
||||
exchange_code=AsyncMock(return_value=SimpleNamespace(access_token="provider-access")),
|
||||
get_user_info=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
id="oauth-user",
|
||||
username="tester",
|
||||
email="user@example.com",
|
||||
email_verified=True,
|
||||
raw={},
|
||||
)
|
||||
),
|
||||
)
|
||||
config = SimpleNamespace(
|
||||
frontend_callback_url="https://app.example.com/auth/callback",
|
||||
display_name="GitHub",
|
||||
is_enabled=True,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._require_module_active",
|
||||
lambda _db: None,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._get_provider_impl",
|
||||
lambda _provider_type: provider,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._get_provider_config",
|
||||
lambda _db, _provider_type: config,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.get_redis_client",
|
||||
AsyncMock(return_value=object()),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.consume_oauth_state",
|
||||
AsyncMock(
|
||||
return_value=OAuthStateData(
|
||||
nonce="state-1",
|
||||
provider_type="github",
|
||||
action="bind",
|
||||
user_id="user-1",
|
||||
client_device_id=None,
|
||||
created_at=123,
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.services.auth.oauth.service.OAuthService._handle_bind",
|
||||
AsyncMock(return_value=SimpleNamespace()),
|
||||
)
|
||||
|
||||
result = await OAuthService.handle_callback(
|
||||
db=db,
|
||||
provider_type="github",
|
||||
state="state-1",
|
||||
code="code-1",
|
||||
error=None,
|
||||
error_description=None,
|
||||
client_ip=None,
|
||||
user_agent="pytest-agent",
|
||||
headers={},
|
||||
)
|
||||
|
||||
parsed = urlparse(result.redirect_url)
|
||||
assert result.refresh_token is None
|
||||
assert parse_qs(parsed.query)["oauth_bound"] == ["GitHub"]
|
||||
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