Files
2026-06-08 18:14:59 +08:00

322 lines
9.3 KiB
Python
Raw Permalink 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.
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
Security Utils - JWT Token工具
用于生成和验证JWT Token
"""
from datetime import datetime, timedelta, timezone
from typing import Optional, Any
from jose import jwt, JWTError
from fastapi import Depends, HTTPException, status, Request
from fastapi.security import OAuth2PasswordBearer
from sqlalchemy.ext.asyncio import AsyncSession
import bcrypt
from app.config import settings
from app.database import get_db
# bcrypt 只处理前 72 字节;与 passlib 默认行为一致
_BCRYPT_ROUNDS = 12
def hash_password(password: str) -> str:
"""使用 bcrypt 生成密码哈希(兼容既有 passlib 产出的 $2b$ 格式)"""
password_bytes = password.encode("utf-8")[:72]
return bcrypt.hashpw(password_bytes, bcrypt.gensalt(rounds=_BCRYPT_ROUNDS)).decode("utf-8")
def verify_password(plain_password: str, hashed_password: str) -> bool:
"""校验明文密码与 bcrypt 哈希是否匹配"""
if not plain_password or not hashed_password:
return False
try:
return bcrypt.checkpw(
plain_password.encode("utf-8")[:72],
hashed_password.encode("utf-8"),
)
except (ValueError, TypeError):
return False
# OAuth2密码流,指定token获取地址(auto_error=False让中间件处理认证)
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/v1/core/auth/login/oauth2", auto_error=False)
def create_access_token(data: dict, expires_delta: Optional[timedelta] = None, device_id: Optional[str] = None) -> str:
"""
创建Access Token
:param data: 要编码的数据
:param expires_delta: 过期时间增量
:param device_id: 设备唯一标识(可选)
:return: JWT Token字符串
"""
to_encode = data.copy()
if expires_delta:
expire = datetime.now(timezone.utc) + expires_delta
else:
expire = datetime.now(timezone.utc) + timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
to_encode.update({
"exp": expire,
"type": "access"
})
# 如果提供了 device_id,添加到 token 中
if device_id:
to_encode["device_id"] = device_id
encoded_jwt = jwt.encode(to_encode, settings.JWT_SECRET_KEY, algorithm=settings.JWT_ALGORITHM)
return encoded_jwt
def create_refresh_token(data: dict, expires_delta: Optional[timedelta] = None, device_id: Optional[str] = None) -> str:
"""
创建Refresh Token
:param data: 要编码的数据
:param expires_delta: 过期时间增量
:param device_id: 设备标识(用于多设备登录)
:return: JWT Token字符串
"""
to_encode = data.copy()
if expires_delta:
expire = datetime.now(timezone.utc) + expires_delta
else:
expire = datetime.now(timezone.utc) + timedelta(days=settings.REFRESH_TOKEN_EXPIRE_DAYS)
to_encode.update({
"exp": expire,
"type": "refresh"
})
# 如果提供了 device_id,添加到 token 中
if device_id:
to_encode["device_id"] = device_id
encoded_jwt = jwt.encode(to_encode, settings.JWT_SECRET_KEY, algorithm=settings.JWT_ALGORITHM)
return encoded_jwt
def decode_token(token: str) -> Optional[dict]:
"""
解码JWT Token
:param token: JWT Token字符串
:return: 解码后的数据或None
"""
try:
payload = jwt.decode(token, settings.JWT_SECRET_KEY, algorithms=[settings.JWT_ALGORITHM])
return payload
except JWTError:
return None
def verify_access_token(token: str) -> Optional[dict]:
"""
验证Access Token
:param token: JWT Token字符串
:return: 解码后的数据或None(如果无效或不是access token
"""
payload = decode_token(token)
if payload and payload.get("type") == "access":
return payload
return None
def verify_refresh_token(token: str) -> Optional[dict]:
"""
验证Refresh Token
:param token: JWT Token字符串
:return: 解码后的数据或None(如果无效或不是refresh token
"""
payload = decode_token(token)
if payload and payload.get("type") == "refresh":
return payload
return None
async def get_current_user(
request: Request,
db: AsyncSession = Depends(get_db)
) -> Any:
"""
获取当前登录用户的依赖函数
注意:此依赖配合AuthMiddleware使用
- 中间件负责验证token有效性
- 此依赖从request.state获取用户ID,然后查询完整用户对象
:param request: 请求对象
:param token: JWT Token(用于Swagger显示锁图标)
:param db: 数据库会话
:return: 当前用户对象
:raises HTTPException: 如果用户不存在
"""
# 优先从request.state获取(中间件已验证)
user_id = getattr(request.state, "user_id", None)
if not user_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无效的认证凭据",
headers={"WWW-Authenticate": "Bearer"},
)
# 延迟导入避免循环依赖
from core.user.service import UserService
user = await UserService.get_by_id(db, user_id)
if user is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户不存在",
headers={"WWW-Authenticate": "Bearer"},
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="用户已被禁用"
)
return user
async def get_current_user_id(
request: Request,
token: str = Depends(oauth2_scheme),
) -> str:
"""
获取当前用户ID的依赖函数(轻量级,不查询数据库)
:param request: 请求对象
:param token: JWT Token(用于Swagger显示锁图标)
:return: 当前用户ID
"""
# 优先从request.state获取(中间件已验证)
user_id = getattr(request.state, "user_id", None)
if not user_id and token:
payload = verify_access_token(token)
if payload:
user_id = payload.get("sub")
if not user_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无效的认证凭据",
headers={"WWW-Authenticate": "Bearer"},
)
return user_id
async def get_current_active_user(
request: Request,
token: str = Depends(oauth2_scheme),
db: AsyncSession = Depends(get_db)
) -> Any:
"""
获取当前活跃用户的依赖函数
:param request: 请求对象
:param token: JWT Token(用于Swagger显示锁图标)
:param db: 数据库会话
:return: 当前活跃用户对象
:raises HTTPException: 如果用户状态不正常
"""
user_id = getattr(request.state, "user_id", None)
if not user_id and token:
payload = verify_access_token(token)
if payload:
user_id = payload.get("sub")
if not user_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无效的认证凭据",
headers={"WWW-Authenticate": "Bearer"},
)
from core.user.service import UserService
user = await UserService.get_by_id(db, user_id)
if user is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户不存在",
headers={"WWW-Authenticate": "Bearer"},
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="用户已被禁用"
)
if user.user_status != 1:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="用户状态异常"
)
return user
async def get_current_superuser(
request: Request,
token: str = Depends(oauth2_scheme),
db: AsyncSession = Depends(get_db)
) -> Any:
"""
获取当前超级管理员的依赖函数
:param request: 请求对象
:param token: JWT Token(用于Swagger显示锁图标)
:param db: 数据库会话
:return: 当前超级管理员对象
:raises HTTPException: 如果不是超级管理员
"""
user_id = getattr(request.state, "user_id", None)
if not user_id and token:
payload = verify_access_token(token)
if payload:
user_id = payload.get("sub")
if not user_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="无效的认证凭据",
headers={"WWW-Authenticate": "Bearer"},
)
from core.user.service import UserService
user = await UserService.get_by_id(db, user_id)
if user is None:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="用户不存在",
headers={"WWW-Authenticate": "Bearer"},
)
if not user.is_active:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="用户已被禁用"
)
if not user.is_superuser:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="需要超级管理员权限"
)
return user