Build lightweight AI agent admin
This commit is contained in:
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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:
|
||||
# 存储权限ID,key为 "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
|
||||
@@ -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()
|
||||
@@ -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 ""
|
||||
@@ -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})
|
||||
)
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user