180 lines
4.6 KiB
Python
180 lines
4.6 KiB
Python
"""
|
||
JWT 工具类
|
||
提供 Token 签发、验证、刷新功能
|
||
"""
|
||
|
||
from datetime import datetime, timedelta
|
||
from typing import Any, Dict, 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()
|