feat: KnowledgeDocModel add field content

This commit is contained in:
martsforever
2025-09-08 15:10:28 +08:00
parent 823e43e8f7
commit 2bfc3bc0e0
2 changed files with 73 additions and 16 deletions
+71 -16
View File
@@ -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查询所有待更新的记录