Files
xiaoxia-saas/packages/domain/auth/jwt_service.py
T
Xiaoxia AI e344fe2e9e
Deploy / Deploy Staging (push) Failing after 5s
Deploy / Deploy Production (push) Has been skipped
Tests / test (push) Failing after 8s
Tests / lint (push) Failing after 7s
feat(auth): add JWT service with sign/verify/refresh functionality
- Implement JWTService class with access_token and refresh_token support
- Add token type validation (access vs refresh)
- Add comprehensive unit tests (9 tests all passed)
- Install PyJWT dependency

Phase 4 Task 1/68 completed
2026-06-17 01:01:22 +08:00

195 lines
5.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
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()