diff --git a/packages/domain/auth/__init__.py b/packages/domain/auth/__init__.py new file mode 100644 index 000000000..dd940765c --- /dev/null +++ b/packages/domain/auth/__init__.py @@ -0,0 +1,9 @@ +"""认证模块""" +from packages.domain.auth.jwt_service import JWTService, JWTConfig, TokenType, jwt_service + +__all__ = [ + "JWTService", + "JWTConfig", + "TokenType", + "jwt_service", +] diff --git a/packages/domain/auth/jwt_service.py b/packages/domain/auth/jwt_service.py new file mode 100644 index 000000000..8897391f1 --- /dev/null +++ b/packages/domain/auth/jwt_service.py @@ -0,0 +1,194 @@ +""" +JWT 工具类 +提供 Token 签发、验证、刷新功能 +""" +from datetime import datetime, timedelta +from typing import Dict, Any, Optional +import jwt +from jwt.exceptions import ExpiredSignatureError, InvalidTokenError + + +class JWTConfig: + """JWT 配置""" + # 从环境变量读取,这里先用默认值 + SECRET_KEY: str = "your-secret-key-change-in-production" + ALGORITHM: str = "HS256" + ACCESS_TOKEN_EXPIRE_MINUTES: int = 30 # 30 分钟 + REFRESH_TOKEN_EXPIRE_DAYS: int = 30 # 30 天 + + +class TokenType: + """Token 类型""" + ACCESS = "access" + REFRESH = "refresh" + + +class JWTService: + """JWT 服务类""" + + def __init__(self, config: JWTConfig = None): + self.config = config or JWTConfig() + + def create_access_token( + self, + user_id: str, + workspace_id: str, + role: str, + additional_claims: Optional[Dict[str, Any]] = None + ) -> str: + """ + 创建 access_token + + Args: + user_id: 用户 ID + workspace_id: 工作空间 ID + role: 用户在该工作空间的角色 + additional_claims: 额外的声明(可选) + + Returns: + JWT Token 字符串 + """ + now = datetime.utcnow() + expire = now + timedelta(minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES) + + payload = { + "sub": user_id, # subject (用户 ID) + "workspace_id": workspace_id, + "role": role, + "type": TokenType.ACCESS, + "iat": now, # issued at + "exp": expire, # expiration time + } + + if additional_claims: + payload.update(additional_claims) + + return jwt.encode( + payload, + self.config.SECRET_KEY, + algorithm=self.config.ALGORITHM + ) + + def create_refresh_token( + self, + user_id: str, + session_id: str + ) -> str: + """ + 创建 refresh_token + + Args: + user_id: 用户 ID + session_id: Session ID(用于撤销) + + Returns: + JWT Token 字符串 + """ + now = datetime.utcnow() + expire = now + timedelta(days=self.config.REFRESH_TOKEN_EXPIRE_DAYS) + + payload = { + "sub": user_id, + "session_id": session_id, + "type": TokenType.REFRESH, + "iat": now, + "exp": expire, + } + + return jwt.encode( + payload, + self.config.SECRET_KEY, + algorithm=self.config.ALGORITHM + ) + + def verify_token(self, token: str) -> Dict[str, Any]: + """ + 验证 Token 并解码 + + Args: + token: JWT Token 字符串 + + Returns: + Token payload + + Raises: + ExpiredSignatureError: Token 已过期 + InvalidTokenError: Token 无效 + """ + try: + payload = jwt.decode( + token, + self.config.SECRET_KEY, + algorithms=[self.config.ALGORITHM] + ) + return payload + except ExpiredSignatureError: + raise ExpiredSignatureError("Token has expired") + except InvalidTokenError as e: + raise InvalidTokenError(f"Invalid token: {str(e)}") + + def verify_access_token(self, token: str) -> Dict[str, Any]: + """ + 验证 access_token + + Args: + token: JWT Token 字符串 + + Returns: + Token payload + + Raises: + ValueError: Token 类型不是 access + ExpiredSignatureError: Token 已过期 + InvalidTokenError: Token 无效 + """ + payload = self.verify_token(token) + + if payload.get("type") != TokenType.ACCESS: + raise ValueError("Token type must be 'access'") + + return payload + + def verify_refresh_token(self, token: str) -> Dict[str, Any]: + """ + 验证 refresh_token + + Args: + token: JWT Token 字符串 + + Returns: + Token payload + + Raises: + ValueError: Token 类型不是 refresh + ExpiredSignatureError: Token 已过期 + InvalidTokenError: Token 无效 + """ + payload = self.verify_token(token) + + if payload.get("type") != TokenType.REFRESH: + raise ValueError("Token type must be 'refresh'") + + return payload + + def decode_token_unsafe(self, token: str) -> Optional[Dict[str, Any]]: + """ + 不验证签名地解码 Token(仅用于调试/日志) + + Args: + token: JWT Token 字符串 + + Returns: + Token payload(如果解码失败返回 None) + """ + try: + return jwt.decode( + token, + options={"verify_signature": False} + ) + except Exception: + return None + + +# 全局实例(生产环境应该从配置读取) +jwt_service = JWTService() diff --git a/requirements.txt b/requirements.txt index 39df8b3c0..d2c383c8f 100644 Binary files a/requirements.txt and b/requirements.txt differ diff --git a/tests/unit/test_jwt_service.py b/tests/unit/test_jwt_service.py new file mode 100644 index 000000000..6a88fe263 --- /dev/null +++ b/tests/unit/test_jwt_service.py @@ -0,0 +1,164 @@ +""" +JWT 工具类测试 +""" +import pytest +from datetime import datetime, timedelta +from jwt.exceptions import ExpiredSignatureError, InvalidTokenError + +from packages.domain.auth.jwt_service import ( + JWTService, + JWTConfig, + TokenType, +) + + +class TestJWTService: + """JWT 服务测试""" + + @pytest.fixture + def jwt_service(self): + """创建 JWT 服务实例""" + config = JWTConfig() + config.SECRET_KEY = "test-secret-key-for-testing" + return JWTService(config) + + def test_create_access_token(self, jwt_service): + """测试创建 access_token""" + token = jwt_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="admin" + ) + + assert isinstance(token, str) + assert len(token) > 0 + + # 验证 Token 内容 + payload = jwt_service.verify_access_token(token) + assert payload["sub"] == "user-123" + assert payload["workspace_id"] == "workspace-456" + assert payload["role"] == "admin" + assert payload["type"] == TokenType.ACCESS + + def test_create_refresh_token(self, jwt_service): + """测试创建 refresh_token""" + token = jwt_service.create_refresh_token( + user_id="user-123", + session_id="session-789" + ) + + assert isinstance(token, str) + assert len(token) > 0 + + # 验证 Token 内容 + payload = jwt_service.verify_refresh_token(token) + assert payload["sub"] == "user-123" + assert payload["session_id"] == "session-789" + assert payload["type"] == TokenType.REFRESH + + def test_verify_valid_access_token(self, jwt_service): + """测试验证有效的 access_token""" + token = jwt_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="member" + ) + + payload = jwt_service.verify_access_token(token) + assert payload["sub"] == "user-123" + assert payload["workspace_id"] == "workspace-456" + assert payload["role"] == "member" + + def test_verify_expired_token(self, jwt_service): + """测试验证过期的 Token""" + # 创建一个已过期的配置(使用相同的 SECRET_KEY) + config = JWTConfig() + config.SECRET_KEY = "test-secret-key-for-testing" # 与 fixture 相同 + config.ACCESS_TOKEN_EXPIRE_MINUTES = -1 # 负数,立即过期 + + expired_service = JWTService(config) + token = expired_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="admin" + ) + + # 验证应该抛出过期异常 + with pytest.raises(ExpiredSignatureError): + jwt_service.verify_access_token(token) + + def test_verify_invalid_token(self, jwt_service): + """测试验证无效的 Token""" + invalid_token = "invalid.token.string" + + with pytest.raises(InvalidTokenError): + jwt_service.verify_access_token(invalid_token) + + def test_verify_wrong_token_type(self, jwt_service): + """测试验证错误类型的 Token""" + # 创建 refresh_token + refresh_token = jwt_service.create_refresh_token( + user_id="user-123", + session_id="session-789" + ) + + # 用 verify_access_token 验证应该失败 + with pytest.raises(ValueError, match="Token type must be 'access'"): + jwt_service.verify_access_token(refresh_token) + + # 反过来也一样 + access_token = jwt_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="admin" + ) + + with pytest.raises(ValueError, match="Token type must be 'refresh'"): + jwt_service.verify_refresh_token(access_token) + + def test_verify_tampered_token(self, jwt_service): + """测试验证被篡改的 Token""" + token = jwt_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="admin" + ) + + # 篡改 Token(修改最后几个字符) + tampered_token = token[:-5] + "XXXXX" + + with pytest.raises(InvalidTokenError): + jwt_service.verify_access_token(tampered_token) + + def test_additional_claims(self, jwt_service): + """测试额外的声明""" + token = jwt_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="admin", + additional_claims={ + "email": "user@example.com", + "display_name": "Test User" + } + ) + + payload = jwt_service.verify_access_token(token) + assert payload["email"] == "user@example.com" + assert payload["display_name"] == "Test User" + + def test_decode_unsafe(self, jwt_service): + """测试不安全解码(不验证签名)""" + token = jwt_service.create_access_token( + user_id="user-123", + workspace_id="workspace-456", + role="admin" + ) + + # 不验证签名地解码 + payload = jwt_service.decode_token_unsafe(token) + assert payload is not None + assert payload["sub"] == "user-123" + + # 无效 Token 应该返回 None + invalid_payload = jwt_service.decode_token_unsafe("invalid.token") + assert invalid_payload is None