Compare commits

...

1 Commits

Author SHA1 Message Date
xiaoxia a81e76112c test(wave209): JWT服务与处理器单测补全 +64测
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 38s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m28s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m32s
Preview Deploy / Deploy Preview Environment (pull_request) Failing after 45s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m16s
AI Code Review / AI Code Review (pull_request) Successful in 2m35s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 3m33s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m54s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m1s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m53s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 7m9s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 7m2s
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Successful in 41s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Has been cancelled
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 46s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1159h52m21s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1159h55m0s
CI/CD Pipeline / Build Production Worker Image (pull_request) Failing after 1159h59m40s
CI/CD Pipeline / Build Production API Image (pull_request) Failing after 1159h59m44s
CI/CD Pipeline / ACR Image Cleanup (pull_request) Failing after 1160h5m20s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1160h5m24s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Failing after 1160h7m44s
CI/CD Pipeline / PR Build Web Image (pull_request) Failing after 1160h7m54s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Failing after 1160h8m4s
CI/CD Pipeline / Frontend Lint (pull_request) Failing after 1160h8m6s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Failing after 1160h11m14s
CI/CD Pipeline / Build Staging API Image (pull_request) Failing after 1160h11m18s
CI/CD Pipeline / Build Production Web Image (pull_request) Failing after 1160h32m47s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1160h38m27s
CI/CD Pipeline / Build Staging Web Image (pull_request) Failing after 1160h44m20s
覆盖范围:
- JWTConfig: 有效/空值/不安全密钥/自定义配置
- JWTService: access/refresh token创建与验证
- 过期token、篡改token、错误密钥验证
- 类型校验(access vs refresh)
- JWTHandler委托层
- 全局handler配置与获取

