From 2bfc3bc0e083a2087327a0aaf9fd7bb93c081d85 Mon Sep 17 00:00:00 2001 From: martsforever Date: Mon, 8 Sep 2025 15:10:28 +0800 Subject: [PATCH] feat: KnowledgeDocModel add field content --- app/middlewares/app_middlewares.py | 2 + app/utils/create_module_service.py | 87 ++++++++++++++++++++++++------ 2 files changed, 73 insertions(+), 16 deletions(-) diff --git a/app/middlewares/app_middlewares.py b/app/middlewares/app_middlewares.py index 9fa6fd0..6229f8e 100644 --- a/app/middlewares/app_middlewares.py +++ b/app/middlewares/app_middlewares.py @@ -34,6 +34,8 @@ def add_app_middlewares(app: FastAPI): async def add_oauth_middleware(request: Request, call_next): if not env.jwt_global_enable or request.method == "OPTIONS": # 没有开启全局的接口认证功能 + request.state.user = None + request.state.token = None return await call_next(request) # 判断接口是否为认证白名单中的接口 diff --git a/app/utils/create_module_service.py b/app/utils/create_module_service.py index 0acc1c7..5713529 100644 --- a/app/utils/create_module_service.py +++ b/app/utils/create_module_service.py @@ -1,11 +1,13 @@ import asyncio +import datetime import json from typing import Type, List, Any, Union from fastapi import FastAPI, APIRouter, HTTPException, Body -from pydantic import create_model +from pydantic import create_model, BaseModel from sqlalchemy import func from sqlmodel import select +from starlette.requests import Request from app.model.BasicModel import BasicModel from app.utils.PageQueryParams import PageQueryParams @@ -13,6 +15,10 @@ from app.utils.db_utils import AsyncSessionDep from app.utils.next_id import next_id +class ModelPublicUser(BaseModel): + id: str + + def create_model_service( #/*@formatter:off*/ Cls: Type[BasicModel], # model实体类 @@ -63,6 +69,25 @@ def create_model_service( detail=f"Invalid filter keys: {invalid_keys}. Valid keys are: {Cls.__annotations__.keys()}" ) + # 检查要新建的行数据 + async def check_insert_row_dict(self, row_dict: dict, user: ModelPublicUser): + # 若未提供id,自动生成唯一id + if row_dict.get("id", None) is None: + row_dict["id"] = await next_id() + + if user: + if not row_dict.get("created_by", None): + row_dict["created_by"] = user.id + + if not row_dict.get("updated_by", None): + row_dict["updated_by"] = user.id + + # 检查要更新的行数据 + async def check_update_row_dict(self, row_dict: dict, user: ModelPublicUser): + if user: + row_dict["updated_by"] = user.id + row_dict["updated_at"] = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") + def add_route( self, app: FastAPI, # FastAPI实例,用来注册路由 @@ -116,10 +141,11 @@ def create_model_service( @router.post("/insert", response_model=ItemResponse) async def _insert( session: AsyncSessionDep, - row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}") + request: Request, + row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}"), ): # 调用item_insert方法执行插入并返回结果 - return {"result": await self.item_insert(session, row_dict)} + return {"result": await self.item_insert(session, row_dict, request.state.user)} # 若启用"batch_insert"端点,注册批量插入接口 if 'batch_insert' in end_points: @@ -127,10 +153,11 @@ def create_model_service( @router.post("/batch_insert", response_model=BatchResponse) async def _batch_insert( session: AsyncSessionDep, - row_dict_list: List[dict] = Body(..., description=f"批量插入的数据数组,字段参考{Cls.__name__}") + row_dict_list: List[dict] = Body(..., description=f"批量插入的数据数组,字段参考{Cls.__name__}"), + request: Request = None, ): # 调用batch_insert方法执行批量插入并返回结果 - return {"result": await self.batch_insert(session, row_dict_list)} + return {"result": await self.batch_insert(session, row_dict_list, request.state.user)} # 若启用"update"端点,注册单条更新接口 if 'update' in end_points: @@ -138,10 +165,11 @@ def create_model_service( @router.post("/update", response_model=ItemResponse) async def _update( session: AsyncSessionDep, - row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}") + row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}"), + request: Request = None, ): # 调用item_update方法执行更新并返回结果 - return {"result": await self.item_update(session, row_dict)} + return {"result": await self.item_update(session, row_dict, request.state.user)} # 若启用"batch_update"端点,注册批量更新接口 if 'batch_update' in end_points: @@ -149,10 +177,11 @@ def create_model_service( @router.post("/batch_update", response_model=BatchResponse) async def _batch_update( session: AsyncSessionDep, - row_dict_list: List[dict] = Body(..., description=f"批量更新的数据数组,字段参考{Cls.__name__}") + row_dict_list: List[dict] = Body(..., description=f"批量更新的数据数组,字段参考{Cls.__name__}"), + request: Request = None, ): # 调用batch_update方法执行批量更新并返回结果 - return {"result": await self.batch_update(session, row_dict_list)} + return {"result": await self.batch_update(session, row_dict_list, request.state.user)} # 若启用"delete"端点,注册单条删除接口 if 'delete' in end_points: @@ -274,13 +303,16 @@ def create_model_service( return item_cls # 单条插入工具方法:新增一条记录 - async def item_insert(self, session: AsyncSessionDep, row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}")): + async def item_insert( + self, + session: AsyncSessionDep, + row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}"), + user: ModelPublicUser = None, + ): if before_insert is not None: await before_insert(row_dict, session) - # 若未提供id,自动生成唯一id - if row_dict.get("id") is None: - row_dict["id"] = await next_id() + await self.check_insert_row_dict(row_dict, user) try: # 使用模型类验证数据并创建实例(校验字段类型和约束) @@ -303,7 +335,12 @@ def create_model_service( return insert_cls # 批量插入工具方法:批量新增记录 - async def batch_insert(self, session: AsyncSessionDep, row_dict_list: List[dict] = Body(..., description=f"批量插入的数据数组,字段参考{Cls.__name__}")): + async def batch_insert( + self, + session: AsyncSessionDep, + row_dict_list: List[dict] = Body(..., description=f"批量插入的数据数组,字段参考{Cls.__name__}"), + user: ModelPublicUser = None, + ): if before_batch_insert is not None: await before_batch_insert(row_dict_list, session) @@ -323,6 +360,9 @@ def create_model_service( for index, id in enumerate(new_id_list): row_dict_list_without_id[index]["id"] = id + for row_dict in row_dict_list_without_id: + await self.check_insert_row_dict(row_dict, user) + try: # 验证所有记录并转换为模型实例列表 insert_cls_list = [Cls.model_validate(row_dict) for row_dict in row_dict_list] @@ -346,11 +386,18 @@ def create_model_service( return refresh_cls_list # 单条更新工具方法:更新一条记录 - async def item_update(self, session: AsyncSessionDep, row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}")): + async def item_update( + self, + session: AsyncSessionDep, + row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}"), + user: ModelPublicUser = None + ): if before_update is not None: await before_update(row_dict, session) + await self.check_update_row_dict(row_dict, user) + # 检查id是否存在(更新必须指定id) if not row_dict.get('id'): raise HTTPException(status_code=400, detail="ID不能为空") @@ -379,11 +426,19 @@ def create_model_service( return update_cls # 批量更新工具方法:批量更新记录 - async def batch_update(self, session: AsyncSessionDep, row_dict_list: List[dict] = Body(..., description=f"批量更新的数据数组,字段参考{Cls.__name__}")): + async def batch_update( + self, + session: AsyncSessionDep, + row_dict_list: List[dict] = Body(..., description=f"批量更新的数据数组,字段参考{Cls.__name__}"), + user: ModelPublicUser = None + ): if before_batch_update is not None: await before_batch_update(row_dict_list, session) + for row_dict in row_dict_list: + await self.check_update_row_dict(row_dict, user) + # 提取所有待更新记录的id update_id_list = [row_dict['id'] for row_dict in row_dict_list] # 根据id查询所有待更新的记录