diff --git a/tests/unit/test_jwt_service.py b/tests/unit/test_jwt_service.py index 7d2d4eaf2..27e3ae634 100755 --- a/tests/unit/test_jwt_service.py +++ b/tests/unit/test_jwt_service.py @@ -1,264 +1,390 @@ -"""JWT 服务单元测试.""" +"""JWT 服务单元测试 — wave130.""" from __future__ import annotations import time +from datetime import datetime, timedelta, timezone +import jwt as pyjwt import pytest from jwt.exceptions import ExpiredSignatureError, InvalidTokenError -from packages.application.auth.jwt_service import JWTConfig, JWTService, TokenType +from packages.application.auth.jwt_service import ( + JWTConfig, + JWTService, + TokenType, +) + +# ── 测试常量 ──────────────────────────────────────────────────────────────── -@pytest.fixture -def jwt_config(): - return JWTConfig( - secret_key="test-secret-key-strong-enough-123456", - algorithm="HS256", - access_token_expire_minutes=30, - refresh_token_expire_days=7, - ) +TEST_SECRET = "test-secret-key-for-unit-testing-only-1234567890" +TEST_ALGORITHM = "HS256" -@pytest.fixture -def jwt_service(jwt_config): - return JWTService(jwt_config) +# ── JWTConfig 配置 ────────────────────────────────────────────────────────── class TestJWTConfig: - """JWTConfig 测试""" + 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_custom_config(self): + config = JWTConfig( + secret_key=TEST_SECRET, + algorithm="HS384", + access_token_expire_minutes=60, + refresh_token_expire_days=30, + ) + assert config.ALGORITHM == "HS384" + assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60 + assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30 def test_empty_secret_raises(self): - """空 secret_key 抛出 ValueError""" - with pytest.raises(ValueError, match="must be provided"): + with pytest.raises(ValueError, match="secret_key must be provided"): JWTConfig(secret_key="") def test_whitespace_secret_raises(self): - """纯空白 secret_key 抛出 ValueError""" - with pytest.raises(ValueError, match="must be provided"): + with pytest.raises(ValueError): JWTConfig(secret_key=" ") - def test_insecure_default_secret_raises(self): - """不安全的默认 secret 抛出 ValueError""" - insecure_secrets = [ + def test_none_secret_raises(self): + with pytest.raises(ValueError): + JWTConfig(secret_key=None) # type: ignore + + @pytest.mark.parametrize( + "bad_secret", + [ "your-secret-key-change-in-production", "your-secret-key", "secret", "changeme", "password", "SECRET", - ] - for secret in insecure_secrets: - with pytest.raises(ValueError, match="insecure"): - JWTConfig(secret_key=secret) - - def test_strong_secret_accepted(self): - """强 secret 可以正常创建""" - config = JWTConfig(secret_key="my-strong-secret-key-1234567890") - assert config.SECRET_KEY == "my-strong-secret-key-1234567890" - - def test_default_values(self): - """默认配置值正确""" - config = JWTConfig(secret_key="test-secret-12345") - assert config.ALGORITHM == "HS256" - assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 15 - assert config.REFRESH_TOKEN_EXPIRE_DAYS == 7 - - def test_custom_expiry_values(self): - """自定义过期时间""" - config = JWTConfig( - secret_key="test-secret-12345", - access_token_expire_minutes=60, - refresh_token_expire_days=30, - ) - assert config.ACCESS_TOKEN_EXPIRE_MINUTES == 60 - assert config.REFRESH_TOKEN_EXPIRE_DAYS == 30 + "Your-Secret-Key", + ], + ) + def test_insecure_defaults_rejected(self, bad_secret): + with pytest.raises(ValueError, match="insecure"): + JWTConfig(secret_key=bad_secret) -class TestTokenType: - """TokenType 测试""" - - def test_access_token_type(self): - """access token 类型值""" - assert TokenType.ACCESS == "access" - - def test_refresh_token_type(self): - """refresh token 类型值""" - assert TokenType.REFRESH == "refresh" +# ── JWTService 初始化 ────────────────────────────────────────────────────── class TestJWTServiceInit: - """JWTService 初始化测试""" + def test_with_config_works(self): + config = JWTConfig(secret_key=TEST_SECRET) + service = JWTService(config) + assert service.config is config def test_none_config_raises(self): - """不传 config 抛出 ValueError""" - with pytest.raises(ValueError, match="requires a JWTConfig"): + with pytest.raises(ValueError, match="JWTService requires"): JWTService(None) - def test_with_config_creates_service(self, jwt_config): - """传入 config 正常创建""" - service = JWTService(jwt_config) - assert service.config is jwt_config + +# ── create_access_token ──────────────────────────────────────────────────── class TestCreateAccessToken: - """create_access_token 测试""" + def setup_method(self): + self.service = JWTService(JWTConfig(secret_key=TEST_SECRET)) - def test_returns_string(self, jwt_service): - """返回非空字符串""" - token = jwt_service.create_access_token(user_id="user_001") + 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_contains_user_id(self, jwt_service): - """payload 包含正确的 user_id(sub字段)""" - token = jwt_service.create_access_token(user_id="user_123") - payload = jwt_service.verify_token(token) - assert payload["sub"] == "user_123" + 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_contains_role(self, jwt_service): - """payload 包含 role""" - token = jwt_service.create_access_token(user_id="user_001", role="admin") - payload = jwt_service.verify_token(token) + 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, jwt_service): - """不传 role 默认为空字符串""" - token = jwt_service.create_access_token(user_id="user_001") - payload = jwt_service.verify_token(token) + 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, jwt_service): - """access token 的 type 为 access""" - token = jwt_service.create_access_token(user_id="user_001") - payload = jwt_service.verify_token(token) + 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_additional_claims(self, jwt_service): - """额外声明被包含在 payload 中""" - token = jwt_service.create_access_token( - user_id="user_001", - additional_claims={"email": "test@example.com", "tenant": "t1"}, - ) - payload = jwt_service.verify_token(token) - assert payload["email"] == "test@example.com" - assert payload["tenant"] == "t1" - - def test_has_iat_and_exp(self, jwt_service): - """payload 包含 iat 和 exp""" - token = jwt_service.create_access_token(user_id="user_001") - payload = jwt_service.verify_token(token) + 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_expiry_correct_duration(self, jwt_service): - """过期时间设置正确""" - token = jwt_service.create_access_token(user_id="user_001") - payload = jwt_service.verify_token(token) - # 30分钟 = 1800秒 - duration = payload["exp"] - payload["iat"] - assert 1790 <= duration <= 1810 # 允许10秒误差 + 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: - """create_refresh_token 测试""" + def setup_method(self): + self.service = JWTService(JWTConfig(secret_key=TEST_SECRET)) - def test_returns_string(self, jwt_service): - """返回非空字符串""" - token = jwt_service.create_refresh_token(user_id="user_001", session_id="sess_001") + def test_creates_valid_token(self): + token = self.service.create_refresh_token(user_id="u1", session_id="sess_001") assert isinstance(token, str) - assert len(token) > 0 + assert token.count(".") == 2 - def test_contains_user_and_session(self, jwt_service): - """包含 user_id 和 session_id""" - token = jwt_service.create_refresh_token(user_id="user_123", session_id="sess_456") - payload = jwt_service.verify_token(token) - assert payload["sub"] == "user_123" - assert payload["session_id"] == "sess_456" + 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, jwt_service): - """refresh token 的 type 为 refresh""" - token = jwt_service.create_refresh_token(user_id="user_001", session_id="s1") - payload = jwt_service.verify_token(token) + 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: - """verify_token 测试""" + def setup_method(self): + self.service = JWTService(JWTConfig(secret_key=TEST_SECRET)) - def test_valid_token(self, jwt_service): - """有效 token 验证通过""" - token = jwt_service.create_access_token(user_id="u1") - payload = jwt_service.verify_token(token) + 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_invalid_token_raises(self, jwt_service): - """无效 token 抛出 InvalidTokenError""" - with pytest.raises(InvalidTokenError): - jwt_service.verify_token("not.a.valid.token") - - def test_empty_token_raises(self, jwt_service): - """空字符串 token 抛出异常""" - with pytest.raises(InvalidTokenError): - jwt_service.verify_token("") - - def test_wrong_secret_fails(self, jwt_config): - """不同密钥的 token 无法验证""" - service1 = JWTService(JWTConfig(secret_key="secret-one-123456")) - service2 = JWTService(JWTConfig(secret_key="secret-two-1234567")) - - token = service1.create_access_token(user_id="u1") - with pytest.raises(InvalidTokenError): - service2.verify_token(token) - - -class TestVerifyAccessToken: - """verify_access_token 测试""" - - def test_valid_access_token(self, jwt_service): - """有效 access token 验证通过""" - token = jwt_service.create_access_token(user_id="u1") - payload = jwt_service.verify_access_token(token) - assert payload["sub"] == "u1" - - def test_refresh_token_fails(self, jwt_service): - """refresh token 不能当 access token 用""" - token = jwt_service.create_refresh_token(user_id="u1", session_id="s1") - with pytest.raises(ValueError, match="Token type must be 'access'"): - jwt_service.verify_access_token(token) - - -class TestVerifyRefreshToken: - """verify_refresh_token 测试""" - - def test_valid_refresh_token(self, jwt_service): - """有效 refresh token 验证通过""" - token = jwt_service.create_refresh_token(user_id="u1", session_id="s1") - payload = jwt_service.verify_refresh_token(token) - assert payload["sub"] == "u1" - assert payload["session_id"] == "s1" - - def test_access_token_fails(self, jwt_service): - """access token 不能当 refresh token 用""" - token = jwt_service.create_access_token(user_id="u1") - with pytest.raises(ValueError, match="Token type must be 'refresh'"): - jwt_service.verify_refresh_token(token) - - -class TestExpiredToken: - """过期 token 测试""" - - def test_expired_access_token_raises(self): - """过期 token 验证抛出 ExpiredSignatureError""" - config = JWTConfig( - secret_key="test-secret-12345", - access_token_expire_minutes=-1, # 立即过期 - ) + 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") - time.sleep(0.1) + # 等待过期(用 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): - service.verify_access_token(token) + 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: + def test_access_value(self): + assert TokenType.ACCESS == "access" + + def test_refresh_value(self): + assert TokenType.REFRESH == "refresh" + + def test_different_types(self): + assert TokenType.ACCESS != TokenType.REFRESH + + +# ── 多算法支持 ────────────────────────────────────────────────────────────── + + +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="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") + with pytest.raises(InvalidTokenError): + service_384.verify_token(token) + + +# ── 边界:空用户ID等 ──────────────────────────────────────────────────────── + + +class TestEdgeCases: + def setup_method(self): + self.service = JWTService(JWTConfig(secret_key=TEST_SECRET)) + + 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_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_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_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_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}"