test(p3-1): 第二十三波 - auth登录/刷新/登出use case 38个单测 #729
Executable
+589
@@ -0,0 +1,589 @@
|
||||
"""Auth login use cases unit tests.
|
||||
|
||||
Covers LoginUseCase, RefreshTokenUseCase, LogoutUseCase, and helper functions.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.auth.login_use_case import (
|
||||
LEGACY_SHA256_HEX_LENGTH,
|
||||
LoginRequest,
|
||||
LoginResponse,
|
||||
LoginUseCase,
|
||||
LogoutRequest,
|
||||
LogoutUseCase,
|
||||
RefreshTokenRequest,
|
||||
RefreshTokenUseCase,
|
||||
_is_legacy_sha256_hash,
|
||||
_legacy_sha256,
|
||||
)
|
||||
from packages.application.auth.password_hasher import password_hasher
|
||||
|
||||
# ── Test helpers ─────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeUser:
|
||||
id: str = "user-123"
|
||||
email: str = "test@example.com"
|
||||
display_name: str = "Test User"
|
||||
username: str = "testuser"
|
||||
password_hash: str = ""
|
||||
email_verified: bool = True
|
||||
last_login_at: datetime | None = None
|
||||
last_login_ip: str | None = None
|
||||
wechat_openid: str | None = None
|
||||
wechat_unionid: str | None = None
|
||||
|
||||
|
||||
class FakeUserRepository:
|
||||
def __init__(self, user: FakeUser | None = None):
|
||||
self._user = user
|
||||
self.saved_user: FakeUser | None = None
|
||||
|
||||
def find_by_email(self, email: str) -> FakeUser | None:
|
||||
if self._user and self._user.email == email:
|
||||
return self._user
|
||||
return None
|
||||
|
||||
def get(self, user_id: str) -> FakeUser | None:
|
||||
if self._user and self._user.id == user_id:
|
||||
return self._user
|
||||
return None
|
||||
|
||||
def save(self, user: FakeUser) -> FakeUser:
|
||||
self.saved_user = user
|
||||
self._user = user
|
||||
return user
|
||||
|
||||
|
||||
class FakeSessionStore:
|
||||
def __init__(self):
|
||||
self._sessions: dict[str, dict] = {}
|
||||
self._refresh_index: dict[str, str] = {} # refresh_token -> session_id
|
||||
self.saved_sessions: list[dict] = []
|
||||
self.deleted_sessions: list[str] = []
|
||||
self.delete_all_called_for: str | None = None
|
||||
self.delete_all_return_value = 0
|
||||
|
||||
def save_session(self, **kwargs) -> bool:
|
||||
session_id = kwargs.get("session_id", "")
|
||||
self._sessions[session_id] = kwargs
|
||||
if kwargs.get("refresh_token"):
|
||||
self._refresh_index[kwargs["refresh_token"]] = session_id
|
||||
self.saved_sessions.append(kwargs)
|
||||
return True
|
||||
|
||||
def get_session_by_refresh_token(self, refresh_token: str) -> dict | None:
|
||||
session_id = self._refresh_index.get(refresh_token)
|
||||
if not session_id:
|
||||
return None
|
||||
return self._sessions.get(session_id)
|
||||
|
||||
def get_refresh_token(self, session_id: str) -> str | None:
|
||||
session = self._sessions.get(session_id)
|
||||
if not session:
|
||||
return None
|
||||
return session.get("refresh_token")
|
||||
|
||||
def delete_session(self, session_id: str) -> bool:
|
||||
self.deleted_sessions.append(session_id)
|
||||
if session_id in self._sessions:
|
||||
session = self._sessions.pop(session_id)
|
||||
rt = session.get("refresh_token")
|
||||
if rt and rt in self._refresh_index:
|
||||
del self._refresh_index[rt]
|
||||
return True
|
||||
return False
|
||||
|
||||
def delete_all_user_sessions(self, user_id: str) -> int:
|
||||
self.delete_all_called_for = user_id
|
||||
count = self.delete_all_return_value
|
||||
# actually clean up
|
||||
to_delete = [sid for sid, s in self._sessions.items() if s.get("user_id") == user_id]
|
||||
for sid in to_delete:
|
||||
self.delete_session(sid)
|
||||
return count or len(to_delete)
|
||||
|
||||
|
||||
# ── Helper function tests ───────────────────────────────
|
||||
|
||||
|
||||
class TestIsLegacySha256Hash:
|
||||
def test_valid_sha256_hex(self):
|
||||
h = hashlib.sha256(b"password").hexdigest()
|
||||
assert _is_legacy_sha256_hash(h) is True
|
||||
|
||||
def test_bcrypt_hash_not_legacy(self):
|
||||
h = password_hasher.hash_password("password")
|
||||
assert _is_legacy_sha256_hash(h) is False
|
||||
|
||||
def test_empty_string(self):
|
||||
assert _is_legacy_sha256_hash("") is False
|
||||
|
||||
def test_short_string(self):
|
||||
assert _is_legacy_sha256_hash("abc123") is False
|
||||
|
||||
def test_64_chars_non_hex(self):
|
||||
s = "g" * 64 # 'g' is not hex
|
||||
assert _is_legacy_sha256_hash(s) is False
|
||||
|
||||
def test_exact_64_hex_chars(self):
|
||||
s = "a" * 64
|
||||
assert _is_legacy_sha256_hash(s) is True
|
||||
|
||||
def test_mixed_case_hex(self):
|
||||
s = "AbCdEf01" * 8 # 64 chars mixed case hex
|
||||
assert len(s) == 64
|
||||
assert _is_legacy_sha256_hash(s) is True
|
||||
|
||||
|
||||
class TestLegacySha256:
|
||||
def test_matches_hashlib(self):
|
||||
password = "mypassword"
|
||||
expected = hashlib.sha256(password.encode()).hexdigest()
|
||||
assert _legacy_sha256(password) == expected
|
||||
|
||||
def test_empty_password(self):
|
||||
assert _legacy_sha256("") == hashlib.sha256(b"").hexdigest()
|
||||
|
||||
def test_unicode_password(self):
|
||||
result = _legacy_sha256("密码测试")
|
||||
assert len(result) == LEGACY_SHA256_HEX_LENGTH
|
||||
assert all(c in "0123456789abcdef" for c in result)
|
||||
|
||||
|
||||
# ── LoginRequest tests ──────────────────────────────────
|
||||
|
||||
|
||||
class TestLoginRequest:
|
||||
def test_email_stripped_and_lowercased(self):
|
||||
req = LoginRequest(email=" Test@Example.COM ", password="pass")
|
||||
assert req.email == "test@example.com"
|
||||
|
||||
def test_default_device_info(self):
|
||||
req = LoginRequest(email="a@b.com", password="pass")
|
||||
assert req.device_info == "Unknown"
|
||||
|
||||
def test_default_ip_address(self):
|
||||
req = LoginRequest(email="a@b.com", password="pass")
|
||||
assert req.ip_address == "unknown"
|
||||
|
||||
def test_custom_device_and_ip(self):
|
||||
req = LoginRequest(email="a@b.com", password="pass", device_info="iPhone", ip_address="1.2.3.4")
|
||||
assert req.device_info == "iPhone"
|
||||
assert req.ip_address == "1.2.3.4"
|
||||
|
||||
|
||||
# ── LoginUseCase tests ──────────────────────────────────
|
||||
|
||||
|
||||
class TestLoginUseCase:
|
||||
def _make_user_with_password(self, password: str = "password123") -> FakeUser:
|
||||
return FakeUser(password_hash=password_hasher.hash_password(password))
|
||||
|
||||
def test_successful_login(self):
|
||||
user = self._make_user_with_password("mypassword")
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="test@example.com", password="mypassword")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert isinstance(response, LoginResponse)
|
||||
assert response.user_id == "user-123"
|
||||
assert response.email == "test@example.com"
|
||||
assert response.username == "testuser"
|
||||
assert response.display_name == "Test User"
|
||||
assert response.access_token
|
||||
assert response.refresh_token
|
||||
assert response.expires_in > 0
|
||||
|
||||
# session was saved
|
||||
assert len(store.saved_sessions) == 1
|
||||
saved = store.saved_sessions[0]
|
||||
assert saved["user_id"] == "user-123"
|
||||
assert saved["device_info"] == "Unknown"
|
||||
assert saved["ip_address"] == "unknown"
|
||||
assert saved["expires_in_seconds"] == 30 * 24 * 3600
|
||||
|
||||
# last_login updated
|
||||
assert repo.saved_user is not None
|
||||
assert repo.saved_user.last_login_at is not None
|
||||
assert repo.saved_user.last_login_ip == "unknown"
|
||||
|
||||
def test_successful_login_with_device_and_ip(self):
|
||||
user = self._make_user_with_password("pass")
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(
|
||||
email="test@example.com",
|
||||
password="pass",
|
||||
device_info="Chrome/Win10",
|
||||
ip_address="192.168.1.1",
|
||||
)
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
saved = store.saved_sessions[0]
|
||||
assert saved["device_info"] == "Chrome/Win10"
|
||||
assert saved["ip_address"] == "192.168.1.1"
|
||||
assert repo.saved_user.last_login_ip == "192.168.1.1"
|
||||
|
||||
def test_empty_email_returns_error(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="", password="pass")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Email is required" in error
|
||||
|
||||
def test_empty_password_returns_error(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="a@b.com", password="")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Password is required" in error
|
||||
|
||||
def test_user_not_found_returns_error(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="nobody@example.com", password="pass")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Invalid email or password" in error
|
||||
|
||||
def test_wrong_password_returns_error(self):
|
||||
user = self._make_user_with_password("correctpassword")
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="test@example.com", password="wrongpassword")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Invalid email or password" in error
|
||||
# no session created
|
||||
assert len(store.saved_sessions) == 0
|
||||
|
||||
def test_legacy_sha256_hash_login_success_and_upgrade(self):
|
||||
password = "oldpassword"
|
||||
legacy_hash = _legacy_sha256(password)
|
||||
assert _is_legacy_sha256_hash(legacy_hash)
|
||||
|
||||
user = FakeUser(password_hash=legacy_hash)
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="test@example.com", password=password)
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
|
||||
# password should have been upgraded to bcrypt
|
||||
assert repo.saved_user is not None
|
||||
assert not _is_legacy_sha256_hash(repo.saved_user.password_hash)
|
||||
assert repo.saved_user.password_hash.startswith("$2b$")
|
||||
|
||||
# new hash should verify correctly
|
||||
assert password_hasher.verify_password(password, repo.saved_user.password_hash)
|
||||
|
||||
def test_legacy_sha256_hash_wrong_password(self):
|
||||
password = "rightpassword"
|
||||
legacy_hash = _legacy_sha256(password)
|
||||
user = FakeUser(password_hash=legacy_hash)
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="test@example.com", password="wrongpassword")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Invalid email or password" in error
|
||||
# password hash unchanged
|
||||
assert repo.saved_user is None
|
||||
|
||||
def test_exception_handling(self):
|
||||
repo = MagicMock()
|
||||
repo.find_by_email.side_effect = RuntimeError("DB connection failed")
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="a@b.com", password="pass")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert error is not None
|
||||
assert "Login failed" in error
|
||||
|
||||
def test_access_token_contains_correct_claims(self):
|
||||
import jwt as pyjwt
|
||||
|
||||
user = self._make_user_with_password("pass")
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
use_case = LoginUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = LoginRequest(email="test@example.com", password="pass")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
|
||||
payload = pyjwt.decode(
|
||||
response.access_token,
|
||||
use_case.jwt_secret_key,
|
||||
algorithms=["HS256"],
|
||||
)
|
||||
assert payload["sub"] == "user-123"
|
||||
assert payload["type"] == "user_auth"
|
||||
assert "sid" in payload
|
||||
assert "iat" in payload
|
||||
assert "exp" in payload
|
||||
|
||||
|
||||
# ── RefreshTokenUseCase tests ───────────────────────────
|
||||
|
||||
|
||||
class TestRefreshTokenUseCase:
|
||||
def test_successful_refresh(self):
|
||||
user = FakeUser()
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
|
||||
# create a session first
|
||||
store.save_session(
|
||||
session_id="sess-1",
|
||||
user_id="user-123",
|
||||
refresh_token="refresh-token-xyz",
|
||||
device_info="Chrome",
|
||||
ip_address="1.2.3.4",
|
||||
expires_in_seconds=3600,
|
||||
)
|
||||
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
req = RefreshTokenRequest(refresh_token="refresh-token-xyz")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert error is None
|
||||
assert response is not None
|
||||
assert response.user_id == "user-123"
|
||||
assert response.access_token
|
||||
# same refresh token returned
|
||||
assert response.refresh_token == "refresh-token-xyz"
|
||||
|
||||
def test_empty_refresh_token(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = RefreshTokenRequest(refresh_token="")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Refresh token is required" in error
|
||||
|
||||
def test_invalid_refresh_token(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
|
||||
req = RefreshTokenRequest(refresh_token="nonexistent-token")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Invalid or expired refresh token" in error
|
||||
|
||||
def test_session_missing_session_id(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
# session with no session_id
|
||||
store._refresh_index["bad-token"] = "bad-sess"
|
||||
store._sessions["bad-sess"] = {"user_id": "user-123"} # no session_id
|
||||
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
req = RefreshTokenRequest(refresh_token="bad-token")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Invalid session data" in error
|
||||
|
||||
def test_session_missing_user_id(self):
|
||||
repo = FakeUserRepository()
|
||||
store = FakeSessionStore()
|
||||
store._refresh_index["bad-token"] = "bad-sess"
|
||||
store._sessions["bad-sess"] = {"session_id": "bad-sess"} # no user_id
|
||||
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
req = RefreshTokenRequest(refresh_token="bad-token")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Invalid session data" in error
|
||||
|
||||
def test_refresh_token_mismatch(self):
|
||||
user = FakeUser()
|
||||
repo = FakeUserRepository(user)
|
||||
store = FakeSessionStore()
|
||||
|
||||
store.save_session(
|
||||
session_id="sess-1",
|
||||
user_id="user-123",
|
||||
refresh_token="original-token",
|
||||
device_info="Chrome",
|
||||
ip_address="1.2.3.4",
|
||||
expires_in_seconds=3600,
|
||||
)
|
||||
|
||||
# Manually add a stale reverse index pointing to same session
|
||||
store._refresh_index["stale-token"] = "sess-1"
|
||||
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
req = RefreshTokenRequest(refresh_token="stale-token")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Refresh token mismatch" in error
|
||||
|
||||
def test_user_not_found(self):
|
||||
repo = FakeUserRepository() # no users
|
||||
store = FakeSessionStore()
|
||||
|
||||
store.save_session(
|
||||
session_id="sess-1",
|
||||
user_id="nonexistent-user",
|
||||
refresh_token="valid-token",
|
||||
device_info="Chrome",
|
||||
ip_address="1.2.3.4",
|
||||
expires_in_seconds=3600,
|
||||
)
|
||||
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
req = RefreshTokenRequest(refresh_token="valid-token")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "User not found" in error
|
||||
|
||||
def test_exception_handling(self):
|
||||
repo = MagicMock()
|
||||
repo.get.side_effect = RuntimeError("DB down")
|
||||
store = FakeSessionStore()
|
||||
|
||||
store.save_session(
|
||||
session_id="sess-1",
|
||||
user_id="user-123",
|
||||
refresh_token="valid-token",
|
||||
device_info="Chrome",
|
||||
ip_address="1.2.3.4",
|
||||
expires_in_seconds=3600,
|
||||
)
|
||||
|
||||
use_case = RefreshTokenUseCase(user_repository=repo, session_store=store)
|
||||
req = RefreshTokenRequest(refresh_token="valid-token")
|
||||
response, error = use_case.execute(req)
|
||||
|
||||
assert response is None
|
||||
assert "Token refresh failed" in error
|
||||
|
||||
|
||||
# ── LogoutUseCase tests ─────────────────────────────────
|
||||
|
||||
|
||||
class TestLogoutUseCase:
|
||||
def test_logout_single_device_success(self):
|
||||
store = FakeSessionStore()
|
||||
store.save_session(
|
||||
session_id="sess-1",
|
||||
user_id="user-123",
|
||||
refresh_token="token1",
|
||||
device_info="Chrome",
|
||||
ip_address="1.2.3.4",
|
||||
expires_in_seconds=3600,
|
||||
)
|
||||
|
||||
use_case = LogoutUseCase(session_store=store)
|
||||
req = LogoutRequest(user_id="user-123", session_id="sess-1")
|
||||
success, error = use_case.execute(req)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
assert "sess-1" in store.deleted_sessions
|
||||
|
||||
def test_logout_single_device_no_session_id(self):
|
||||
store = FakeSessionStore()
|
||||
use_case = LogoutUseCase(session_store=store)
|
||||
req = LogoutRequest(user_id="user-123", session_id=None)
|
||||
success, error = use_case.execute(req)
|
||||
|
||||
assert success is False
|
||||
assert "Session ID is required" in error
|
||||
|
||||
def test_logout_single_device_session_not_found(self):
|
||||
store = FakeSessionStore()
|
||||
use_case = LogoutUseCase(session_store=store)
|
||||
req = LogoutRequest(user_id="user-123", session_id="nonexistent")
|
||||
success, error = use_case.execute(req)
|
||||
|
||||
assert success is False
|
||||
assert "Session not found" in error
|
||||
|
||||
def test_logout_all_devices(self):
|
||||
store = FakeSessionStore()
|
||||
store.delete_all_return_value = 3
|
||||
|
||||
use_case = LogoutUseCase(session_store=store)
|
||||
req = LogoutRequest(user_id="user-123", logout_all_devices=True)
|
||||
success, error = use_case.execute(req)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
assert store.delete_all_called_for == "user-123"
|
||||
|
||||
def test_logout_all_with_empty_session_id(self):
|
||||
store = FakeSessionStore()
|
||||
use_case = LogoutUseCase(session_store=store)
|
||||
# logout_all should work even without session_id
|
||||
req = LogoutRequest(user_id="user-123", session_id=None, logout_all_devices=True)
|
||||
success, error = use_case.execute(req)
|
||||
|
||||
assert success is True
|
||||
assert error is None
|
||||
|
||||
def test_exception_handling(self):
|
||||
store = MagicMock()
|
||||
store.delete_session.side_effect = RuntimeError("Redis down")
|
||||
|
||||
use_case = LogoutUseCase(session_store=store)
|
||||
req = LogoutRequest(user_id="user-123", session_id="sess-1")
|
||||
success, error = use_case.execute(req)
|
||||
|
||||
assert success is False
|
||||
assert "Logout failed" in error
|
||||
Reference in New Issue
Block a user