Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 28a2322f73 |
+298
-387
@@ -1,68 +1,63 @@
|
||||
"""JWT 服务与处理器单元测试."""
|
||||
"""JWT 服务单元测试 — wave130."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import jwt
|
||||
import jwt as pyjwt
|
||||
import pytest
|
||||
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
|
||||
|
||||
from packages.application.auth.jwt_handler import (
|
||||
JWTHandler,
|
||||
configure_jwt_handler,
|
||||
get_jwt_handler,
|
||||
)
|
||||
from packages.application.auth.jwt_service import (
|
||||
JWTConfig,
|
||||
JWTService,
|
||||
TokenType,
|
||||
)
|
||||
|
||||
# ── 测试常量 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
TEST_SECRET = "test-secret-key-for-unit-testing-only-not-for-production"
|
||||
STRONG_SECRET = "x" * 32 # 满足长度要求的测试密钥
|
||||
# ── 测试常量 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
# ── JWTConfig 测试 ───────────────────────────────────────────────────────────
|
||||
TEST_SECRET = "test-secret-key-for-unit-testing-only-1234567890"
|
||||
TEST_ALGORITHM = "HS256"
|
||||
|
||||
|
||||
# ── JWTConfig 配置 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestJWTConfig:
|
||||
"""JWTConfig 配置类测试"""
|
||||
|
||||
def test_init_with_valid_secret(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET)
|
||||
assert config.SECRET_KEY == STRONG_SECRET
|
||||
def test_normal_config(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET)
|
||||
assert config.SECRET_KEY == TEST_SECRET
|
||||
assert config.ALGORITHM == "HS256"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7
|
||||
|
||||
def test_init_custom_values(self):
|
||||
def test_custom_config(self):
|
||||
config = JWTConfig(
|
||||
secret_key=STRONG_SECRET,
|
||||
secret_key=TEST_SECRET,
|
||||
algorithm="HS384",
|
||||
access_token_expire_minutes=60,
|
||||
refresh_token_expire_days=14,
|
||||
refresh_token_expire_days=30,
|
||||
)
|
||||
assert config.SECRET_KEY == STRONG_SECRET
|
||||
assert config.ALGORITHM == "HS384"
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 14
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30
|
||||
|
||||
def test_empty_secret_raises(self):
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
JWTConfig(secret_key="")
|
||||
|
||||
def test_whitespace_only_secret_raises(self):
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
def test_whitespace_secret_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
JWTConfig(secret_key=" ")
|
||||
|
||||
def test_none_secret_raises(self):
|
||||
with pytest.raises(ValueError, match="secret_key must be provided"):
|
||||
JWTConfig(secret_key=None)
|
||||
with pytest.raises(ValueError):
|
||||
JWTConfig(secret_key=None) # type: ignore
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"insecure_secret",
|
||||
"bad_secret",
|
||||
[
|
||||
"your-secret-key-change-in-production",
|
||||
"your-secret-key",
|
||||
@@ -73,407 +68,323 @@ class TestJWTConfig:
|
||||
"Your-Secret-Key",
|
||||
],
|
||||
)
|
||||
def test_insecure_default_secret_raises(self, insecure_secret):
|
||||
def test_insecure_defaults_rejected(self, bad_secret):
|
||||
with pytest.raises(ValueError, match="insecure"):
|
||||
JWTConfig(secret_key=insecure_secret)
|
||||
|
||||
def test_zero_expire_minutes_allowed(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=0)
|
||||
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 0
|
||||
|
||||
def test_negative_expire_days_allowed(self):
|
||||
# 配置类不校验合理性,由业务层判断
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, refresh_token_expire_days=-1)
|
||||
assert config.REFRESH_TOKEN_EXPIRE_DAYS == -1
|
||||
JWTConfig(secret_key=bad_secret)
|
||||
|
||||
|
||||
# ── JWTService 初始化测试 ────────────────────────────────────────────────────
|
||||
# ── JWTService 初始化 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestJWTServiceInit:
|
||||
"""JWTService 初始化测试"""
|
||||
|
||||
def test_init_with_config(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET)
|
||||
def test_with_config_works(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET)
|
||||
service = JWTService(config)
|
||||
assert service.config is config
|
||||
|
||||
def test_init_none_config_raises(self):
|
||||
with pytest.raises(ValueError, match="JWTService requires a JWTConfig"):
|
||||
def test_none_config_raises(self):
|
||||
with pytest.raises(ValueError, match="JWTService requires"):
|
||||
JWTService(None)
|
||||
|
||||
|
||||
# ── TokenType 测试 ───────────────────────────────────────────────────────────
|
||||
# ── create_access_token ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateAccessToken:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
|
||||
def test_creates_valid_jwt(self):
|
||||
token = self.service.create_access_token(user_id="user123")
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
# JWT 格式:xxx.yyy.zzz
|
||||
assert token.count(".") == 2
|
||||
|
||||
def test_payload_contains_user_id(self):
|
||||
token = self.service.create_access_token(user_id="user_001")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["sub"] == "user_001"
|
||||
|
||||
def test_payload_contains_role(self):
|
||||
token = self.service.create_access_token(user_id="u1", role="admin")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["role"] == "admin"
|
||||
|
||||
def test_default_role_empty(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["role"] == ""
|
||||
|
||||
def test_token_type_is_access(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["type"] == TokenType.ACCESS
|
||||
|
||||
def test_has_iat_and_exp(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert "iat" in payload
|
||||
assert "exp" in payload
|
||||
assert payload["exp"] > payload["iat"]
|
||||
|
||||
def test_expiration_correct(self):
|
||||
"""过期时间大约等于当前时间 + 配置的分钟数."""
|
||||
config = JWTConfig(secret_key=TEST_SECRET, access_token_expire_minutes=30)
|
||||
service = JWTService(config)
|
||||
before = datetime.now(timezone.utc)
|
||||
token = service.create_access_token(user_id="u1")
|
||||
after = datetime.now(timezone.utc)
|
||||
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
exp = datetime.fromtimestamp(payload["exp"], tz=timezone.utc)
|
||||
|
||||
min_expected = before + timedelta(minutes=30) - timedelta(seconds=1)
|
||||
max_expected = after + timedelta(minutes=30) + timedelta(seconds=1)
|
||||
assert min_expected <= exp <= max_expected
|
||||
|
||||
def test_additional_claims_included(self):
|
||||
extra = {"email": "test@example.com", "org_id": "org_001", "level": 5}
|
||||
token = self.service.create_access_token(user_id="u1", additional_claims=extra)
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["email"] == "test@example.com"
|
||||
assert payload["org_id"] == "org_001"
|
||||
assert payload["level"] == 5
|
||||
|
||||
def test_additional_claims_none(self):
|
||||
token = self.service.create_access_token(user_id="u1", additional_claims=None)
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert "email" not in payload
|
||||
|
||||
def test_signed_with_correct_key(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
# 用正确的密钥可以解码
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["sub"] == "u1"
|
||||
# 用错误的密钥无法解码
|
||||
with pytest.raises(InvalidTokenError):
|
||||
pyjwt.decode(token, "wrong-secret", algorithms=["HS256"])
|
||||
|
||||
|
||||
# ── create_refresh_token ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateRefreshToken:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
|
||||
def test_creates_valid_token(self):
|
||||
token = self.service.create_refresh_token(user_id="u1", session_id="sess_001")
|
||||
assert isinstance(token, str)
|
||||
assert token.count(".") == 2
|
||||
|
||||
def test_payload_contains_session_id(self):
|
||||
token = self.service.create_refresh_token(user_id="u1", session_id="sess_abc")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["session_id"] == "sess_abc"
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
def test_token_type_is_refresh(self):
|
||||
token = self.service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
assert payload["type"] == TokenType.REFRESH
|
||||
|
||||
def test_refresh_expiration_days(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET, refresh_token_expire_days=7)
|
||||
service = JWTService(config)
|
||||
before = datetime.now(timezone.utc)
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
after = datetime.now(timezone.utc)
|
||||
|
||||
payload = pyjwt.decode(token, TEST_SECRET, algorithms=["HS256"])
|
||||
exp = datetime.fromtimestamp(payload["exp"], tz=timezone.utc)
|
||||
|
||||
min_exp = before + timedelta(days=7) - timedelta(seconds=1)
|
||||
max_exp = after + timedelta(days=7, seconds=1)
|
||||
assert min_exp <= exp <= max_exp
|
||||
|
||||
|
||||
# ── verify_token ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerifyToken:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
|
||||
def test_valid_token_returns_payload(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
payload = self.service.verify_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
def test_expired_token_raises(self):
|
||||
# 创建一个 1 秒过期的 token
|
||||
config = JWTConfig(secret_key=TEST_SECRET, access_token_expire_minutes=1)
|
||||
service = JWTService(config)
|
||||
token = service.create_access_token(user_id="u1")
|
||||
|
||||
# 等待过期(用 pyjwt 直接构造过期 token 更可靠)
|
||||
expired_payload = {
|
||||
"sub": "u1",
|
||||
"type": "access",
|
||||
"exp": datetime.now(timezone.utc) - timedelta(seconds=10),
|
||||
}
|
||||
expired_token = pyjwt.encode(expired_payload, TEST_SECRET, algorithm="HS256")
|
||||
|
||||
with pytest.raises(ExpiredSignatureError, match="expired"):
|
||||
self.service.verify_token(expired_token)
|
||||
|
||||
def test_invalid_token_raises(self):
|
||||
with pytest.raises(InvalidTokenError, match="Invalid token"):
|
||||
self.service.verify_token("not-a-valid-jwt-token")
|
||||
|
||||
def test_wrong_signature_raises(self):
|
||||
token = pyjwt.encode({"sub": "u1"}, "different-secret", algorithm="HS256")
|
||||
with pytest.raises(InvalidTokenError):
|
||||
self.service.verify_token(token)
|
||||
|
||||
def test_tampered_payload_raises(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
# 尝试篡改:JWT 有签名保护,篡改会导致验证失败
|
||||
parts = token.split(".")
|
||||
assert len(parts) == 3
|
||||
# 把 payload 部分替换(不会成功,因为签名不对)
|
||||
import base64
|
||||
|
||||
fake_payload = base64.urlsafe_b64encode(b'{"sub":"admin","role":"admin"}').rstrip(b"=").decode()
|
||||
tampered = f"{parts[0]}.{fake_payload}.{parts[2]}"
|
||||
with pytest.raises(InvalidTokenError):
|
||||
self.service.verify_token(tampered)
|
||||
|
||||
|
||||
# ── verify_access_token ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerifyAccessToken:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
|
||||
def test_access_token_passes(self):
|
||||
token = self.service.create_access_token(user_id="u1", role="user")
|
||||
payload = self.service.verify_access_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["type"] == "access"
|
||||
|
||||
def test_refresh_token_rejected(self):
|
||||
token = self.service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
with pytest.raises(ValueError, match="Token type must be 'access'"):
|
||||
self.service.verify_access_token(token)
|
||||
|
||||
def test_expired_token_raises(self):
|
||||
expired_payload = {
|
||||
"sub": "u1",
|
||||
"type": "access",
|
||||
"exp": datetime.now(timezone.utc) - timedelta(seconds=10),
|
||||
}
|
||||
token = pyjwt.encode(expired_payload, TEST_SECRET, algorithm="HS256")
|
||||
with pytest.raises(ExpiredSignatureError):
|
||||
self.service.verify_access_token(token)
|
||||
|
||||
|
||||
# ── verify_refresh_token ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerifyRefreshToken:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
|
||||
def test_refresh_token_passes(self):
|
||||
token = self.service.create_refresh_token(user_id="u1", session_id="sess_001")
|
||||
payload = self.service.verify_refresh_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["session_id"] == "sess_001"
|
||||
|
||||
def test_access_token_rejected(self):
|
||||
token = self.service.create_access_token(user_id="u1")
|
||||
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
|
||||
self.service.verify_refresh_token(token)
|
||||
|
||||
def test_has_session_id(self):
|
||||
token = self.service.create_refresh_token(user_id="u1", session_id="custom_sess")
|
||||
payload = self.service.verify_refresh_token(token)
|
||||
assert payload["session_id"] == "custom_sess"
|
||||
|
||||
|
||||
# ── TokenType 常量 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTokenType:
|
||||
"""TokenType 常量测试"""
|
||||
|
||||
def test_access_value(self):
|
||||
assert TokenType.ACCESS == "access"
|
||||
|
||||
def test_refresh_value(self):
|
||||
assert TokenType.REFRESH == "refresh"
|
||||
|
||||
def test_access_and_refresh_different(self):
|
||||
def test_different_types(self):
|
||||
assert TokenType.ACCESS != TokenType.REFRESH
|
||||
|
||||
|
||||
# ── JWTService create_access_token 测试 ─────────────────────────────────────
|
||||
# ── 多算法支持 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateAccessToken:
|
||||
"""创建 access_token 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def service(self):
|
||||
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
|
||||
|
||||
def test_creates_valid_jwt_string(self, service):
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_token_contains_user_id_as_sub(self, service):
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["sub"] == "user-123"
|
||||
|
||||
def test_token_type_is_access(self, service):
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["type"] == TokenType.ACCESS
|
||||
|
||||
def test_default_role_is_empty_string(self, service):
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["role"] == ""
|
||||
|
||||
def test_custom_role(self, service):
|
||||
token = service.create_access_token(user_id="user-123", role="admin")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["role"] == "admin"
|
||||
|
||||
def test_has_iat_and_exp(self, service):
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert "iat" in payload
|
||||
assert "exp" in payload
|
||||
assert payload["exp"] > payload["iat"]
|
||||
|
||||
def test_expire_matches_config(self, service):
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
iat = datetime.fromtimestamp(payload["iat"], tz=timezone.utc)
|
||||
exp = datetime.fromtimestamp(payload["exp"], tz=timezone.utc)
|
||||
delta = exp - iat
|
||||
assert delta.total_seconds() == 15 * 60 # 15分钟
|
||||
|
||||
def test_custom_expire_time(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=30)
|
||||
class TestDifferentAlgorithms:
|
||||
def test_hs384_works(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET * 2, algorithm="HS384")
|
||||
service = JWTService(config)
|
||||
token = service.create_access_token(user_id="user-123")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
delta = payload["exp"] - payload["iat"]
|
||||
assert delta == 30 * 60
|
||||
|
||||
def test_additional_claims(self, service):
|
||||
extra = {"custom_field": "value", "another": 42}
|
||||
token = service.create_access_token(user_id="user-123", additional_claims=extra)
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["custom_field"] == "value"
|
||||
assert payload["another"] == 42
|
||||
|
||||
def test_additional_claims_can_override_standard(self, service):
|
||||
# additional_claims 可以覆盖标准字段(由调用者负责)
|
||||
token = service.create_access_token(
|
||||
user_id="user-123",
|
||||
additional_claims={"sub": "overridden"},
|
||||
)
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["sub"] == "overridden"
|
||||
|
||||
def test_additional_claims_none_is_same_as_empty(self, service):
|
||||
token = service.create_access_token(user_id="user-123", additional_claims=None)
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["sub"] == "user-123"
|
||||
|
||||
def test_uses_correct_algorithm(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, algorithm="HS384")
|
||||
service = JWTService(config)
|
||||
token = service.create_access_token(user_id="u1")
|
||||
# 用 HS256 解码应该失败
|
||||
with pytest.raises(InvalidTokenError):
|
||||
jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
# 用 HS384 解码应该成功
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS384"])
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
|
||||
# ── JWTService create_refresh_token 测试 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestCreateRefreshToken:
|
||||
"""创建 refresh_token 测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def service(self):
|
||||
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
|
||||
|
||||
def test_creates_valid_string(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
assert isinstance(token, str)
|
||||
assert len(token) > 0
|
||||
|
||||
def test_contains_user_id_and_session_id(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="sess-abc")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["session_id"] == "sess-abc"
|
||||
|
||||
def test_token_type_is_refresh(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert payload["type"] == TokenType.REFRESH
|
||||
|
||||
def test_has_iat_and_exp(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
assert "iat" in payload
|
||||
assert "exp" in payload
|
||||
assert payload["exp"] > payload["iat"]
|
||||
|
||||
def test_expire_matches_config_days(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
delta = payload["exp"] - payload["iat"]
|
||||
assert delta == 7 * 24 * 60 * 60 # 7天
|
||||
|
||||
def test_custom_refresh_expire_days(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, refresh_token_expire_days=30)
|
||||
service = JWTService(config)
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
|
||||
delta = payload["exp"] - payload["iat"]
|
||||
assert delta == 30 * 24 * 60 * 60
|
||||
|
||||
|
||||
# ── JWTService verify_token 测试 ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerifyToken:
|
||||
"""通用 Token 验证测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def service(self):
|
||||
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
|
||||
|
||||
def test_verify_valid_access_token(self, service):
|
||||
token = service.create_access_token(user_id="u1")
|
||||
payload = service.verify_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["type"] == TokenType.ACCESS
|
||||
|
||||
def test_verify_valid_refresh_token(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = service.verify_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["session_id"] == "s1"
|
||||
|
||||
def test_verify_expired_token_raises(self, service):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=0)
|
||||
svc = JWTService(config)
|
||||
token = svc.create_access_token(user_id="u1")
|
||||
# 0 分钟过期,立即过期
|
||||
time.sleep(0.1) # 稍微等一下确保过期
|
||||
with pytest.raises(ExpiredSignatureError):
|
||||
svc.verify_token(token)
|
||||
|
||||
def test_verify_wrong_secret_raises(self, service):
|
||||
token = service.create_access_token(user_id="u1")
|
||||
other_service = JWTService(JWTConfig(secret_key="different-secret-1234567890"))
|
||||
with pytest.raises(InvalidTokenError):
|
||||
other_service.verify_token(token)
|
||||
|
||||
def test_verify_tampered_token_raises(self, service):
|
||||
token = service.create_access_token(user_id="u1")
|
||||
# 篡改 token 中间部分
|
||||
parts = token.split(".")
|
||||
assert len(parts) == 3
|
||||
tampered = parts[0] + "." + parts[1][:-1] + "A." + parts[2]
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service.verify_token(tampered)
|
||||
|
||||
def test_verify_empty_string_raises(self, service):
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service.verify_token("")
|
||||
|
||||
def test_verify_garbage_string_raises(self, service):
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service.verify_token("not.a.valid.jwt.token")
|
||||
|
||||
def test_verify_returns_dict(self, service):
|
||||
token = service.create_access_token(user_id="u1", role="admin")
|
||||
payload = service.verify_token(token)
|
||||
assert isinstance(payload, dict)
|
||||
assert "sub" in payload
|
||||
assert "role" in payload
|
||||
|
||||
|
||||
# ── JWTService verify_access_token 测试 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerifyAccessToken:
|
||||
"""Access Token 专属验证测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def service(self):
|
||||
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
|
||||
|
||||
def test_valid_access_token_passes(self, service):
|
||||
token = service.create_access_token(user_id="u1", role="admin")
|
||||
payload = service.verify_access_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["role"] == "admin"
|
||||
|
||||
def test_refresh_token_fails_type_check(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
with pytest.raises(ValueError, match="Token type must be 'access'"):
|
||||
service.verify_access_token(token)
|
||||
|
||||
def test_token_without_type_field_raises(self, service):
|
||||
# 手动构造一个没有 type 字段的 token
|
||||
payload_data = {"sub": "u1", "iat": 1000, "exp": 9999999999}
|
||||
token = jwt.encode(payload_data, STRONG_SECRET, algorithm="HS256")
|
||||
with pytest.raises(ValueError, match="Token type must be 'access'"):
|
||||
service.verify_access_token(token)
|
||||
|
||||
def test_expired_access_token_raises_expired_error(self, service):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, access_token_expire_minutes=0)
|
||||
svc = JWTService(config)
|
||||
token = svc.create_access_token(user_id="u1")
|
||||
time.sleep(0.1)
|
||||
with pytest.raises(ExpiredSignatureError):
|
||||
svc.verify_access_token(token)
|
||||
|
||||
|
||||
# ── JWTService verify_refresh_token 测试 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerifyRefreshToken:
|
||||
"""Refresh Token 专属验证测试"""
|
||||
|
||||
@pytest.fixture
|
||||
def service(self):
|
||||
return JWTService(JWTConfig(secret_key=STRONG_SECRET))
|
||||
|
||||
def test_valid_refresh_token_passes(self, service):
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
payload = service.verify_refresh_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["session_id"] == "s1"
|
||||
|
||||
def test_access_token_fails_type_check(self, service):
|
||||
token = service.create_access_token(user_id="u1")
|
||||
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
|
||||
service.verify_refresh_token(token)
|
||||
|
||||
def test_token_without_type_field_raises(self, service):
|
||||
payload_data = {"sub": "u1", "session_id": "s1", "iat": 1000, "exp": 9999999999}
|
||||
token = jwt.encode(payload_data, STRONG_SECRET, algorithm="HS256")
|
||||
with pytest.raises(ValueError, match="Token type must be 'refresh'"):
|
||||
service.verify_refresh_token(token)
|
||||
|
||||
def test_expired_refresh_token_raises(self):
|
||||
config = JWTConfig(secret_key=STRONG_SECRET, refresh_token_expire_days=0)
|
||||
def test_hs512_works(self):
|
||||
config = JWTConfig(secret_key=TEST_SECRET * 3, algorithm="HS512")
|
||||
service = JWTService(config)
|
||||
token = service.create_refresh_token(user_id="u1", session_id="s1")
|
||||
# 0天过期,应该立即使exp <= iat
|
||||
with pytest.raises(ExpiredSignatureError):
|
||||
service.verify_refresh_token(token)
|
||||
|
||||
|
||||
# ── JWTHandler 委托层测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestJWTHandler:
|
||||
"""JWTHandler 委托层测试"""
|
||||
|
||||
def test_init_creates_handler(self):
|
||||
handler = JWTHandler(secret_key=STRONG_SECRET)
|
||||
assert handler is not None
|
||||
|
||||
def test_create_and_verify_access_token(self):
|
||||
handler = JWTHandler(secret_key=STRONG_SECRET)
|
||||
token = handler.create_access_token(user_id="u1", role="user")
|
||||
payload = handler.verify_access_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
assert payload["role"] == "user"
|
||||
|
||||
def test_verify_token_generic(self):
|
||||
handler = JWTHandler(secret_key=STRONG_SECRET)
|
||||
token = handler.create_access_token(user_id="u1")
|
||||
payload = handler.verify_token(token)
|
||||
token = service.create_access_token(user_id="u1")
|
||||
payload = service.verify_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
|
||||
def test_custom_algorithm(self):
|
||||
handler = JWTHandler(secret_key=STRONG_SECRET, algorithm="HS384")
|
||||
token = handler.create_access_token(user_id="u1")
|
||||
payload = handler.verify_access_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
def test_algorithm_mismatch_fails(self):
|
||||
config_hs256 = JWTConfig(secret_key=TEST_SECRET, algorithm="HS256")
|
||||
config_hs384 = JWTConfig(secret_key=TEST_SECRET, algorithm="HS384")
|
||||
service_256 = JWTService(config_hs256)
|
||||
service_384 = JWTService(config_hs384)
|
||||
|
||||
def test_custom_expire_minutes(self):
|
||||
handler = JWTHandler(secret_key=STRONG_SECRET, access_token_expire_minutes=45)
|
||||
token = handler.create_access_token(user_id="u1")
|
||||
payload = handler.verify_access_token(token)
|
||||
delta = payload["exp"] - payload["iat"]
|
||||
assert delta == 45 * 60
|
||||
|
||||
def test_additional_claims_passthrough(self):
|
||||
handler = JWTHandler(secret_key=STRONG_SECRET)
|
||||
extra = {"org_id": "org-1", "plan": "pro"}
|
||||
token = handler.create_access_token("u1", additional_claims={"org_id": "org-1"})
|
||||
payload = handler.verify_access_token(
|
||||
token := handler.create_access_token("u1", additional_claims={"org_id": "org-1"})
|
||||
)
|
||||
# 这里直接测试更简洁
|
||||
payload = handler.verify_access_token(handler.create_access_token("u1", additional_claims={"x": 1}))
|
||||
assert payload["x"] == 1
|
||||
token = service_256.create_access_token(user_id="u1")
|
||||
with pytest.raises(InvalidTokenError):
|
||||
service_384.verify_token(token)
|
||||
|
||||
|
||||
# ── 全局 JWT handler 测试 ───────────────────────────────────────────────────
|
||||
# ── 边界:空用户ID等 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGlobalJWTHandler:
|
||||
"""全局 JWT Handler 配置与获取测试"""
|
||||
class TestEdgeCases:
|
||||
def setup_method(self):
|
||||
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
|
||||
|
||||
def test_configure_creates_handler(self):
|
||||
handler = configure_jwt_handler(secret_key=STRONG_SECRET)
|
||||
assert isinstance(handler, JWTHandler)
|
||||
def test_empty_user_id(self):
|
||||
token = self.service.create_access_token(user_id="")
|
||||
payload = self.service.verify_access_token(token)
|
||||
assert payload["sub"] == ""
|
||||
|
||||
def test_get_after_configure_works(self):
|
||||
configure_jwt_handler(secret_key=STRONG_SECRET)
|
||||
handler = get_jwt_handler()
|
||||
assert isinstance(handler, JWTHandler)
|
||||
token = handler.create_access_token(user_id="u1")
|
||||
payload = handler.verify_access_token(token)
|
||||
assert payload["sub"] == "u1"
|
||||
def test_long_user_id(self):
|
||||
long_id = "x" * 1000
|
||||
token = self.service.create_access_token(user_id=long_id)
|
||||
payload = self.service.verify_access_token(token)
|
||||
assert payload["sub"] == long_id
|
||||
|
||||
def test_get_before_configure_raises(self):
|
||||
# 重置全局状态(通过设置 None 模拟未配置)
|
||||
import packages.application.auth.jwt_handler as mod
|
||||
def test_special_chars_in_user_id(self):
|
||||
uid = "user@#$%^&*()_+-=[]{}|;:',.<>?/`~"
|
||||
token = self.service.create_access_token(user_id=uid)
|
||||
payload = self.service.verify_access_token(token)
|
||||
assert payload["sub"] == uid
|
||||
|
||||
mod._default_handler = None
|
||||
with pytest.raises(RuntimeError, match="JWT handler not configured"):
|
||||
get_jwt_handler()
|
||||
def test_unicode_user_id(self):
|
||||
uid = "用户_测试_123_🎉"
|
||||
token = self.service.create_access_token(user_id=uid)
|
||||
payload = self.service.verify_access_token(token)
|
||||
assert payload["sub"] == uid
|
||||
|
||||
def test_configure_returns_same_as_get(self):
|
||||
h1 = configure_jwt_handler(secret_key=STRONG_SECRET)
|
||||
h2 = get_jwt_handler()
|
||||
assert h1 is h2
|
||||
|
||||
def test_reconfigure_replaces_handler(self):
|
||||
h1 = configure_jwt_handler(secret_key=STRONG_SECRET)
|
||||
h2 = configure_jwt_handler(secret_key=STRONG_SECRET + "_new")
|
||||
assert h1 is not h2
|
||||
assert get_jwt_handler() is h2
|
||||
def test_many_additional_claims(self):
|
||||
claims = {f"key_{i}": f"value_{i}" for i in range(50)}
|
||||
token = self.service.create_access_token(user_id="u1", additional_claims=claims)
|
||||
payload = self.service.verify_access_token(token)
|
||||
for i in range(50):
|
||||
assert payload[f"key_{i}"] == f"value_{i}"
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""验证码服务单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
@@ -9,9 +8,9 @@ import pytest
|
||||
|
||||
from packages.application.auth.verification_code_service import (
|
||||
CODE_TYPE_EMAIL_BIND,
|
||||
CODE_TYPE_EMAIL_LOGIN,
|
||||
CODE_TYPE_PHONE_BIND,
|
||||
DAILY_LIMIT,
|
||||
DEFAULT_TTL_SECONDS,
|
||||
MAX_ATTEMPTS,
|
||||
RESEND_COOLDOWN_SECONDS,
|
||||
VerificationCodeService,
|
||||
@@ -21,298 +20,560 @@ from packages.application.auth.verification_code_service import (
|
||||
)
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
|
||||
# ── Test Fixtures ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
"""mock 验证码仓储."""
|
||||
repo = MagicMock()
|
||||
repo.find_latest.return_value = None
|
||||
repo.count_today.return_value = 0
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def code_service(mock_repo):
|
||||
return VerificationCodeService(mock_repo)
|
||||
def service(mock_repo):
|
||||
"""验证码服务实例."""
|
||||
return VerificationCodeService(repo=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_code():
|
||||
code = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
ttl_seconds=300,
|
||||
def _make_code(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
code="123456",
|
||||
ttl=300,
|
||||
used=False,
|
||||
attempts=0,
|
||||
created_at=None,
|
||||
):
|
||||
"""创建一个测试用验证码实体."""
|
||||
now = created_at or datetime.now(timezone.utc)
|
||||
vc = VerificationCode(
|
||||
id="test-code-id",
|
||||
recipient=recipient,
|
||||
code=code,
|
||||
code_type=code_type,
|
||||
expires_at=now + timedelta(seconds=ttl),
|
||||
used_at=now if used else None,
|
||||
attempts=attempts,
|
||||
created_at=now,
|
||||
)
|
||||
return code
|
||||
return vc
|
||||
|
||||
|
||||
class TestVerificationCodeServiceGenerate:
|
||||
# ── generate 方法测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerate:
|
||||
"""generate 方法测试"""
|
||||
|
||||
def test_generate_success(self, code_service, mock_repo, sample_code):
|
||||
"""生成验证码成功"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
def test_generate_success(self, service, mock_repo):
|
||||
"""成功生成验证码."""
|
||||
code, error = service.generate("user@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
assert error is None
|
||||
assert code is not None
|
||||
assert code.recipient == "test@example.com"
|
||||
assert code.recipient == "user@example.com"
|
||||
assert code.code_type == CODE_TYPE_EMAIL_BIND
|
||||
assert len(code.code) == 6
|
||||
assert code.code.isdigit()
|
||||
assert not code.is_used
|
||||
mock_repo.save.assert_called_once()
|
||||
|
||||
def test_generate_empty_recipient(self, code_service):
|
||||
"""空接收方返回错误"""
|
||||
code, error = code_service.generate("", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "接收方不能为空" in error
|
||||
def test_generate_with_custom_code(self, service, mock_repo):
|
||||
"""使用自定义验证码."""
|
||||
code, error = service.generate("user@example.com", CODE_TYPE_EMAIL_LOGIN, custom_code="999999")
|
||||
|
||||
def test_generate_invalid_type(self, code_service):
|
||||
"""无效验证码类型返回错误"""
|
||||
code, error = code_service.generate("test@example.com", "invalid_type")
|
||||
assert error is None
|
||||
assert code.code == "999999"
|
||||
|
||||
def test_generate_custom_ttl(self, service, mock_repo):
|
||||
"""自定义 TTL."""
|
||||
code, _ = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=600)
|
||||
delta = code.expires_at - code.created_at
|
||||
assert delta.total_seconds() == 600
|
||||
|
||||
def test_generate_default_ttl(self, service, mock_repo):
|
||||
"""默认 TTL."""
|
||||
code, _ = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
delta = code.expires_at - code.created_at
|
||||
assert delta.total_seconds() == 300 # 默认5分钟
|
||||
|
||||
def test_generate_empty_recipient(self, service):
|
||||
"""空接收方."""
|
||||
code, error = service.generate("", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "不能为空" in error
|
||||
|
||||
def test_generate_whitespace_recipient(self, service):
|
||||
"""全空白接收方."""
|
||||
code, error = service.generate(" ", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "不能为空" in error
|
||||
|
||||
def test_generate_invalid_type(self, service):
|
||||
"""无效验证码类型."""
|
||||
code, error = service.generate("u@e.com", "invalid_type")
|
||||
assert code is None
|
||||
assert "无效的验证码类型" in error
|
||||
|
||||
def test_generate_cooldown(self, code_service, mock_repo, sample_code):
|
||||
"""冷却期内返回频控错误"""
|
||||
# 最新的验证码刚创建10秒前
|
||||
sample_code.created_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
mock_repo.count_today.return_value = 1
|
||||
def test_generate_recipient_stripped(self, service, mock_repo):
|
||||
"""接收方前后空格会被清理."""
|
||||
code, _ = service.generate(" user@e.com ", CODE_TYPE_EMAIL_BIND)
|
||||
assert code.recipient == "user@e.com"
|
||||
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
def test_generate_phone_code(self, service, mock_repo):
|
||||
"""手机验证码生成."""
|
||||
code, error = service.generate("13800138000", CODE_TYPE_PHONE_BIND)
|
||||
assert error is None
|
||||
assert code.code_type == CODE_TYPE_PHONE_BIND
|
||||
assert len(code.code) == 6
|
||||
|
||||
|
||||
# ── generate 频控测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGenerateRateLimit:
|
||||
"""generate 频控测试"""
|
||||
|
||||
def test_cooldown_active_rejects(self, service, mock_repo):
|
||||
"""冷却期内拒绝重发."""
|
||||
recent = _make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=10))
|
||||
mock_repo.find_latest.return_value = recent
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "发送太频繁" in error
|
||||
assert "秒后再试" in error
|
||||
# 等待时间应该接近 50 秒 (60-10)
|
||||
match = re.search(r"(\d+)\s*秒", error)
|
||||
assert match
|
||||
wait = int(match.group(1))
|
||||
assert 45 <= wait <= 55
|
||||
|
||||
def test_generate_daily_limit_exceeded(self, code_service, mock_repo):
|
||||
"""超过每日上限返回错误"""
|
||||
mock_repo.find_latest.return_value = None # 没有冷却期问题
|
||||
def test_cooldown_expired_allows(self, service, mock_repo):
|
||||
"""冷却期过后允许重发."""
|
||||
old = _make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=120))
|
||||
mock_repo.find_latest.return_value = old
|
||||
mock_repo.count_today.return_value = 1
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert error is None
|
||||
assert code is not None
|
||||
|
||||
def test_daily_limit_reached(self, service, mock_repo):
|
||||
"""达到每日上限."""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT
|
||||
|
||||
code, error = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "今日发送次数已达上限" in error
|
||||
|
||||
def test_generate_recipient_stripped(self, code_service, mock_repo, sample_code):
|
||||
"""recipient 会被 strip"""
|
||||
def test_daily_limit_one_below_allows(self, service, mock_repo):
|
||||
"""未达到上限时允许."""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = DAILY_LIMIT - 1
|
||||
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert error is None
|
||||
assert code is not None
|
||||
|
||||
def test_custom_daily_limit(self, mock_repo):
|
||||
"""自定义每日上限."""
|
||||
svc = VerificationCodeService(repo=mock_repo, daily_limit=3)
|
||||
mock_repo.count_today.return_value = 3
|
||||
|
||||
code, error = svc.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
assert "已达上限" in error
|
||||
|
||||
def test_custom_cooldown(self, mock_repo):
|
||||
"""自定义冷却时间."""
|
||||
svc = VerificationCodeService(repo=mock_repo, resend_cooldown=30)
|
||||
recent = _make_code(created_at=datetime.now(timezone.utc) - timedelta(seconds=10))
|
||||
mock_repo.find_latest.return_value = recent
|
||||
|
||||
code, error = svc.generate("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert code is None
|
||||
match = re.search(r"(\d+)\s*秒", error)
|
||||
assert match
|
||||
wait = int(match.group(1))
|
||||
assert 15 <= wait <= 25
|
||||
|
||||
def test_cooldown_different_types_independent(self, service, mock_repo):
|
||||
"""不同类型的验证码冷却独立."""
|
||||
# email_bind 类型有一个近期验证码
|
||||
recent = _make_code(code_type=CODE_TYPE_EMAIL_BIND)
|
||||
mock_repo.find_latest.side_effect = lambda r, t: recent if t == CODE_TYPE_EMAIL_BIND else None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code_service.generate(" test@example.com ", CODE_TYPE_EMAIL_BIND)
|
||||
|
||||
# 传给 repo 的应该是 strip 后的值
|
||||
save_call = mock_repo.save.call_args[0][0]
|
||||
assert save_call.recipient == "test@example.com"
|
||||
|
||||
def test_generate_custom_code(self, code_service, mock_repo):
|
||||
"""使用自定义验证码"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, _ = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, custom_code="123456")
|
||||
assert code.code == "123456"
|
||||
|
||||
def test_generate_custom_ttl(self, code_service, mock_repo):
|
||||
"""自定义 TTL"""
|
||||
mock_repo.find_latest.return_value = None
|
||||
mock_repo.count_today.return_value = 0
|
||||
mock_repo.save.return_value = None
|
||||
|
||||
code, _ = code_service.generate("test@example.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=600)
|
||||
# email_login 类型应该可以正常发送
|
||||
code, error = service.generate("u@e.com", CODE_TYPE_EMAIL_LOGIN)
|
||||
assert error is None
|
||||
assert code is not None
|
||||
|
||||
|
||||
class TestVerificationCodeServiceVerify:
|
||||
# ── verify 方法测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerify:
|
||||
"""verify 方法测试"""
|
||||
|
||||
def test_verify_success(self, code_service, mock_repo, sample_code):
|
||||
"""验证成功"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
def test_verify_success(self, service, mock_repo):
|
||||
"""验证码正确."""
|
||||
code = _make_code(code="654321")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code)
|
||||
|
||||
assert success is True
|
||||
ok, error = service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "654321")
|
||||
assert ok is True
|
||||
assert error is None
|
||||
assert sample_code.is_used is True
|
||||
assert code.is_used # 标记为已使用
|
||||
assert mock_repo.save.call_count >= 2 # increment + mark_used
|
||||
|
||||
def test_verify_wrong_code(self, code_service, mock_repo, sample_code):
|
||||
"""验证码错误"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
def test_verify_wrong_code(self, service, mock_repo):
|
||||
"""验证码错误."""
|
||||
code = _make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "wrongcode")
|
||||
|
||||
assert success is False
|
||||
ok, error = service.verify("test@e.com", CODE_TYPE_EMAIL_BIND, "000000")
|
||||
assert ok is False
|
||||
assert "验证码错误" in error
|
||||
assert not code.is_used # 不标记为已使用
|
||||
assert code.attempts == 1 # 尝试次数+1
|
||||
|
||||
def test_verify_not_found(self, code_service, mock_repo):
|
||||
"""验证码不存在"""
|
||||
def test_verify_no_code_found(self, service, mock_repo):
|
||||
"""找不到验证码."""
|
||||
mock_repo.find_latest.return_value = None
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
|
||||
assert success is False
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert ok is False
|
||||
assert "不存在或已过期" in error
|
||||
|
||||
def test_verify_expired(self, code_service, mock_repo):
|
||||
"""验证码已过期"""
|
||||
expired_code = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type=CODE_TYPE_EMAIL_BIND,
|
||||
ttl_seconds=1, # 1秒过期
|
||||
)
|
||||
# 手动设置过期时间
|
||||
expired_code.expires_at = datetime.now(timezone.utc) - timedelta(seconds=10)
|
||||
mock_repo.find_latest.return_value = expired_code
|
||||
def test_verify_empty_params(self, service):
|
||||
"""参数为空."""
|
||||
ok, error = service.verify("", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert ok is False
|
||||
assert "参数不完整" in error
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, expired_code.code)
|
||||
ok2, error2 = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "")
|
||||
assert ok2 is False
|
||||
assert "参数不完整" in error2
|
||||
|
||||
assert success is False
|
||||
assert "已过期" in error
|
||||
def test_verify_whitespace_params(self, service, mock_repo):
|
||||
"""参数前后空格会被清理."""
|
||||
code = _make_code(recipient="u@e.com", code="111111")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
def test_verify_already_used(self, code_service, mock_repo, sample_code):
|
||||
"""验证码已使用"""
|
||||
sample_code.mark_used()
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
ok, error = service.verify(" u@e.com ", CODE_TYPE_EMAIL_BIND, " 111111 ")
|
||||
assert ok is True
|
||||
assert error is None
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code)
|
||||
def test_verify_already_used(self, service, mock_repo):
|
||||
"""验证码已使用."""
|
||||
code = _make_code(used=True)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
assert success is False
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "已使用" in error
|
||||
|
||||
def test_verify_max_attempts_exceeded(self, code_service, mock_repo, sample_code):
|
||||
"""尝试次数过多"""
|
||||
# 先把尝试次数加到超过上限
|
||||
for _ in range(MAX_ATTEMPTS + 1):
|
||||
sample_code.increment_attempts()
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
def test_verify_expired(self, service, mock_repo):
|
||||
"""验证码已过期."""
|
||||
code = _make_code(ttl=-60) # 已过期
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code)
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "已过期" in error
|
||||
|
||||
assert success is False
|
||||
def test_verify_too_many_attempts(self, service, mock_repo):
|
||||
"""尝试次数过多."""
|
||||
code = _make_code(attempts=MAX_ATTEMPTS + 1)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "验证次数过多" in error
|
||||
|
||||
def test_verify_empty_params(self, code_service):
|
||||
"""空参数返回错误"""
|
||||
success, error = code_service.verify("", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert success is False
|
||||
assert "参数不完整" in error
|
||||
def test_verify_attempts_increment_each_time(self, service, mock_repo):
|
||||
"""每次错误尝试都增加尝试次数."""
|
||||
code = _make_code(code="123456", attempts=0)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
success, error = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "")
|
||||
assert success is False
|
||||
assert "参数不完整" in error
|
||||
for _ in range(3):
|
||||
service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "wrong")
|
||||
|
||||
def test_verify_increments_attempts(self, code_service, mock_repo, sample_code):
|
||||
"""验证会增加尝试次数"""
|
||||
initial_attempts = sample_code.attempts
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
assert code.attempts == 3
|
||||
|
||||
code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, "wrong")
|
||||
def test_verify_without_consume(self, service, mock_repo):
|
||||
"""验证成功但不标记为已使用(consume=False)."""
|
||||
code = _make_code(code="999999")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
assert sample_code.attempts == initial_attempts + 1
|
||||
|
||||
def test_verify_no_consume(self, code_service, mock_repo, sample_code):
|
||||
"""consume=False 时不标记为已使用"""
|
||||
mock_repo.find_latest.return_value = sample_code
|
||||
|
||||
success, _ = code_service.verify("test@example.com", CODE_TYPE_EMAIL_BIND, sample_code.code, consume=False)
|
||||
|
||||
assert success is True
|
||||
assert sample_code.is_used is False
|
||||
|
||||
|
||||
class TestVerifyPhone:
|
||||
"""validate_phone 函数测试"""
|
||||
|
||||
def test_valid_phone(self):
|
||||
"""有效手机号"""
|
||||
ok, err = validate_phone("13800000001")
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "999999", consume=False)
|
||||
assert ok is True
|
||||
assert err == ""
|
||||
assert error is None
|
||||
assert not code.is_used # 不标记为已使用
|
||||
|
||||
def test_valid_phone_with_plus86(self):
|
||||
"""带 +86 前缀的手机号"""
|
||||
ok, err = validate_phone("+8613800000001")
|
||||
def test_verify_consume_default_true(self, service, mock_repo):
|
||||
"""默认 consume=True."""
|
||||
code = _make_code(code="123456")
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, "123456")
|
||||
assert code.is_used
|
||||
|
||||
def test_verify_used_checked_before_attempts(self, service, mock_repo):
|
||||
"""已使用优先于其他检查."""
|
||||
code = _make_code(used=True, attempts=0)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, error = service.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "已使用" in error
|
||||
# attempts 会被 increment,但错误原因是已使用
|
||||
assert code.attempts == 1
|
||||
|
||||
def test_custom_max_attempts(self, mock_repo):
|
||||
"""自定义最大尝试次数."""
|
||||
svc = VerificationCodeService(repo=mock_repo, max_attempts=2)
|
||||
code = _make_code(attempts=2)
|
||||
mock_repo.find_latest.return_value = code
|
||||
|
||||
ok, error = svc.verify("u@e.com", CODE_TYPE_EMAIL_BIND, code.code)
|
||||
assert ok is False
|
||||
assert "验证次数过多" in error
|
||||
|
||||
|
||||
# ── validate_phone 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidatePhone:
|
||||
"""手机号格式校验测试"""
|
||||
|
||||
def test_valid_11_digit(self):
|
||||
"""标准11位手机号."""
|
||||
ok, msg = validate_phone("13800138000")
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
def test_valid_with_plus_86(self):
|
||||
"""带+86前缀."""
|
||||
ok, msg = validate_phone("+8613800138000")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_phone_short(self):
|
||||
"""太短的手机号"""
|
||||
ok, err = validate_phone("123")
|
||||
def test_invalid_too_short(self):
|
||||
"""位数不足."""
|
||||
ok, msg = validate_phone("1380013800")
|
||||
assert ok is False
|
||||
assert "格式不正确" in err
|
||||
assert "格式不正确" in msg
|
||||
|
||||
def test_invalid_phone_wrong_prefix(self):
|
||||
"""号段不对的手机号"""
|
||||
ok, err = validate_phone("11000000000")
|
||||
def test_invalid_too_long(self):
|
||||
"""位数过多."""
|
||||
ok, msg = validate_phone("138001380001")
|
||||
assert ok is False
|
||||
|
||||
def test_empty_phone(self):
|
||||
"""空手机号"""
|
||||
ok, err = validate_phone("")
|
||||
def test_invalid_starts_with_2(self):
|
||||
"""开头不是1."""
|
||||
ok, msg = validate_phone("23800138000")
|
||||
assert ok is False
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_phone_with_spaces(self):
|
||||
"""带空格的手机号会被 strip"""
|
||||
ok, _ = validate_phone(" 13800000001 ")
|
||||
def test_invalid_starts_with_12(self):
|
||||
"""第二位不在3-9."""
|
||||
ok, msg = validate_phone("12800138000")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_empty(self):
|
||||
"""空字符串."""
|
||||
ok, msg = validate_phone("")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_invalid_whitespace_only(self):
|
||||
"""仅空白."""
|
||||
ok, msg = validate_phone(" ")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_valid_all_prefixes_3_to_9(self):
|
||||
"""第二位3-9都有效."""
|
||||
for n in range(3, 10):
|
||||
ok, _ = validate_phone(f"1{n}800138000")
|
||||
assert ok is True, f"1{n} prefix should be valid"
|
||||
|
||||
def test_invalid_contains_letters(self):
|
||||
"""包含字母."""
|
||||
ok, msg = validate_phone("13800abc000")
|
||||
assert ok is False
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
"""前后空格会被清理."""
|
||||
ok, msg = validate_phone(" 13800138000 ")
|
||||
assert ok is True
|
||||
|
||||
|
||||
# ── normalize_phone 测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestNormalizePhone:
|
||||
"""normalize_phone 函数测试"""
|
||||
"""手机号标准化测试"""
|
||||
|
||||
def test_removes_plus86(self):
|
||||
"""去掉 +86 前缀"""
|
||||
assert normalize_phone("+8613800000001") == "13800000001"
|
||||
def test_strip_plus_86(self):
|
||||
"""去掉+86前缀."""
|
||||
assert normalize_phone("+8613800138000") == "13800138000"
|
||||
|
||||
def test_no_prefix_stays_same(self):
|
||||
"""没有前缀保持不变"""
|
||||
assert normalize_phone("13800000001") == "13800000001"
|
||||
"""无前缀保持不变."""
|
||||
assert normalize_phone("13800138000") == "13800138000"
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
"""去掉两端空白"""
|
||||
assert normalize_phone(" 13800000001 ") == "13800000001"
|
||||
"""清理前后空格."""
|
||||
assert normalize_phone(" 13800138000 ") == "13800138000"
|
||||
|
||||
def test_plus_86_with_spaces(self):
|
||||
"""带空格的+86."""
|
||||
assert normalize_phone(" +8613800138000 ") == "13800138000"
|
||||
|
||||
|
||||
# ── validate_email 测试 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidateEmail:
|
||||
"""validate_email 函数测试"""
|
||||
"""邮箱格式校验测试"""
|
||||
|
||||
def test_valid_email(self):
|
||||
"""有效邮箱"""
|
||||
ok, err = validate_email("test@example.com")
|
||||
def test_valid_simple(self):
|
||||
"""标准邮箱."""
|
||||
ok, msg = validate_email("user@example.com")
|
||||
assert ok is True
|
||||
assert err == ""
|
||||
assert msg == ""
|
||||
|
||||
def test_valid_email_with_subdomain(self):
|
||||
"""带子域名的邮箱"""
|
||||
ok, _ = validate_email("user@mail.example.com")
|
||||
def test_valid_with_dots(self):
|
||||
"""带点号的用户名."""
|
||||
ok, _ = validate_email("user.name@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_valid_email_with_plus(self):
|
||||
"""带 + 号的邮箱"""
|
||||
def test_valid_with_plus(self):
|
||||
"""带加号的邮箱."""
|
||||
ok, _ = validate_email("user+tag@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_email_no_at(self):
|
||||
"""没有 @ 的邮箱"""
|
||||
ok, err = validate_email("notanemail")
|
||||
assert ok is False
|
||||
assert "格式不正确" in err
|
||||
|
||||
def test_invalid_email_no_domain(self):
|
||||
"""没有域名的邮箱"""
|
||||
ok, err = validate_email("user@")
|
||||
assert ok is False
|
||||
|
||||
def test_empty_email(self):
|
||||
"""空邮箱"""
|
||||
ok, err = validate_email("")
|
||||
assert ok is False
|
||||
assert "不能为空" in err
|
||||
|
||||
def test_email_with_spaces(self):
|
||||
"""带空格的邮箱会被 strip"""
|
||||
ok, _ = validate_email(" test@example.com ")
|
||||
def test_valid_with_underscore(self):
|
||||
"""带下划线."""
|
||||
ok, _ = validate_email("user_name@example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_valid_subdomain(self):
|
||||
"""多级域名."""
|
||||
ok, _ = validate_email("user@mail.example.com")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_no_at(self):
|
||||
"""没有@."""
|
||||
ok, msg = validate_email("userexample.com")
|
||||
assert ok is False
|
||||
assert "格式不正确" in msg
|
||||
|
||||
def test_invalid_empty_local(self):
|
||||
"""@前为空."""
|
||||
ok, _ = validate_email("@example.com")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_empty_domain(self):
|
||||
"""@后为空."""
|
||||
ok, _ = validate_email("user@")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_no_tld(self):
|
||||
"""没有顶级域名."""
|
||||
ok, _ = validate_email("user@example")
|
||||
assert ok is False
|
||||
|
||||
def test_invalid_empty(self):
|
||||
"""空字符串."""
|
||||
ok, msg = validate_email("")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_invalid_spaces_only(self):
|
||||
"""仅空白."""
|
||||
ok, msg = validate_email(" ")
|
||||
assert ok is False
|
||||
assert "不能为空" in msg
|
||||
|
||||
def test_strips_whitespace(self):
|
||||
"""前后空格会被清理."""
|
||||
ok, msg = validate_email(" user@e.com ")
|
||||
assert ok is True
|
||||
|
||||
def test_invalid_special_chars(self):
|
||||
"""特殊字符."""
|
||||
ok, _ = validate_email("user name@e.com")
|
||||
assert ok is False
|
||||
|
||||
def test_valid_numbers(self):
|
||||
"""数字邮箱."""
|
||||
ok, _ = validate_email("12345@example.com")
|
||||
assert ok is True
|
||||
|
||||
|
||||
# ── VerificationCode 实体辅助验证 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVerificationCodeEntity:
|
||||
"""VerificationCode 实体属性测试"""
|
||||
|
||||
def test_is_expired_false_when_fresh(self):
|
||||
code = _make_code(ttl=300)
|
||||
assert code.is_expired is False
|
||||
|
||||
def test_is_expired_true_when_past(self):
|
||||
code = _make_code(ttl=-1)
|
||||
assert code.is_expired is True
|
||||
|
||||
def test_is_used_false_initially(self):
|
||||
code = _make_code()
|
||||
assert code.is_used is False
|
||||
|
||||
def test_is_used_after_mark_used(self):
|
||||
code = _make_code()
|
||||
code.mark_used()
|
||||
assert code.is_used is True
|
||||
assert code.used_at is not None
|
||||
|
||||
def test_is_valid_fresh(self):
|
||||
code = _make_code()
|
||||
assert code.is_valid is True
|
||||
|
||||
def test_is_valid_when_expired(self):
|
||||
code = _make_code(ttl=-100)
|
||||
assert code.is_valid is False
|
||||
|
||||
def test_is_valid_when_used(self):
|
||||
code = _make_code(used=True)
|
||||
assert code.is_valid is False
|
||||
|
||||
def test_increment_attempts(self):
|
||||
code = _make_code(attempts=0)
|
||||
code.increment_attempts()
|
||||
assert code.attempts == 1
|
||||
code.increment_attempts()
|
||||
assert code.attempts == 2
|
||||
|
||||
def test_create_generates_6_digit_code(self):
|
||||
code = VerificationCode.create("u@e.com", CODE_TYPE_EMAIL_BIND)
|
||||
assert len(code.code) == 6
|
||||
assert code.code.isdigit()
|
||||
|
||||
def test_create_custom_code(self):
|
||||
code = VerificationCode.create("u@e.com", CODE_TYPE_EMAIL_BIND, custom_code="555555")
|
||||
assert code.code == "555555"
|
||||
|
||||
def test_create_strips_recipient(self):
|
||||
code = VerificationCode.create(" u@e.com ", CODE_TYPE_EMAIL_BIND)
|
||||
assert code.recipient == "u@e.com"
|
||||
|
||||
def test_create_sets_expiry(self):
|
||||
code = VerificationCode.create("u@e.com", CODE_TYPE_EMAIL_BIND, ttl_seconds=120)
|
||||
delta = code.expires_at - code.created_at
|
||||
assert delta.total_seconds() == 120
|
||||
|
||||
Reference in New Issue
Block a user