1
0
forked from erp-dev/erp
Files
erpnew/flower/middleware.py
2026-01-12 13:45:57 +08:00

197 lines
6.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
自定义中间件模块
"""
import json
import logging
from django.conf import settings
from rest_framework_simplejwt.authentication import JWTAuthentication
logger = logging.getLogger(__name__)
class ApiAuditLogMiddleware:
"""
API审计日志中间件
用于记录特定URL前缀的POST请求保存"创建"操作的历史现场。
特性:
- 只记录POST请求
- 通过URL前缀白名单过滤
- 支持multipart/form-data请求文件字段只记录元信息
- 通过Celery异步写入不阻塞API响应
- 在AuthenticationMiddleware之后运行可获取request.user
配置项settings.py
- API_AUDIT_LOG_ENABLED: 是否启用默认True
- API_AUDIT_LOG_URL_PREFIXES: URL前缀白名单列表
"""
def __init__(self, get_response):
self.get_response = get_response
self.enabled = getattr(settings, 'API_AUDIT_LOG_ENABLED', True)
self.url_prefixes = getattr(settings, 'API_AUDIT_LOG_URL_PREFIXES', [])
def __call__(self, request):
# 检查是否启用
if not self.enabled:
return self.get_response(request)
# 只处理POST请求
if request.method != 'POST':
return self.get_response(request)
# 检查URL是否在白名单中
if not self._is_url_matched(request.path):
logger.debug('ApiAuditLog: URL不匹配白名单, path=%s, prefixes=%s', request.path, self.url_prefixes)
return self.get_response(request)
logger.info('ApiAuditLog: 捕获到POST请求, path=%s', request.path)
# 在请求进入视图前,缓存请求体数据
# 注意request.body只能读取一次需要在这里缓存
request_data = self._extract_request_data(request)
query_params = dict(request.GET)
# 执行视图
response = self.get_response(request)
# 获取用户信息此时DRF的JWT认证已完成
user_id, username = self._get_user_info(request)
logger.info('ApiAuditLog: 准备发送任务, user=%s, status=%s', username, response.status_code)
# 异步保存审计日志
self._save_audit_log(
url=request.path,
method=request.method,
request_data=request_data,
query_params=query_params,
user_id=user_id,
username=username,
response_status=response.status_code,
)
return response
def _is_url_matched(self, path: str) -> bool:
"""检查URL是否匹配白名单前缀"""
for prefix in self.url_prefixes:
if path.startswith(prefix):
return True
return False
def _extract_request_data(self, request) -> dict:
"""
提取请求数据
对于multipart/form-data请求分别处理表单字段和文件字段。
文件字段只记录元信息(文件名、大小、类型),不记录二进制内容。
"""
content_type = request.content_type or ''
if 'multipart/form-data' in content_type:
# multipart请求分别处理表单字段和文件
data = {}
# 处理普通表单字段
for key, values in request.POST.lists():
if len(values) == 1:
data[key] = values[0]
else:
data[key] = values
# 处理文件字段:只记录元信息
for key, files in request.FILES.lists():
file_infos = []
for f in files:
file_infos.append({
'_type': 'file',
'name': f.name,
'size': f.size,
'content_type': f.content_type,
})
if len(file_infos) == 1:
data[key] = file_infos[0]
else:
data[key] = file_infos
return data
elif 'application/json' in content_type:
# JSON请求解析body
try:
return json.loads(request.body.decode('utf-8'))
except (json.JSONDecodeError, UnicodeDecodeError):
return {'_raw': request.body.decode('utf-8', errors='replace')}
elif 'application/x-www-form-urlencoded' in content_type:
# 表单请求
data = {}
for key, values in request.POST.lists():
if len(values) == 1:
data[key] = values[0]
else:
data[key] = values
return data
else:
# 其他类型:尝试作为文本记录
try:
body = request.body.decode('utf-8')
if body:
return {'_raw': body}
return {}
except UnicodeDecodeError:
return {'_raw': '<binary data>'}
def _get_user_info(self, request) -> tuple[int | None, str]:
"""
获取用户信息
优先从request.user获取Session认证或DRF已设置
如果未认证尝试手动进行JWT认证。
"""
# 尝试从request.user获取
if hasattr(request, 'user') and request.user.is_authenticated:
return request.user.id, request.user.username
# 尝试手动JWT认证
try:
auth = JWTAuthentication()
user_auth_tuple = auth.authenticate(request)
if user_auth_tuple:
user = user_auth_tuple[0]
return user.id, user.username
except Exception:
pass
return None, ''
def _save_audit_log(
self,
url: str,
method: str,
request_data: dict,
query_params: dict,
user_id: int | None,
username: str,
response_status: int | None,
):
"""通过Celery异步保存审计日志"""
try:
from api_v1.tasks import save_api_audit_log
save_api_audit_log.delay(
url=url,
method=method,
request_data=request_data,
query_params=query_params,
user_id=user_id,
username=username,
response_status=response_status,
)
except Exception as e:
# 日志保存失败不应影响正常请求
logger.error('发送API审计日志任务失败: %s', str(e), exc_info=True)