From a81e76112cdbf8a83ec4cdb709647df4cb29caaa Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Thu, 30 Jul 2026 07:14:00 +0800 Subject: [PATCH] =?UTF-8?q?test(wave209):=20JWT=E6=9C=8D=E5=8A=A1=E4=B8=8E?= =?UTF-8?q?=E5=A4=84=E7=90=86=E5=99=A8=E5=8D=95=E6=B5=8B=E8=A1=A5=E5=85=A8?= =?UTF-8?q?=20+64=E6=B5=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 覆盖范围: - JWTConfig: 有效/空值/不安全密钥/自定义配置 - JWTService: access/refresh token创建与验证 - 过期token、篡改token、错误密钥验证 - 类型校验(access vs refresh) - JWTHandler委托层 - 全局handler配置与获取 64 test cases, 10 test classes --- tests/unit/test_jwt_service.py | 697 +++++++++++++++++++-------------- 1 file changed, 393 insertions(+), 304 deletions(-) diff --git a/tests/unit/test_jwt_service.py b/tests/unit/test_jwt_service.py index 27e3ae634..4c4d5c65b 100755 --- a/tests/unit/test_jwt_service.py +++ b/tests/unit/test_jwt_service.py @@ -1,63 +1,68 @@ -"""JWT 服务单元测试 — wave130.""" - -from __future__ import annotations +"""JWT 服务与处理器单元测试.""" import time from datetime import datetime, timedelta, timezone -import jwt as pyjwt +import jwt 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 # 满足长度要求的测试密钥 -TEST_SECRET = "test-secret-key-for-unit-testing-only-1234567890" -TEST_ALGORITHM = "HS256" - - -# ── JWTConfig 配置 ────────────────────────────────────────────────────────── +# ── JWTConfig 测试 ─────────────────────────────────────────────────────────── class TestJWTConfig: - def test_normal_config(self): - config = JWTConfig(secret_key=TEST_SECRET) - assert config.SECRET_KEY == TEST_SECRET + """JWTConfig 配置类测试""" + + 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.ACCESS_TOKEN_EXPIRE_MINUTES == 15 assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7 - def test_custom_config(self): + def test_init_custom_values(self): config = JWTConfig( - secret_key=TEST_SECRET, + secret_key=STRONG_SECRET, algorithm="HS384", 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.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): with pytest.raises(ValueError, match="secret_key must be provided"): JWTConfig(secret_key="") - def test_whitespace_secret_raises(self): - with pytest.raises(ValueError): + def test_whitespace_only_secret_raises(self): + with pytest.raises(ValueError, match="secret_key must be provided"): JWTConfig(secret_key=" ") def test_none_secret_raises(self): - with pytest.raises(ValueError): - JWTConfig(secret_key=None) # type: ignore + with pytest.raises(ValueError, match="secret_key must be provided"): + JWTConfig(secret_key=None) @pytest.mark.parametrize( - "bad_secret", + "insecure_secret", [ "your-secret-key-change-in-production", "your-secret-key", @@ -68,323 +73,407 @@ class TestJWTConfig: "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"): - 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: - def test_with_config_works(self): - config = JWTConfig(secret_key=TEST_SECRET) + """JWTService 初始化测试""" + + def test_init_with_config(self): + config = JWTConfig(secret_key=STRONG_SECRET) service = JWTService(config) assert service.config is config - def test_none_config_raises(self): - with pytest.raises(ValueError, match="JWTService requires"): + def test_init_none_config_raises(self): + with pytest.raises(ValueError, match="JWTService requires a JWTConfig"): JWTService(None) -# ── 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 常量 ────────────────────────────────────────────────────────── +# ── TokenType 测试 ─────────────────────────────────────────────────────────── class TestTokenType: + """TokenType 常量测试""" + def test_access_value(self): assert TokenType.ACCESS == "access" def test_refresh_value(self): assert TokenType.REFRESH == "refresh" - def test_different_types(self): + def test_access_and_refresh_different(self): assert TokenType.ACCESS != TokenType.REFRESH -# ── 多算法支持 ────────────────────────────────────────────────────────────── +# ── JWTService create_access_token 测试 ───────────────────────────────────── -class TestDifferentAlgorithms: - def test_hs384_works(self): - config = JWTConfig(secret_key=TEST_SECRET * 2, algorithm="HS384") +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) + 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") - payload = service.verify_token(token) - 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") + # 用 HS256 解码应该失败 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: - def setup_method(self): - self.service = JWTService(JWTConfig(secret_key=TEST_SECRET)) +class TestCreateRefreshToken: + """创建 refresh_token 测试""" - def test_empty_user_id(self): - token = self.service.create_access_token(user_id="") - payload = self.service.verify_access_token(token) - assert payload["sub"] == "" + @pytest.fixture + def service(self): + return JWTService(JWTConfig(secret_key=STRONG_SECRET)) - 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_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_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 + 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_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_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_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}" + 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) + 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