forked from erp-dev/erp
201 lines
6.8 KiB
Python
201 lines
6.8 KiB
Python
"""
|
||
自定义中间件模块
|
||
"""
|
||
import io
|
||
import json
|
||
import logging
|
||
|
||
from django.conf import settings
|
||
from django.http import QueryDict
|
||
from django.http.multipartparser import MultiPartParserError
|
||
from django.utils.datastructures import MultiValueDict
|
||
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)
|