Files
xiaoxia-saas/packages/application/auth/jwt_service.py
T
CI Bot 9c6c477f55
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 2m22s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 2m24s
CI/CD Pipeline / Integration Tests (pull_request) Failing after 37s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 4m3s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
fix(backend): Phase 1 后端代码清理与修复
P0 关键修复:
- P0-1: 注册接口添加 RateLimitMiddleware 限流保护
- P0-3: /metrics 端点添加 JWT 认证(移除匿名访问)
- P0-4: 修复 Celery 任务名冲突(generation_task vs generate_video)
- P1-5: JWT logout token 黑名单机制

P1 修复:
- P1-1: forgot_password 硬编码 localhost → 使用 settings.APP_BASE_URL
- P1-2: generation.py 直接创建 DB 连接 → 使用依赖注入
- P1-6: Image.open() 未关闭 → 统一使用 with 语句
- P1-7: 订阅续费事务修复

P2 代码质量:
- P2-1: 修复 EditingMode 枚举重复定义 → 统一引用 shared 包
- P2-2: 修复 SMTP_FRON_NAME → SMTP_FROM_NAME 拼写
- P2-3: UserModel subscription_quota 类型统一为 float
- P2-4: .env.production DATABASE_MAX_OVERFLOW 30 → 10
- 清理 15 处 except:pass(保留 2 处有注释说明的)
- 禁用 SVG 上传(XSS 风险)
- 删除 decode_token_unsafe() 不安全函数
- 简化 /ready 端点
- 删除 8 处死代码、10 个空文件/模块
- 合并 3 对 100% 重复函数
- 对齐 6 个废弃环境变量

v2 修复(代码审查后):
- 修复密码重置路由路径: /password/forgot → /forgot-password,
  /password/reset → /reset-password(与前端 API 对齐)
- 合并 _check_project_access: asset_libraries.py 和 edit_plans.py
  中的重复函数统一到 _helpers.py(含空字符串守卫 + 中文错误信息)
- 顺手修复: HTTPException 统一从 fastapi 导入(替换 starlette 导入)
- OSS_ENDPOINT 拼写修复拆分为单独 PR,本 PR 不包含
2026-07-13 13:50:52 +08:00

225 lines
6.4 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 Any, Dict, Optional
import jwt
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
class JWTConfig:
"""JWT 配置"""
def __init__(
self,
secret_key: str,
algorithm: str = "HS256",
access_token_expire_minutes: int = 15,
refresh_token_expire_days: int = 7,
):
"""
初始化 JWT 配置
Args:
secret_key: JWT 签名密钥(必须从环境变量或配置注入,不允许默认值)
algorithm: 加密算法,默认 HS256
access_token_expire_minutes: Access Token 过期时间(分钟)
refresh_token_expire_days: Refresh Token 过期时间(天)
Raises:
ValueError: 如果 secret_key 为空或包含不安全默认值
"""
if not secret_key or secret_key.strip() == "":
raise ValueError("JWT secret_key must be provided and cannot be empty") # noqa: E501
insecure_defaults = [
"your-secret-key-change-in-production",
"your-secret-key",
"secret",
"changeme",
"password",
]
if secret_key.lower() in [d.lower() for d in insecure_defaults]:
raise ValueError( # noqa: E501
f"JWT secret_key '{secret_key}' is insecure. " "Please provide a strong random secret."
)
self.SECRET_KEY: str = secret_key
self.ALGORITHM: str = algorithm
self.ACCESS_TOKEN_EXPIRE_MINUTES: int = access_token_expire_minutes
self.REFRESH_TOKEN_EXPIRE_DAYS: int = refresh_token_expire_days
class TokenType:
"""Token 类型"""
ACCESS = "access"
REFRESH = "refresh"
class JWTService:
"""JWT 服务类"""
def __init__(self, config: JWTConfig = None):
if config is None:
raise ValueError( # noqa: E501
"JWTService requires a JWTConfig instance. " # noqa: E501
"Please provide a configured JWTConfig with a valid " # noqa: E501
"secret_key."
)
self.config = config
def create_access_token(
self,
user_id: str,
role: str = "",
additional_claims: Optional[Dict[str, Any]] = None,
) -> str:
"""
创建 access_token
Args:
user_id: 用户 ID
role: 用户角色(admin/user/guest
additional_claims: 额外的声明信息
Returns:
JWT Token 字符串
"""
now = datetime.utcnow()
expire = now + timedelta(minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES) # noqa: E501
payload = {
"sub": user_id, # subject (用户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( # noqa: E501
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(如果解码失败返回 None
"""
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
# 全局实例(生产环境必须从配置读取有效的 secret_key)
# jwt_service = JWTService() # 不再允许无参数实例化
# Lazy singleton - created with settings on first access
_jwt_service_instance = None
def _get_jwt_service():
global _jwt_service_instance
if _jwt_service_instance is None:
from app.config import settings
kw = dict(secret_key=settings.JWT_SECRET_KEY)
if hasattr(settings, "JWT_ALGORITHM"):
kw["algorithm"] = settings.JWT_ALGORITHM
if hasattr(settings, "JWT_ACCESS_TOKEN_EXPIRE_MINUTES"):
kw["access_token_expire_minutes"] = settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES
if hasattr(settings, "JWT_REFRESH_TOKEN_EXPIRE_DAYS"):
kw["refresh_token_expire_days"] = settings.JWT_REFRESH_TOKEN_EXPIRE_DAYS
_jwt_service_instance = JWTService(JWTConfig(**kw))
return _jwt_service_instance
class _JWTServiceProxy:
def __getattr__(self, name):
return getattr(_get_jwt_service(), name)
jwt_service = _JWTServiceProxy()