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
+2
View File
@@ -34,6 +34,8 @@ def add_app_middlewares(app: FastAPI):
async def add_oauth_middleware(request: Request, call_next): async def add_oauth_middleware(request: Request, call_next):
if not env.jwt_global_enable or request.method == "OPTIONS": if not env.jwt_global_enable or request.method == "OPTIONS":
# 没有开启全局的接口认证功能 # 没有开启全局的接口认证功能
request.state.user = None
request.state.token = None
return await call_next(request) return await call_next(request)
# 判断接口是否为认证白名单中的接口 # 判断接口是否为认证白名单中的接口
+71 -16
View File
@@ -1,11 +1,13 @@
import asyncio import asyncio
import datetime
import json import json
from typing import Type, List, Any, Union from typing import Type, List, Any, Union
from fastapi import FastAPI, APIRouter, HTTPException, Body from fastapi import FastAPI, APIRouter, HTTPException, Body
from pydantic import create_model from pydantic import create_model, BaseModel
from sqlalchemy import func from sqlalchemy import func
from sqlmodel import select from sqlmodel import select
from starlette.requests import Request
from app.model.BasicModel import BasicModel from app.model.BasicModel import BasicModel
from app.utils.PageQueryParams import PageQueryParams 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 from app.utils.next_id import next_id
class ModelPublicUser(BaseModel):
id: str
def create_model_service( def create_model_service(
#/*@formatter:off*/ #/*@formatter:off*/
Cls: Type[BasicModel], # model实体类 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()}" 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( def add_route(
self, self,
app: FastAPI, # FastAPI实例,用来注册路由 app: FastAPI, # FastAPI实例,用来注册路由
@@ -116,10 +141,11 @@ def create_model_service(
@router.post("/insert", response_model=ItemResponse) @router.post("/insert", response_model=ItemResponse)
async def _insert( async def _insert(
session: AsyncSessionDep, session: AsyncSessionDep,
row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}") request: Request,
row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}"),
): ):
# 调用item_insert方法执行插入并返回结果 # 调用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"端点,注册批量插入接口 # 若启用"batch_insert"端点,注册批量插入接口
if 'batch_insert' in end_points: if 'batch_insert' in end_points:
@@ -127,10 +153,11 @@ def create_model_service(
@router.post("/batch_insert", response_model=BatchResponse) @router.post("/batch_insert", response_model=BatchResponse)
async def _batch_insert( async def _batch_insert(
session: AsyncSessionDep, 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方法执行批量插入并返回结果 # 调用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"端点,注册单条更新接口 # 若启用"update"端点,注册单条更新接口
if 'update' in end_points: if 'update' in end_points:
@@ -138,10 +165,11 @@ def create_model_service(
@router.post("/update", response_model=ItemResponse) @router.post("/update", response_model=ItemResponse)
async def _update( async def _update(
session: AsyncSessionDep, session: AsyncSessionDep,
row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}") row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}"),
request: Request = None,
): ):
# 调用item_update方法执行更新并返回结果 # 调用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"端点,注册批量更新接口 # 若启用"batch_update"端点,注册批量更新接口
if 'batch_update' in end_points: if 'batch_update' in end_points:
@@ -149,10 +177,11 @@ def create_model_service(
@router.post("/batch_update", response_model=BatchResponse) @router.post("/batch_update", response_model=BatchResponse)
async def _batch_update( async def _batch_update(
session: AsyncSessionDep, 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方法执行批量更新并返回结果 # 调用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"端点,注册单条删除接口 # 若启用"delete"端点,注册单条删除接口
if 'delete' in end_points: if 'delete' in end_points:
@@ -274,13 +303,16 @@ def create_model_service(
return item_cls 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: if before_insert is not None:
await before_insert(row_dict, session) await before_insert(row_dict, session)
# 若未提供id,自动生成唯一id await self.check_insert_row_dict(row_dict, user)
if row_dict.get("id") is None:
row_dict["id"] = await next_id()
try: try:
# 使用模型类验证数据并创建实例(校验字段类型和约束) # 使用模型类验证数据并创建实例(校验字段类型和约束)
@@ -303,7 +335,12 @@ def create_model_service(
return insert_cls 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: if before_batch_insert is not None:
await before_batch_insert(row_dict_list, session) await before_batch_insert(row_dict_list, session)
@@ -323,6 +360,9 @@ def create_model_service(
for index, id in enumerate(new_id_list): for index, id in enumerate(new_id_list):
row_dict_list_without_id[index]["id"] = id 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: try:
# 验证所有记录并转换为模型实例列表 # 验证所有记录并转换为模型实例列表
insert_cls_list = [Cls.model_validate(row_dict) for row_dict in row_dict_list] 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 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: if before_update is not None:
await before_update(row_dict, session) await before_update(row_dict, session)
await self.check_update_row_dict(row_dict, user)
# 检查id是否存在(更新必须指定id) # 检查id是否存在(更新必须指定id)
if not row_dict.get('id'): if not row_dict.get('id'):
raise HTTPException(status_code=400, detail="ID不能为空") raise HTTPException(status_code=400, detail="ID不能为空")
@@ -379,11 +426,19 @@ def create_model_service(
return update_cls 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: if before_batch_update is not None:
await before_batch_update(row_dict_list, session) 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 # 提取所有待更新记录的id
update_id_list = [row_dict['id'] for row_dict in row_dict_list] update_id_list = [row_dict['id'] for row_dict in row_dict_list]
# 根据id查询所有待更新的记录 # 根据id查询所有待更新的记录