Files
H3ConuMS-v2/backend/app/middleware/audit_middleware.py
T
v6ole f1f8518985 ```
feat(auth): 添加用户权限获取接口并完善JWT令牌角色信息

- 在JWT令牌中添加用户角色信息
- 新增get_my_permissions接口用于获取当前用户权限码列表
- 重构认证回调逻辑,增加错误日志记录
- 更新用户信息获取接口使用Authorization头验证
```
2026-04-06 00:40:08 +08:00

110 lines
3.6 KiB
Python

"""审计日志中间件:拦截所有 API 请求,异步写入审计日志"""
import json
import time
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from starlette.responses import Response
from app.core.security import verify_token
# 不记录审计日志的路径前缀
_SKIP_PATHS = {
"/health",
"/docs",
"/redoc",
"/openapi.json",
"/api/auth/login", # 仅获取登录 URL,无用户身份
"/api/auth/permissions", # 高频只读
"/api/stats/",
"/api/olt/regions",
"/api/olt/new-devices",
"/api/olt/duplicate-macs",
}
# 只记录写操作 + 登录回调 + 特定查询(GET 默认跳过,以下 GET 例外)
_ALWAYS_LOG_METHODS = {"POST", "PUT", "DELETE", "PATCH"}
_LOG_GET_PATHS = {
"/api/auth/profile",
}
def _should_log(method: str, path: str) -> bool:
for skip in _SKIP_PATHS:
if path.startswith(skip):
return False
if method in _ALWAYS_LOG_METHODS:
return True
if method == "GET":
return path in _LOG_GET_PATHS
return False
def _extract_token_payload(request: Request) -> dict:
auth = request.headers.get("Authorization", "")
if auth.startswith("Bearer "):
payload = verify_token(auth[7:])
if payload:
return payload
return {}
class AuditMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next) -> Response:
method = request.method
path = request.url.path
if not _should_log(method, path):
return await call_next(request)
# 读取请求体(只读一次,需要重新构造)
request_params = None
try:
body_bytes = await request.body()
if body_bytes:
try:
request_params = json.loads(body_bytes)
# 脱敏:移除密码字段
if isinstance(request_params, dict):
for k in ("password", "passwd", "secret"):
if k in request_params:
request_params[k] = "***"
except Exception:
pass
except Exception:
pass
start_time = time.time()
response = await call_next(request)
duration_ms = int((time.time() - start_time) * 1000)
# 异步写日志(不等待)
try:
payload = _extract_token_payload(request)
user_id = payload.get("sub", "anonymous")
username = payload.get("username", "anonymous")
user_role = payload.get("role", "")
ip_address = ""
forwarded = request.headers.get("x-forwarded-for")
if forwarded:
ip_address = forwarded.split(",")[0].strip()
elif request.client:
ip_address = request.client.host
from app.tasks.audit_tasks import create_audit_log_task
create_audit_log_task.delay(
user_id=str(user_id),
username=username,
user_role=user_role,
method=method,
path=path,
ip_address=ip_address,
user_agent=request.headers.get("user-agent", "")[:500],
status_code=response.status_code,
request_params=request_params,
response_data=None, # 不捕获响应体(性能考虑)
error_message=None if response.status_code < 400 else f"HTTP {response.status_code}",
)
except Exception:
pass # 中间件异常绝不影响主响应
return response