Build lightweight AI agent admin

This commit is contained in:
Codex
2026-06-08 18:14:59 +08:00
commit e164840f43
2530 changed files with 435693 additions and 0 deletions
View File
+449
View File
@@ -0,0 +1,449 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
Auth Middleware - 全局认证和鉴权中间件
功能:
1. 认证(Authentication):验证JWT Token的有效性
2. 鉴权(Authorization):基于API路径的动态权限检查
"""
from typing import List, Optional, Callable
import re
from fastapi import Request, HTTPException, status
from fastapi.responses import JSONResponse
from starlette.middleware.base import BaseHTTPMiddleware
from utils.security import verify_access_token
from utils.context import set_current_user_context, clear_current_user_context
# 默认白名单路由(不需要认证)
DEFAULT_WHITE_LIST = [
# 认证相关
"/api/core/login",
"/api/core/refresh_token",
# 文档相关
"/docs",
"/redoc",
"/openapi.json",
# 健康检查
"/",
"/health",
# UI配置(前端初始化时需要无认证访问)
"/api/core/ui_config/preferences",
]
# OAuth白名单正则模式
OAUTH_WHITE_LIST_PATTERNS = [
r"^/api/core/oauth/.*/authorize$", # OAuth授权URL获取
r"^/api/core/oauth/.*/callback$", # OAuth回调处理
]
# WebSocket白名单正则模式(WebSocket自己处理Token认证)
WEBSOCKET_WHITE_LIST_PATTERNS = [
r"^/ws/.*", # 所有WebSocket路径
]
# 白名单正则模式(支持通配符)
DEFAULT_WHITE_LIST_PATTERNS = [
r"^/docs.*",
r"^/redoc.*",
r"^/api/core/applications/code/.*", # 根据编码获取应用(子应用初始化需要)
r"^/api/core/file_manager/stream/.*",
r"^/api/core/file_manager/signature/info/.*", # 签名令牌信息(移动端无需登录)
r"^/api/core/file_manager/signature/upload/.*", # 上传签名图片(移动端无需登录,需验证令牌)
r"^/api/core/file_manager/signature/complete/.*", # 完成签名(移动端无需登录)
r"^/api/core/contract/mobile/.*", # 合同移动端签署(无需登录)
r"^/api/core/dingtalk-sync/callback$", # 钉钉事件回调(钉钉服务器推送,无需登录)
r"^/api/core/wecom-sync/callback$", # 企业微信事件回调(企业微信服务器推送,无需登录)
r"^/api/core/feishu-sync/callback$", # 飞书事件回调(飞书服务器推送,无需登录)
*OAUTH_WHITE_LIST_PATTERNS, # OAuth相关接口
*WEBSOCKET_WHITE_LIST_PATTERNS, # WebSocket相关接口
]
# 允许使用Query参数传递Token的API路径模式(出于安全考虑,仅限特定接口)
# 主要用于文件下载、流式传输等无法设置Header的场景
QUERY_TOKEN_ALLOWED_PATTERNS = [
r"^/api/core/file_manager/proxy/.*", # 文件代理访问
r"^/api/core/file_manager/file/download.*", # 文件下载
]
# 鉴权白名单(不需要权限检查,但需要认证)
DEFAULT_PERMISSION_WHITE_LIST = [
"/api/core/userinfo", # 获取当前用户信息
"/api/core/logout", # 登出
"/api/core/user/change-password", # 修改密码
]
# 鉴权白名单正则模式
DEFAULT_PERMISSION_WHITE_LIST_PATTERNS = [
r"^/api/core/dict/type/.*/data$", # 字典数据查询(所有登录用户都可以访问)
]
class AuthPermissionMiddleware(BaseHTTPMiddleware):
"""
全局认证和鉴权中间件
功能:
- 认证:验证JWT Token的有效性
- 鉴权:基于API路径的动态权限检查
工作流程:
1. 检查是否在白名单中,是则直接放行
2. 验证JWT Token,获取用户信息
3. 根据请求的API路径和方法,查找Permission表中是否有对应的权限记录
4. 如果有权限记录,检查用户的角色是否关联了该权限
5. 如果用户角色有该权限,则放行;否则返回403
6. 如果Permission表中没有该API的权限记录,则默认放行
"""
def __init__(
self,
app,
white_list: Optional[List[str]] = None,
white_list_patterns: Optional[List[str]] = None,
permission_white_list: Optional[List[str]] = None,
permission_white_list_patterns: Optional[List[str]] = None,
enable_permission_check: bool = True,
):
"""
初始化中间件
:param app: FastAPI应用
:param white_list: 认证白名单路由列表(精确匹配,不需要认证)
:param white_list_patterns: 认证白名单正则模式列表(不需要认证)
:param permission_white_list: 鉴权白名单路由列表(需要认证,但不需要权限检查)
:param permission_white_list_patterns: 鉴权白名单正则模式列表(需要认证,但不需要权限检查)
:param enable_permission_check: 是否启用权限检查(默认True)
"""
super().__init__(app)
self.white_list = set(white_list or DEFAULT_WHITE_LIST)
self.white_list_patterns = [
re.compile(p) for p in (white_list_patterns or DEFAULT_WHITE_LIST_PATTERNS)
]
self.permission_white_list = set(permission_white_list or DEFAULT_PERMISSION_WHITE_LIST)
self.permission_white_list_patterns = [
re.compile(p) for p in (permission_white_list_patterns or DEFAULT_PERMISSION_WHITE_LIST_PATTERNS)
]
self.enable_permission_check = enable_permission_check
def is_white_listed(self, path: str) -> bool:
"""
检查路径是否在白名单中(不需要认证)
:param path: 请求路径
:return: 是否在白名单中
"""
# 精确匹配
if path in self.white_list:
return True
# 正则匹配
for pattern in self.white_list_patterns:
if pattern.match(path):
return True
return False
def is_permission_white_listed(self, path: str) -> bool:
"""
检查路径是否在鉴权白名单中(需要认证,但不需要权限检查)
:param path: 请求路径
:return: 是否在鉴权白名单中
"""
# 精确匹配
if path in self.permission_white_list:
return True
# 正则匹配
for pattern in self.permission_white_list_patterns:
if pattern.match(path):
return True
return False
def _is_query_token_allowed(self, path: str) -> bool:
"""
检查路径是否允许使用Query参数传递Token
出于安全考虑,仅允许特定接口使用Query Token
"""
for pattern in [re.compile(p) for p in QUERY_TOKEN_ALLOWED_PATTERNS]:
if pattern.match(path):
return True
return False
def _extract_token(self, request: Request) -> str | None:
"""
从请求中提取Token
支持两种方式:
1. Authorization Header: Bearer <token>(所有接口)
2. Query参数: ?token=<token>(仅限特定接口,如文件下载)
"""
# 优先从Authorization头获取
auth_header = request.headers.get("Authorization")
if auth_header:
parts = auth_header.split()
if len(parts) == 2 and parts[0].lower() == "bearer":
return parts[1]
# Query参数方式仅限特定接口
path = request.url.path
if self._is_query_token_allowed(path):
token = request.query_params.get("token")
if token:
return token
return None
async def _verify_api_token(self, raw_token: str):
"""
验证 API Token (Personal Access Token)
:return: (user_id, username) 或 (None, None)
"""
from app.database import AsyncSessionLocal
from core.api_token.service import ApiTokenService
async with AsyncSessionLocal() as db:
token_record = await ApiTokenService.verify_token(db, raw_token)
if not token_record:
return None, None
from core.user.service import UserService
user = await UserService.get_by_id(db, token_record.user_id)
if not user or not user.is_active:
return None, None
return user.id, user.username
async def dispatch(self, request: Request, call_next: Callable):
"""处理请求"""
path = request.url.path
method = request.method
# 白名单路由直接放行
if self.is_white_listed(path):
return await call_next(request)
# OPTIONS请求放行(CORS预检)
if method == "OPTIONS":
return await call_next(request)
# ========== 认证(Authentication==========
token = self._extract_token(request)
if not token:
return JSONResponse(
status_code=status.HTTP_401_UNAUTHORIZED,
content={"detail": "未提供认证凭据"},
headers={"WWW-Authenticate": "Bearer"},
)
# 检查是否为 API Token (Personal Access Token)
is_api_token = token.startswith("zqpat_")
user_id = None
username = None
if is_api_token:
try:
user_id, username = await self._verify_api_token(token)
except Exception:
import logging
logging.getLogger(__name__).exception("API Token验证异常")
user_id, username = None, None
if not user_id:
return JSONResponse(
status_code=status.HTTP_401_UNAUTHORIZED,
content={"detail": "无效或过期的API Token"},
headers={"WWW-Authenticate": "Bearer"},
)
else:
payload = verify_access_token(token)
if not payload:
return JSONResponse(
status_code=status.HTTP_401_UNAUTHORIZED,
content={"detail": "无效或过期的Token"},
headers={"WWW-Authenticate": "Bearer"},
)
# 检查 access token 是否在 Redis 白名单中(防止已登出的token被使用)
user_id = payload.get("sub")
device_id = payload.get("device_id")
username = payload.get("username")
if user_id and device_id:
from utils.redis import RedisClient
redis = await RedisClient.get_client()
access_token_key = f"access_token:{user_id}:{device_id}"
token_exists = await redis.exists(access_token_key)
if not token_exists:
return JSONResponse(
status_code=status.HTTP_401_UNAUTHORIZED,
content={"detail": "Token已失效,请重新登录"},
headers={"WWW-Authenticate": "Bearer"},
)
# 从 Redis 缓存获取用户动态信息(角色、部门、超管等)
from utils.user_info_cache import get_cached_user_info, load_user_info_from_db
user_info = await get_cached_user_info(user_id)
if not user_info:
user_info = await load_user_info_from_db(user_id)
if not user_info:
return JSONResponse(
status_code=status.HTTP_401_UNAUTHORIZED,
content={"detail": "用户不存在"},
headers={"WWW-Authenticate": "Bearer"},
)
if not username:
username = user_info.get("username")
role_ids = user_info.get("role_ids", [])
role_id = role_ids[0] if role_ids else None
dept_id = user_info.get("dept_id")
is_superuser = user_info.get("is_superuser", False)
# 将用户信息存入request.state
request.state.user_id = user_id
request.state.username = username
request.state.role_id = role_id
request.state.role_ids = role_ids
request.state.dept_id = dept_id
request.state.is_superuser = is_superuser
request.state.token_payload = {} if is_api_token else payload
# 设置完整用户信息和请求信息到上下文(供Service层使用)
set_current_user_context(
user_id=user_id,
role_id=role_id,
role_ids=role_ids,
dept_id=dept_id,
is_superuser=is_superuser,
username=username,
request_path=path,
http_method=method
)
try:
# ========== 鉴权(Authorization==========
if self.enable_permission_check:
# 检查是否在鉴权白名单中
if self.is_permission_white_listed(path):
# 在鉴权白名单中,跳过权限检查
pass
# 超级管理员跳过权限检查
elif is_superuser:
pass
else:
# 普通用户需要进行权限检查
from app.database import AsyncSessionLocal
from utils.permission import check_api_permission
async with AsyncSessionLocal() as db:
has_permission, error_msg = await check_api_permission(
db=db,
user_id=user_id,
role_ids=role_ids,
is_superuser=is_superuser,
request_path=path,
http_method=method,
)
if not has_permission:
return JSONResponse(
status_code=status.HTTP_403_FORBIDDEN,
content={"detail": error_msg or "权限不足"},
)
response = await call_next(request)
return response
finally:
# 请求结束后清除上下文
clear_current_user_context()
def get_auth_middleware(
white_list: Optional[List[str]] = None,
white_list_patterns: Optional[List[str]] = None,
permission_white_list: Optional[List[str]] = None,
permission_white_list_patterns: Optional[List[str]] = None,
) -> type:
"""
获取配置好的认证+鉴权中间件类
使用方式:
app.add_middleware(get_auth_middleware(
white_list=["/public"], # 不需要认证的接口
permission_white_list=["/api/my/profile"] # 需要认证但不需要权限检查的接口
))
:param white_list: 额外的认证白名单路由(不需要认证)
:param white_list_patterns: 额外的认证白名单正则模式(不需要认证)
:param permission_white_list: 额外的鉴权白名单路由(需要认证,但不需要权限检查)
:param permission_white_list_patterns: 额外的鉴权白名单正则模式(需要认证,但不需要权限检查)
:return: 配置好的中间件类
"""
merged_white_list = list(DEFAULT_WHITE_LIST)
if white_list:
merged_white_list.extend(white_list)
merged_patterns = list(DEFAULT_WHITE_LIST_PATTERNS)
if white_list_patterns:
merged_patterns.extend(white_list_patterns)
merged_permission_white_list = list(DEFAULT_PERMISSION_WHITE_LIST)
if permission_white_list:
merged_permission_white_list.extend(permission_white_list)
merged_permission_patterns = list(DEFAULT_PERMISSION_WHITE_LIST_PATTERNS)
if permission_white_list_patterns:
merged_permission_patterns.extend(permission_white_list_patterns)
class ConfiguredMiddleware(AuthPermissionMiddleware):
def __init__(self, app):
super().__init__(
app,
white_list=merged_white_list,
white_list_patterns=merged_patterns,
permission_white_list=merged_permission_white_list,
permission_white_list_patterns=merged_permission_patterns,
enable_permission_check=True,
)
return ConfiguredMiddleware
def get_auth_permission_middleware(
white_list: Optional[List[str]] = None,
white_list_patterns: Optional[List[str]] = None,
permission_white_list: Optional[List[str]] = None,
permission_white_list_patterns: Optional[List[str]] = None,
) -> type:
"""
获取配置好的认证+鉴权中间件类(别名函数,与 get_auth_middleware 相同)
使用方式:
app.add_middleware(get_auth_permission_middleware(
white_list=["/public"], # 不需要认证的接口
permission_white_list=["/api/my/profile"] # 需要认证但不需要权限检查的接口
))
:param white_list: 额外的认证白名单路由(不需要认证)
:param white_list_patterns: 额外的认证白名单正则模式(不需要认证)
:param permission_white_list: 额外的鉴权白名单路由(需要认证,但不需要权限检查)
:param permission_white_list_patterns: 额外的鉴权白名单正则模式(需要认证,但不需要权限检查)
:return: 配置好的中间件类
"""
return get_auth_middleware(
white_list=white_list,
white_list_patterns=white_list_patterns,
permission_white_list=permission_white_list,
permission_white_list_patterns=permission_white_list_patterns,
)
+112
View File
@@ -0,0 +1,112 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
客户端信息提取工具
从HTTP请求中提取客户端IP、浏览器、操作系统等信息
"""
import hashlib
from fastapi import Request
def get_client_info(request: Request) -> dict:
"""
从请求中提取客户端信息
Args:
request: FastAPI Request对象
Returns:
dict: 包含以下字段:
- login_ip: 客户端IP地址
- user_agent: 完整的User-Agent字符串
- browser_type: 浏览器类型
- os_type: 操作系统类型
- device_type: 设备类型 (desktop/mobile/tablet)
"""
# 获取IP地址
login_ip = request.headers.get("X-Forwarded-For", "")
if login_ip:
login_ip = login_ip.split(",")[0].strip()
else:
login_ip = request.client.host if request.client else "0.0.0.0"
# 获取User-Agent
user_agent = request.headers.get("User-Agent", "")
# 解析浏览器和操作系统
browser_type = None
os_type = None
device_type = "desktop"
ua_lower = user_agent.lower()
# 浏览器检测
if "chrome" in ua_lower and "edg" not in ua_lower:
browser_type = "Chrome"
elif "firefox" in ua_lower:
browser_type = "Firefox"
elif "safari" in ua_lower and "chrome" not in ua_lower:
browser_type = "Safari"
elif "edg" in ua_lower:
browser_type = "Edge"
elif "msie" in ua_lower or "trident" in ua_lower:
browser_type = "IE"
# 操作系统检测
if "windows" in ua_lower:
os_type = "Windows"
elif "mac os" in ua_lower or "macos" in ua_lower:
os_type = "macOS"
elif "linux" in ua_lower:
os_type = "Linux"
elif "android" in ua_lower:
os_type = "Android"
device_type = "mobile"
elif "iphone" in ua_lower or "ipad" in ua_lower:
os_type = "iOS"
device_type = "mobile" if "iphone" in ua_lower else "tablet"
return {
"login_ip": login_ip,
"user_agent": user_agent,
"browser_type": browser_type,
"os_type": os_type,
"device_type": device_type,
}
def get_client_ip(request: Request) -> str:
"""
仅获取客户端IP地址
Args:
request: FastAPI Request对象
Returns:
str: 客户端IP地址
"""
login_ip = request.headers.get("X-Forwarded-For", "")
if login_ip:
return login_ip.split(",")[0].strip()
return request.client.host if request.client else "0.0.0.0"
def get_device_id(request: Request) -> str:
"""
生成设备唯一标识
基于 User-Agent 和 IP 地址生成设备ID,用于多设备登录场景
Args:
request: FastAPI Request对象
Returns:
str: 设备唯一标识(MD5哈希值)
"""
user_agent = request.headers.get("User-Agent", "")
client_ip = get_client_ip(request)
# 使用 User-Agent + IP 生成设备标识
device_string = f"{user_agent}:{client_ip}"
device_id = hashlib.md5(device_string.encode()).hexdigest()
return device_id
+77
View File
@@ -0,0 +1,77 @@
"""
上下文管理器 - 用于在整个请求生命周期中共享数据
"""
from contextvars import ContextVar
from typing import Optional, Dict, Any
# 当前用户信息上下文
current_user_context: ContextVar[Optional[Dict[str, Any]]] = ContextVar('current_user', default=None)
def get_current_user_id_from_context() -> Optional[str]:
"""
从上下文中获取当前用户ID
:return: 用户ID,如果未设置则返回None
"""
user_info = current_user_context.get()
return user_info.get('user_id') if user_info else None
def get_current_user_info_from_context() -> Optional[Dict[str, Any]]:
"""
从上下文中获取当前用户完整信息
:return: 用户信息字典,包含 user_id, role_id, dept_id, is_superuser 等
"""
return current_user_context.get()
def set_current_user_context(
user_id: str,
role_id: Optional[str] = None,
dept_id: Optional[str] = None,
is_superuser: bool = False,
username: Optional[str] = None,
request_path: Optional[str] = None,
http_method: Optional[str] = None,
role_ids: Optional[list] = None
) -> None:
"""
设置当前用户信息到上下文
:param user_id: 用户ID
:param role_id: 角色ID(单个,向后兼容)
:param dept_id: 部门ID
:param is_superuser: 是否超级管理员
:param username: 用户名
:param request_path: 请求路径
:param http_method: HTTP方法
:param role_ids: 角色ID列表(多个角色)
"""
# 如果提供了 role_ids,使用它;否则从 role_id 构建
if role_ids is None and role_id:
role_ids = [role_id]
elif role_ids is None:
role_ids = []
current_user_context.set({
'user_id': user_id,
'role_id': role_id, # 保持向后兼容
'role_ids': role_ids, # 新增:支持多角色
'dept_id': dept_id,
'is_superuser': is_superuser,
'username': username,
'request_path': request_path,
'http_method': http_method
})
def clear_current_user_context() -> None:
"""清除当前用户上下文"""
current_user_context.set(None)
# 保持向后兼容的别名
set_current_user_id_context = lambda user_id: set_current_user_context(user_id)
clear_current_user_id_context = clear_current_user_context
+117
View File
@@ -0,0 +1,117 @@
from typing import Optional
from fastapi import Request
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from utils.context import get_current_user_info_from_context
from utils.permission import apply_data_scope_filter
def get_data_scope_filter(request: Request):
"""
FastAPI依赖函数:获取当前请求的数据权限过滤参数
使用方式:
@router.get("/users")
async def get_users(
data_scope = Depends(get_data_scope_filter),
db: AsyncSession = Depends(get_db)
):
items, total = await UserService.get_list_with_data_scope(
db=db,
**data_scope, # 自动展开为 role_id, is_superuser, user_id, user_dept_id, request_path, http_method
page=page,
page_size=page_size
)
:param request: FastAPI Request对象
:return: 数据权限参数字典
"""
# 从上下文获取用户信息(由中间件设置)
user_info = get_current_user_info_from_context()
if not user_info:
# 如果上下文中没有用户信息,返回默认值(不应该发生,因为有认证中间件)
return {
"role_id": None,
"is_superuser": False,
"user_id": None,
"user_dept_id": None,
"request_path": request.url.path,
"http_method": request.method
}
return {
"role_id": user_info.get("role_id"),
"is_superuser": user_info.get("is_superuser", False),
"user_id": user_info.get("user_id"),
"user_dept_id": user_info.get("dept_id"),
"request_path": request.url.path,
"http_method": request.method
}
async def get_data_scope_dict(
request: Request,
db: AsyncSession = Depends(get_db)
):
"""
FastAPI依赖函数:获取当前请求的数据权限过滤条件字典
使用方式:
@router.get("/users")
async def get_users(
data_scope_dict = Depends(get_data_scope_dict),
db: AsyncSession = Depends(get_db)
):
# data_scope_dict 包含 filter_type, user_id, dept_id, dept_ids 等
if data_scope_dict['filter_type'] == 'self':
query = query.where(User.id == data_scope_dict['user_id'])
...
:param request: FastAPI Request对象
:param db: 数据库会话
:return: 数据权限过滤条件字典
"""
# 从上下文获取用户信息(由中间件设置)
user_info = get_current_user_info_from_context()
if not user_info:
# 如果上下文中没有用户信息,返回默认的全部数据权限
return {
'scope': 0,
'filter_type': 'all',
'user_id': None,
'dept_id': None,
'dept_ids': None,
}
return await apply_data_scope_filter(
db=db,
role_id=user_info.get("role_id"),
is_superuser=user_info.get("is_superuser", False),
user_id=user_info.get("user_id"),
user_dept_id=user_info.get("dept_id"),
request_path=request.url.path,
http_method=request.method
)
def get_current_user_id(current_user = Depends(get_current_user)) -> str:
"""
FastAPI依赖函数:获取当前用户ID
使用方式:
@router.post("/posts")
async def create_post(
data: PostCreate,
current_user_id: str = Depends(get_current_user_id),
db: AsyncSession = Depends(get_db)
):
post = await PostService.create(db, data, current_user_id=current_user_id)
return post
:param current_user: 当前用户
:return: 用户ID
"""
return current_user.id
+144
View File
@@ -0,0 +1,144 @@
from io import BytesIO
from typing import List, Dict, Any, Type, Optional
from openpyxl import Workbook, load_workbook
from openpyxl.styles import Font, Alignment, Border, Side, PatternFill
from pydantic import BaseModel
class ExcelHandler:
"""Excel处理工具类"""
# 表头样式
HEADER_FONT = Font(bold=True, color="FFFFFF")
HEADER_FILL = PatternFill(start_color="4472C4", end_color="4472C4", fill_type="solid")
HEADER_ALIGNMENT = Alignment(horizontal="center", vertical="center")
# 边框样式
THIN_BORDER = Border(
left=Side(style="thin"),
right=Side(style="thin"),
top=Side(style="thin"),
bottom=Side(style="thin")
)
@classmethod
def export_to_excel(
cls,
data: List[Dict[str, Any]],
columns: Dict[str, str],
sheet_name: str = "Sheet1"
) -> BytesIO:
"""
导出数据到Excel
:param data: 数据列表,每个元素是一个字典
:param columns: 列映射,格式为 {字段名: 显示名}
:param sheet_name: 工作表名称
:return: Excel文件的BytesIO对象
"""
wb = Workbook()
ws = wb.active
ws.title = sheet_name
# 写入表头
headers = list(columns.values())
field_names = list(columns.keys())
for col_idx, header in enumerate(headers, 1):
cell = ws.cell(row=1, column=col_idx, value=header)
cell.font = cls.HEADER_FONT
cell.fill = cls.HEADER_FILL
cell.alignment = cls.HEADER_ALIGNMENT
cell.border = cls.THIN_BORDER
# 写入数据
for row_idx, row_data in enumerate(data, 2):
for col_idx, field in enumerate(field_names, 1):
value = row_data.get(field, "")
cell = ws.cell(row=row_idx, column=col_idx, value=value)
cell.border = cls.THIN_BORDER
cell.alignment = Alignment(vertical="center")
# 自动调整列宽
for col_idx, header in enumerate(headers, 1):
max_length = len(str(header))
for row in ws.iter_rows(min_row=2, min_col=col_idx, max_col=col_idx):
for cell in row:
if cell.value:
max_length = max(max_length, len(str(cell.value)))
ws.column_dimensions[ws.cell(row=1, column=col_idx).column_letter].width = min(max_length + 2, 50)
# 保存到BytesIO
output = BytesIO()
wb.save(output)
output.seek(0)
return output
@classmethod
def import_from_excel(
cls,
file_content: bytes,
columns: Dict[str, str],
schema: Optional[Type[BaseModel]] = None
) -> List[Dict[str, Any]]:
"""
从Excel导入数据
:param file_content: Excel文件内容
:param columns: 列映射,格式为 {字段名: 显示名}
:param schema: 可选的Pydantic Schema用于数据验证
:return: 数据列表
"""
wb = load_workbook(filename=BytesIO(file_content), read_only=True)
ws = wb.active
# 读取表头,建立显示名到字段名的映射
header_to_field = {v: k for k, v in columns.items()}
rows = list(ws.iter_rows(values_only=True))
if not rows:
return []
# 第一行是表头
headers = rows[0]
field_indices = {}
for idx, header in enumerate(headers):
if header in header_to_field:
field_indices[idx] = header_to_field[header]
# 读取数据行
result = []
for row in rows[1:]:
if not any(row): # 跳过空行
continue
row_data = {}
for idx, field_name in field_indices.items():
value = row[idx] if idx < len(row) else None
row_data[field_name] = value
# 如果提供了schema,进行数据验证
if schema:
try:
validated = schema(**row_data)
row_data = validated.model_dump()
except Exception:
continue # 跳过验证失败的行
result.append(row_data)
wb.close()
return result
@classmethod
def generate_template(
cls,
columns: Dict[str, str],
sheet_name: str = "Sheet1"
) -> BytesIO:
"""
生成导入模板
:param columns: 列映射,格式为 {字段名: 显示名}
:param sheet_name: 工作表名称
:return: Excel模板文件的BytesIO对象
"""
return cls.export_to_excel([], columns, sheet_name)
+61
View File
@@ -0,0 +1,61 @@
"""
FastAPI 0.121.x OpenAPI schema 生成兼容性补丁
修复 _remap_definitions_and_field_mappings 中 KeyError: '$ref' 的问题。
当可选请求体参数(Body/Form 参数带有默认值)生成 JSON Schema 时,
Pydantic v2 会产生 {"allOf": [{"$ref": "..."}], "default": None} 格式,
而 FastAPI 0.121.x 的实现假定 schema 顶层直接有 "$ref" 键,导致崩溃。
此补丁在读取 $ref 时兼容 allOf 包裹格式。
升级 FastAPI 到修复此问题的版本后可移除本补丁。
"""
import fastapi._compat.v2 as _v2
_original_remap = _v2._remap_definitions_and_field_mappings
def _extract_ref(schema: dict) -> str | None:
"""从 schema 中提取 $ref,兼容直接引用和 allOf 包裹两种格式"""
if "$ref" in schema:
return schema["$ref"]
if "allOf" in schema:
for item in schema["allOf"]:
if isinstance(item, dict) and "$ref" in item:
return item["$ref"]
return None
def _patched_remap(*, model_name_map, definitions, field_mapping):
old_name_to_new_name_map = {}
for field_key, schema in field_mapping.items():
model = field_key[0].type_
if model not in model_name_map:
continue
new_name = model_name_map[model]
ref = _extract_ref(schema)
if ref is None:
continue
old_name = ref.split("/")[-1]
if old_name in {f"{new_name}-Input", f"{new_name}-Output"}:
continue
old_name_to_new_name_map[old_name] = new_name
new_field_mapping = {}
for field_key, schema in field_mapping.items():
new_schema = _v2._replace_refs(
schema=schema, old_name_to_new_name_map=old_name_to_new_name_map
)
new_field_mapping[field_key] = new_schema
new_definitions = {}
for key, value in definitions.items():
new_key = old_name_to_new_name_map.get(key, key)
new_value = _v2._replace_refs(
schema=value, old_name_to_new_name_map=old_name_to_new_name_map
)
new_definitions[new_key] = new_value
return new_field_mapping, new_definitions
_v2._remap_definitions_and_field_mappings = _patched_remap
+418
View File
@@ -0,0 +1,418 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
Permission Utils - 基于API路径的动态权限鉴权
工作原理:
1. 用户访问某个API时,根据请求的路径和方法,查找Permission表中是否有对应的权限记录
2. 如果有权限记录,检查用户的角色是否关联了该权限
3. 如果用户角色有该权限,则放行;否则返回403
4. 如果Permission表中没有该API的权限记录,则默认放行(未配置权限的API不做限制)
"""
import json
import re
from typing import List, Optional, Dict, Any, Set
from fastapi import HTTPException, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from utils.redis import RedisClient
# Redis 缓存 key 前缀和 TTL
API_PERMISSION_CACHE_KEY = "api_permission:path_map" # Hash: (path:method) -> permission_id
ROLE_PERMISSION_CACHE_PREFIX = "api_permission:role:" # role_permission:role:{role_id} -> Set[permission_id]
API_PERMISSION_CACHE_TTL = 300 # API路径映射缓存 5分钟
ROLE_PERMISSION_CACHE_TTL = 300 # 角色权限缓存 5分钟
# HTTP方法映射(与Permission模型中的定义一致)
HTTP_METHOD_MAP = {
'GET': 0,
'POST': 1,
'PUT': 2,
'DELETE': 3,
'PATCH': 4,
'ALL': 5,
}
class APIPermissionChecker:
"""
基于API路径的动态权限检查器
在AuthMiddleware中调用,根据请求的API路径和方法检查用户是否有权限
"""
def __init__(self):
pass
async def load_permissions_cache(self, db: AsyncSession):
"""
加载所有API权限到 Redis 缓存
建议在应用启动时调用,或者定期刷新
"""
from core.permission.model import Permission
result = await db.execute(
select(Permission).where(
Permission.is_active == True, # noqa: E712
Permission.is_deleted == False, # noqa: E712
Permission.permission_type == 1, # 只缓存API权限
Permission.api_path.isnot(None)
)
)
permissions = result.scalars().all()
redis = await RedisClient.get_client()
# 先删除旧缓存
await redis.delete(API_PERMISSION_CACHE_KEY)
cache_data = {}
for perm in permissions:
if perm.api_path:
# 存储权限IDkey为 "path:method"
cache_key = f"{perm.api_path}:{perm.http_method}"
cache_data[cache_key] = perm.id
# 如果是ALL方法,也存储到各个具体方法
if perm.http_method == 5: # ALL
for method_code in [0, 1, 2, 3, 4]:
key = f"{perm.api_path}:{method_code}"
if key not in cache_data:
cache_data[key] = perm.id
if cache_data:
await redis.hset(API_PERMISSION_CACHE_KEY, mapping=cache_data)
await redis.expire(API_PERMISSION_CACHE_KEY, API_PERMISSION_CACHE_TTL)
async def clear_cache(self):
"""清除所有权限相关的 Redis 缓存"""
redis = await RedisClient.get_client()
# 清除 API 路径映射缓存
await redis.delete(API_PERMISSION_CACHE_KEY)
# 清除所有角色权限缓存
cursor = 0
while True:
cursor, keys = await redis.scan(cursor, match=f"{ROLE_PERMISSION_CACHE_PREFIX}*", count=100)
if keys:
await redis.delete(*keys)
if cursor == 0:
break
def _match_path(self, request_path: str, permission_path: str) -> bool:
"""
匹配请求路径和权限路径
支持路径参数,如 /api/user/{id} 匹配 /api/user/123
"""
# 将权限路径中的{xxx}替换为正则表达式
pattern = re.sub(r'\{[^}]+\}', r'[^/]+', permission_path)
pattern = f'^{pattern}$'
return bool(re.match(pattern, request_path))
async def find_permission_id(self, request_path: str, http_method: str) -> Optional[str]:
"""
根据请求路径和方法查找对应的权限ID(从 Redis 缓存查询)
:param request_path: 请求路径,如 /api/core/user
:param http_method: HTTP方法,如 GET, POST
:return: 权限ID,如果没有找到则返回None
"""
method_code = HTTP_METHOD_MAP.get(http_method.upper(), 0)
redis = await RedisClient.get_client()
# 1. 精确匹配
cache_key = f"{request_path}:{method_code}"
perm_id = await redis.hget(API_PERMISSION_CACHE_KEY, cache_key)
if perm_id:
return perm_id
# 2. 尝试匹配ALL方法
cache_key_all = f"{request_path}:5"
perm_id = await redis.hget(API_PERMISSION_CACHE_KEY, cache_key_all)
if perm_id:
return perm_id
# 3. 路径参数匹配(需要获取所有缓存条目)
all_entries = await redis.hgetall(API_PERMISSION_CACHE_KEY)
if all_entries:
for entry_key, entry_perm_id in all_entries.items():
# entry_key 格式: "path:method"
# 需要从末尾分离出 method(最后一个冒号后的数字)
last_colon = entry_key.rfind(':')
if last_colon == -1:
continue
perm_path = entry_key[:last_colon]
try:
perm_method = int(entry_key[last_colon + 1:])
except ValueError:
continue
if perm_method in (method_code, 5): # 匹配具体方法或ALL
if '{' in perm_path and self._match_path(request_path, perm_path):
return entry_perm_id
return None
async def check_permission(
self,
db: AsyncSession,
user_id: str,
role_id: Optional[str] = None,
is_superuser: bool = False,
request_path: str = "",
http_method: str = "",
role_ids: Optional[List[str]] = None,
) -> tuple[bool, str]:
"""
检查用户是否有访问指定API的权限(支持多角色)
:param db: 数据库会话
:param user_id: 用户ID
:param role_id: 角色ID(单个,向后兼容)
:param is_superuser: 是否超级管理员
:param request_path: 请求路径
:param http_method: HTTP方法
:param role_ids: 角色ID列表(多个角色)
:return: (是否有权限, 错误信息)
"""
# 超级管理员跳过权限检查
if is_superuser:
return True, ""
# 确保缓存已加载
redis = await RedisClient.get_client()
cache_exists = await redis.exists(API_PERMISSION_CACHE_KEY)
if not cache_exists:
await self.load_permissions_cache(db)
# 查找该API对应的权限
permission_id = await self.find_permission_id(request_path, http_method)
# 如果该API没有配置权限,默认放行
if not permission_id:
return True, ""
# 处理角色ID列表
if role_ids is None:
role_ids = [role_id] if role_id else []
# 用户没有角色,无权限
if not role_ids:
return False, "用户未分配角色,无权访问此接口"
# 检查用户的任一角色是否有该权限(只要有一个角色有权限即可)
for rid in role_ids:
has_permission = await self._check_role_has_permission(db, rid, permission_id)
if has_permission:
return True, ""
return False, "权限不足,无权访问此接口"
async def _get_role_permission_ids(
self,
db: AsyncSession,
role_id: str,
) -> Set[str]:
"""
获取角色的所有权限ID集合(Redis 缓存)
"""
redis = await RedisClient.get_client()
cache_key = f"{ROLE_PERMISSION_CACHE_PREFIX}{role_id}"
# 从 Redis 获取缓存
cached = await redis.get(cache_key)
# if cached:
# try:
# return set(json.loads(cached))
# except (json.JSONDecodeError, TypeError):
# pass
# 缓存未命中,从数据库加载
from core.role.model import Role
result = await db.execute(
select(Role)
.options(selectinload(Role.permissions))
.where(
Role.id == role_id,
Role.status == True, # noqa: E712
Role.is_deleted == False # noqa: E712
)
)
role = result.scalar_one_or_none()
perm_ids: Set[str] = set()
if role and role.permissions:
for perm in role.permissions:
if perm.is_active:
perm_ids.add(perm.id)
# 写入 Redis 缓存
await redis.set(cache_key, json.dumps(list(perm_ids)), ex=ROLE_PERMISSION_CACHE_TTL)
return perm_ids
async def _check_role_has_permission(
self,
db: AsyncSession,
role_id: str,
permission_id: str
) -> bool:
"""
检查角色是否有指定权限
"""
perm_ids = await self._get_role_permission_ids(db, role_id)
return permission_id in perm_ids
# 全局权限检查器实例
api_permission_checker = APIPermissionChecker()
async def check_api_permission(
db: AsyncSession,
user_id: str,
role_id: Optional[str] = None,
is_superuser: bool = False,
request_path: str = "",
http_method: str = "",
role_ids: Optional[List[str]] = None,
) -> tuple[bool, str]:
"""
检查API权限的便捷函数(支持多角色)
:return: (是否有权限, 错误信息)
"""
return await api_permission_checker.check_permission(
db, user_id, role_id, is_superuser, request_path, http_method, role_ids
)
async def refresh_permission_cache(db: AsyncSession):
"""
刷新权限缓存
当权限数据变更时调用此函数
"""
await api_permission_checker.load_permissions_cache(db)
async def clear_permission_cache():
"""
清除权限缓存
"""
await api_permission_checker.clear_cache()
async def clear_role_permission_cache(role_id: str):
"""
清除指定角色的权限缓存
:param role_id: 角色ID
"""
redis = await RedisClient.get_client()
await redis.delete(f"{ROLE_PERMISSION_CACHE_PREFIX}{role_id}")
async def clear_all_role_permission_cache():
"""
清除所有角色的权限缓存
"""
redis = await RedisClient.get_client()
cursor = 0
while True:
cursor, keys = await redis.scan(cursor, match=f"{ROLE_PERMISSION_CACHE_PREFIX}*", count=100)
if keys:
await redis.delete(*keys)
if cursor == 0:
break
async def get_user_api_permissions(
db: AsyncSession,
role_id: Optional[str] = None,
is_superuser: bool = False,
role_ids: Optional[List[str]] = None,
) -> Set[str]:
"""
获取用户的API权限列表(返回API路径集合,支持多角色)
"""
if is_superuser:
return {"*"} # 超级管理员有所有权限
# 处理角色ID列表
if role_ids is None:
role_ids = [role_id] if role_id else []
if not role_ids:
return set()
from core.role.model import Role
api_paths = set()
# 合并所有角色的权限
for rid in role_ids:
result = await db.execute(
select(Role)
.options(selectinload(Role.permissions))
.where(
Role.id == rid,
Role.status == True, # noqa: E712
Role.is_deleted == False # noqa: E712
)
)
role = result.scalar_one_or_none()
if role and role.permissions:
for perm in role.permissions:
if perm.is_active and perm.api_path:
api_paths.add(perm.api_path)
return api_paths
async def get_user_permission_codes(
db: AsyncSession,
role_id: Optional[str] = None,
is_superuser: bool = False,
role_ids: Optional[List[str]] = None,
) -> Set[str]:
"""
获取用户的权限代码列表(返回权限代码集合,支持多角色)
"""
if is_superuser:
return {"*"} # 超级管理员有所有权限
# 处理角色ID列表
if role_ids is None:
role_ids = [role_id] if role_id else []
if not role_ids:
return set()
from core.role.model import Role
permission_codes = set()
# 合并所有角色的权限
for rid in role_ids:
result = await db.execute(
select(Role)
.options(selectinload(Role.permissions))
.where(
Role.id == rid,
Role.status == True, # noqa: E712
Role.is_deleted == False # noqa: E712
)
)
role = result.scalar_one_or_none()
if role and role.permissions:
for perm in role.permissions:
if perm.is_active and perm.code:
permission_codes.add(perm.code)
return permission_codes
+251
View File
@@ -0,0 +1,251 @@
"""
Redis缓存模块
提供Redis连接管理和缓存操作工具类
"""
import json
from typing import Optional, Any, Union, List
from contextlib import asynccontextmanager
from redis import asyncio as aioredis
from redis.asyncio import Redis
from app.config import settings
class RedisClient:
"""Redis客户端管理器"""
_client: Optional[Redis] = None
@classmethod
async def get_client(cls) -> Redis:
"""获取Redis客户端实例(单例模式)"""
if cls._client is None:
cls._client = await aioredis.from_url(
settings.REDIS_URL,
encoding="utf-8",
decode_responses=True
)
return cls._client
@classmethod
async def close(cls) -> None:
"""关闭Redis连接"""
if cls._client:
await cls._client.close()
cls._client = None
class CacheManager:
"""
缓存管理器
提供通用的缓存操作方法
"""
def __init__(self, prefix: str = ""):
"""
初始化缓存管理器
:param prefix: 缓存key前缀(会自动添加全局前缀)
"""
self.prefix = f"{settings.CACHE_PREFIX}{prefix}"
def _make_key(self, key: str) -> str:
"""生成完整的缓存key"""
return f"{self.prefix}{key}"
async def get(self, key: str) -> Optional[Any]:
"""
获取缓存
:param key: 缓存key
:return: 缓存值,不存在返回None
"""
client = await RedisClient.get_client()
value = await client.get(self._make_key(key))
if value:
try:
return json.loads(value)
except json.JSONDecodeError:
return value
return None
async def set(
self,
key: str,
value: Any,
expire: Optional[int] = None
) -> bool:
"""
设置缓存
:param key: 缓存key
:param value: 缓存值(自动JSON序列化)
:param expire: 过期时间(秒),默认使用配置值
:return: 是否成功
"""
client = await RedisClient.get_client()
expire = expire or settings.CACHE_DEFAULT_EXPIRE
if isinstance(value, (dict, list)):
value = json.dumps(value, ensure_ascii=False, default=str)
elif not isinstance(value, str):
value = json.dumps(value, default=str)
return await client.set(self._make_key(key), value, ex=expire)
async def delete(self, key: str) -> int:
"""
删除缓存
:param key: 缓存key
:return: 删除的key数量
"""
client = await RedisClient.get_client()
return await client.delete(self._make_key(key))
async def delete_pattern(self, pattern: str) -> int:
"""
根据模式删除缓存
:param pattern: 匹配模式(如 user:*
:return: 删除的key数量
"""
client = await RedisClient.get_client()
full_pattern = self._make_key(pattern)
keys = []
async for key in client.scan_iter(match=full_pattern):
keys.append(key)
if keys:
return await client.delete(*keys)
return 0
async def exists(self, key: str) -> bool:
"""
检查缓存是否存在
:param key: 缓存key
:return: 是否存在
"""
client = await RedisClient.get_client()
return await client.exists(self._make_key(key)) > 0
async def expire(self, key: str, seconds: int) -> bool:
"""
设置过期时间
:param key: 缓存key
:param seconds: 过期时间(秒)
:return: 是否成功
"""
client = await RedisClient.get_client()
return await client.expire(self._make_key(key), seconds)
async def ttl(self, key: str) -> int:
"""
获取剩余过期时间
:param key: 缓存key
:return: 剩余秒数,-1表示永不过期,-2表示不存在
"""
client = await RedisClient.get_client()
return await client.ttl(self._make_key(key))
async def incr(self, key: str, amount: int = 1) -> int:
"""
自增
:param key: 缓存key
:param amount: 增加量
:return: 增加后的值
"""
client = await RedisClient.get_client()
return await client.incrby(self._make_key(key), amount)
async def decr(self, key: str, amount: int = 1) -> int:
"""
自减
:param key: 缓存key
:param amount: 减少量
:return: 减少后的值
"""
client = await RedisClient.get_client()
return await client.decrby(self._make_key(key), amount)
async def hget(self, name: str, key: str) -> Optional[Any]:
"""获取Hash字段值"""
client = await RedisClient.get_client()
value = await client.hget(self._make_key(name), key)
if value:
try:
return json.loads(value)
except json.JSONDecodeError:
return value
return None
async def hset(self, name: str, key: str, value: Any) -> int:
"""设置Hash字段值"""
client = await RedisClient.get_client()
if isinstance(value, (dict, list)):
value = json.dumps(value, ensure_ascii=False, default=str)
elif not isinstance(value, str):
value = json.dumps(value, default=str)
return await client.hset(self._make_key(name), key, value)
async def hdel(self, name: str, *keys: str) -> int:
"""删除Hash字段"""
client = await RedisClient.get_client()
return await client.hdel(self._make_key(name), *keys)
async def hgetall(self, name: str) -> dict:
"""获取Hash所有字段"""
client = await RedisClient.get_client()
data = await client.hgetall(self._make_key(name))
result = {}
for k, v in data.items():
try:
result[k] = json.loads(v)
except json.JSONDecodeError:
result[k] = v
return result
async def lpush(self, key: str, *values: Any) -> int:
"""列表左侧插入"""
client = await RedisClient.get_client()
serialized = [
json.dumps(v, ensure_ascii=False, default=str) if isinstance(v, (dict, list)) else str(v)
for v in values
]
return await client.lpush(self._make_key(key), *serialized)
async def rpush(self, key: str, *values: Any) -> int:
"""列表右侧插入"""
client = await RedisClient.get_client()
serialized = [
json.dumps(v, ensure_ascii=False, default=str) if isinstance(v, (dict, list)) else str(v)
for v in values
]
return await client.rpush(self._make_key(key), *serialized)
async def lrange(self, key: str, start: int = 0, end: int = -1) -> List[Any]:
"""获取列表范围"""
client = await RedisClient.get_client()
values = await client.lrange(self._make_key(key), start, end)
result = []
for v in values:
try:
result.append(json.loads(v))
except json.JSONDecodeError:
result.append(v)
return result
# 默认缓存管理器实例
cache = CacheManager()
async def get_redis() -> Redis:
"""FastAPI依赖注入用:获取Redis客户端"""
return await RedisClient.get_client()
+43
View File
@@ -0,0 +1,43 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""敏感信息加解密(数据库连接密码等)"""
import base64
import hashlib
import logging
from app.config import settings
logger = logging.getLogger(__name__)
_fernet = None
def _get_fernet():
global _fernet
if _fernet is not None:
return _fernet
try:
from cryptography.fernet import Fernet
except ImportError as exc:
raise RuntimeError("cryptography package is required for secret encryption") from exc
secret = getattr(settings, "DB_CONN_SECRET_KEY", None) or settings.JWT_SECRET_KEY
key = base64.urlsafe_b64encode(hashlib.sha256(secret.encode("utf-8")).digest())
_fernet = Fernet(key)
return _fernet
def encrypt_secret(value: str) -> str:
if not value:
return ""
return _get_fernet().encrypt(value.encode("utf-8")).decode("utf-8")
def decrypt_secret(value: str) -> str:
if not value:
return ""
try:
return _get_fernet().decrypt(value.encode("utf-8")).decode("utf-8")
except Exception as exc:
logger.warning("Failed to decrypt secret: %s", exc)
return ""
+321
View File
@@ -0,0 +1,321 @@
#!/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
@@ -0,0 +1,37 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""将 :name 命名参数 SQL 编译为带字面量的方言 SQL,供外部库驱动执行"""
from typing import Any, Dict
from sqlalchemy import bindparam, text
from sqlalchemy.dialects import mysql, mssql, oracle, postgresql
_DIALECTS = {
"postgresql": postgresql.dialect(),
"postgres": postgresql.dialect(),
"mysql": mysql.dialect(),
"sqlserver": mssql.dialect(),
"mssql": mssql.dialect(),
"oracle": oracle.dialect(),
}
def compile_sql_with_named_params(
sql: str,
params: Dict[str, Any],
db_type: str,
) -> str:
"""
将 :param 绑定为字面量后返回可执行 SQL 字符串。
注意:仅用于已校验过的参数(表单/数据源内部使用),勿用于未过滤的用户原始 SQL。
"""
dialect = _DIALECTS.get((db_type or "postgresql").lower(), postgresql.dialect())
stmt = text(sql)
if params:
stmt = stmt.bindparams(
*[bindparam(key, value=value) for key, value in params.items()]
)
return str(
stmt.compile(dialect=dialect, compile_kwargs={"literal_binds": True})
)
+114
View File
@@ -0,0 +1,114 @@
"""
用户信息缓存模块
将用户的动态权限信息(role_ids, dept_id, is_superuser)缓存到 Redis
中间件从 Redis 获取而非从 JWT token 中读取,确保角色变更后权限立即生效。
缓存 key: user_info:{user_id}
缓存内容: {"role_ids": [...], "dept_id": "...", "is_superuser": true/false}
默认 TTL: 300秒(5分钟),缓存未命中时从数据库加载
"""
import json
from typing import Optional, Dict, Any, List
from utils.redis import RedisClient
USER_INFO_CACHE_PREFIX = "user_info:"
USER_INFO_CACHE_TTL = 300 # 5分钟
async def get_cached_user_info(user_id: str) -> Optional[Dict[str, Any]]:
"""
从 Redis 获取用户信息缓存
:param user_id: 用户ID
:return: 用户信息字典,未命中返回 None
"""
redis = await RedisClient.get_client()
value = await redis.get(f"{USER_INFO_CACHE_PREFIX}{user_id}")
if value:
try:
return json.loads(value)
except json.JSONDecodeError:
return None
return None
async def set_cached_user_info(
user_id: str,
role_ids: List[str],
dept_id: Optional[str],
is_superuser: bool
) -> None:
"""
写入用户信息缓存到 Redis
:param user_id: 用户ID
:param role_ids: 角色ID列表
:param dept_id: 部门ID
:param is_superuser: 是否超级管理员
"""
redis = await RedisClient.get_client()
data = json.dumps({
"role_ids": role_ids,
"dept_id": dept_id,
"is_superuser": is_superuser,
}, ensure_ascii=False)
await redis.set(
f"{USER_INFO_CACHE_PREFIX}{user_id}",
data,
ex=USER_INFO_CACHE_TTL
)
async def delete_cached_user_info(user_id: str) -> None:
"""
删除用户信息缓存(角色/部门等变更时调用)
:param user_id: 用户ID
"""
redis = await RedisClient.get_client()
await redis.delete(f"{USER_INFO_CACHE_PREFIX}{user_id}")
async def delete_all_cached_user_info() -> None:
"""
删除所有用户信息缓存(角色权限配置变更时调用)
"""
redis = await RedisClient.get_client()
cursor = 0
while True:
cursor, keys = await redis.scan(cursor, match=f"{USER_INFO_CACHE_PREFIX}*", count=100)
if keys:
await redis.delete(*keys)
if cursor == 0:
break
async def load_user_info_from_db(user_id: str) -> Optional[Dict[str, Any]]:
"""
从数据库加载用户信息并写入缓存
:param user_id: 用户ID
:return: 用户信息字典,用户不存在返回 None
"""
from app.database import AsyncSessionLocal
from core.user.service import UserService
async with AsyncSessionLocal() as db:
user = await UserService.get_by_id(db, user_id)
if not user:
return None
role_ids = await UserService.get_user_role_ids(db, user_id)
info = {
"role_ids": role_ids,
"dept_id": user.dept_id,
"is_superuser": user.is_superuser,
}
# 写入缓存
await set_cached_user_info(user_id, role_ids, user.dept_id, user.is_superuser)
return info