64 test cases, 10 test classes
2026-07-30 07:14:00 +08:00
+393 -304
View File
@@ -1,63 +1,68 @@
"""JWT 服务单元测试 — wave130.""" """JWT 服务与处理器单元测试."""
from __future__ import annotations
import time import time
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
import jwt as pyjwt import jwt
import pytest import pytest
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError 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 ( from packages.application.auth.jwt_service import (
JWTConfig, JWTConfig,
JWTService, JWTService,
TokenType, TokenType,
) )
# ── 测试常量 ──────────────────────────────────────────────────────────────── # ── 测试常量 ──────────────────────────────────────────────────────────────────
TEST_SECRET = "test-secret-key-for-unit-testing-only-not-for-production"
STRONG_SECRET = "x" * 32 # 满足长度要求的测试密钥
TEST_SECRET = "test-secret-key-for-unit-testing-only-1234567890" # ── JWTConfig 测试 ───────────────────────────────────────────────────────────
TEST_ALGORITHM = "HS256"
# ── JWTConfig 配置 ──────────────────────────────────────────────────────────
class TestJWTConfig: class TestJWTConfig:
def test_normal_config(self): """JWTConfig 配置类测试"""
config = JWTConfig(secret_key=TEST_SECRET)
assert config.SECRET_KEY == TEST_SECRET def test_init_with_valid_secret(self):
config = JWTConfig(secret_key=STRONG_SECRET)
assert config.SECRET_KEY == STRONG_SECRET
assert config.ALGORITHM == "HS256" assert config.ALGORITHM == "HS256"
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15 assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7 assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7
def test_custom_config(self): def test_init_custom_values(self):
config = JWTConfig( config = JWTConfig(
secret_key=TEST_SECRET, secret_key=STRONG_SECRET,
algorithm="HS384", algorithm="HS384",
access_token_expire_minutes=60, access_token_expire_minutes=60,
refresh_token_expire_days=30, refresh_token_expire_days=14,
) )
assert config.SECRET_KEY == STRONG_SECRET
assert config.ALGORITHM == "HS384" assert config.ALGORITHM == "HS384"
assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60 assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60
assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30 assert config.REFRESH_TOKEN_EXPIRE_DAYS == 14
def test_empty_secret_raises(self): def test_empty_secret_raises(self):
with pytest.raises(ValueError, match="secret_key must be provided"): with pytest.raises(ValueError, match="secret_key must be provided"):
JWTConfig(secret_key="") JWTConfig(secret_key="")
def test_whitespace_secret_raises(self): def test_whitespace_only_secret_raises(self):
with pytest.raises(ValueError): with pytest.raises(ValueError, match="secret_key must be provided"):
JWTConfig(secret_key=" ") JWTConfig(secret_key=" ")
def test_none_secret_raises(self): def test_none_secret_raises(self):
with pytest.raises(ValueError): with pytest.raises(ValueError, match="secret_key must be provided"):
JWTConfig(secret_key=None) # type: ignore JWTConfig(secret_key=None)
@pytest.mark.parametrize( @pytest.mark.parametrize(
"bad_secret", "insecure_secret",
[ [
"your-secret-key-change-in-production", "your-secret-key-change-in-production",
"your-secret-key", "your-secret-key",
@@ -68,323 +73,407 @@ class TestJWTConfig:
"Your-Secret-Key", "Your-Secret-Key",
], ],
) )
def test_insecure_defaults_rejected(self, bad_secret): def test_insecure_default_secret_raises(self, insecure_secret):
with pytest.raises(ValueError, match="insecure"): with pytest.raises(ValueError, match="insecure"):
JWTConfig(secret_key=bad_secret) 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
# ── JWTService 初始化 ────────────────────────────────────────────────────── # ── JWTService 初始化测试 ────────────────────────────────────────────────────
class TestJWTServiceInit: class TestJWTServiceInit:
def test_with_config_works(self): """JWTService 初始化测试"""
config = JWTConfig(secret_key=TEST_SECRET)
def test_init_with_config(self):
config = JWTConfig(secret_key=STRONG_SECRET)
service = JWTService(config) service = JWTService(config)
assert service.config is config assert service.config is config
def test_none_config_raises(self): def test_init_none_config_raises(self):
with pytest.raises(ValueError, match="JWTService requires"): with pytest.raises(ValueError, match="JWTService requires a JWTConfig"):
JWTService(None) JWTService(None)
# ── create_access_token ──────────────────────────────────────────────────── # ── TokenType 测试 ───────────────────────────────────────────────────────────
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: class TestTokenType:
"""TokenType 常量测试"""
def test_access_value(self): def test_access_value(self):
assert TokenType.ACCESS == "access" assert TokenType.ACCESS == "access"
def test_refresh_value(self): def test_refresh_value(self):
assert TokenType.REFRESH == "refresh" assert TokenType.REFRESH == "refresh"
def test_different_types(self): def test_access_and_refresh_different(self):
assert TokenType.ACCESS != TokenType.REFRESH assert TokenType.ACCESS != TokenType.REFRESH
# ── 多算法支持 ────────────────────────────────────────────────────────────── # ── JWTService create_access_token 测试 ─────────────────────────────────────
class TestDifferentAlgorithms: class TestCreateAccessToken:
def test_hs384_works(self): """创建 access_token 测试"""
config = JWTConfig(secret_key=TEST_SECRET * 2, algorithm="HS384")
@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)
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) service = JWTService(config)
token = service.create_access_token(user_id="u1") token = service.create_access_token(user_id="u1")
payload = service.verify_token(token) # 用 HS256 解码应该失败
assert payload["sub"] == "u1"
def test_hs512_works(self):
config = JWTConfig(secret_key=TEST_SECRET * 3, algorithm="HS512")
service = JWTService(config)
token = service.create_access_token(user_id="u1")
payload = service.verify_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)
token = service_256.create_access_token(user_id="u1")
with pytest.raises(InvalidTokenError): with pytest.raises(InvalidTokenError):
service_384.verify_token(token) jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
# 用 HS384 解码应该成功
payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS384"])
assert payload["sub"] == "u1"
# ── 边界:空用户ID等 ──────────────────────────────────────────────────────── # ── JWTService create_refresh_token 测试 ────────────────────────────────────
class TestEdgeCases: class TestCreateRefreshToken:
def setup_method(self): """创建 refresh_token 测试"""
self.service = JWTService(JWTConfig(secret_key=TEST_SECRET))
def test_empty_user_id(self): @pytest.fixture
token = self.service.create_access_token(user_id="") def service(self):
payload = self.service.verify_access_token(token) return JWTService(JWTConfig(secret_key=STRONG_SECRET))
assert payload["sub"] == ""
def test_long_user_id(self): def test_creates_valid_string(self, service):
long_id = "x" * 1000 token = service.create_refresh_token(user_id="u1", session_id="s1")
token = self.service.create_access_token(user_id=long_id) assert isinstance(token, str)
payload = self.service.verify_access_token(token) assert len(token) > 0
assert payload["sub"] == long_id
def test_special_chars_in_user_id(self): def test_contains_user_id_and_session_id(self, service):
uid = "user@#$%^&*()_+-=[]{}|;:',.<>?/`~" token = service.create_refresh_token(user_id="u1", session_id="sess-abc")
token = self.service.create_access_token(user_id=uid) payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
payload = self.service.verify_access_token(token) assert payload["sub"] == "u1"
assert payload["sub"] == uid assert payload["session_id"] == "sess-abc"
def test_unicode_user_id(self): def test_token_type_is_refresh(self, service):
uid = "用户_测试_123_🎉" token = service.create_refresh_token(user_id="u1", session_id="s1")
token = self.service.create_access_token(user_id=uid) payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
payload = self.service.verify_access_token(token) assert payload["type"] == TokenType.REFRESH
assert payload["sub"] == uid
def test_many_additional_claims(self): def test_has_iat_and_exp(self, service):
claims = {f"key_{i}": f"value_{i}" for i in range(50)} token = service.create_refresh_token(user_id="u1", session_id="s1")
token = self.service.create_access_token(user_id="u1", additional_claims=claims) payload = jwt.decode(token, STRONG_SECRET, algorithms=["HS256"])
payload = self.service.verify_access_token(token) assert "iat" in payload
for i in range(50): assert "exp" in payload
assert payload[f"key_{i}"] == f"value_{i}" 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)
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)
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_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
# ── 全局 JWT handler 测试 ───────────────────────────────────────────────────
class TestGlobalJWTHandler:
"""全局 JWT Handler 配置与获取测试"""
def test_configure_creates_handler(self):
handler = configure_jwt_handler(secret_key=STRONG_SECRET)
assert isinstance(handler, JWTHandler)
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_get_before_configure_raises(self):
# 重置全局状态(通过设置 None 模拟未配置)
import packages.application.auth.jwt_handler as mod
mod._default_handler = None
with pytest.raises(RuntimeError, match="JWT handler not configured"):
get_jwt_handler()
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