From 867f1f9e33781f58e619794fe363c97ed766ceab Mon Sep 17 00:00:00 2001 From: martsforever Date: Sat, 23 Aug 2025 01:46:16 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=B0=81=E8=A3=85=E5=87=BD=E6=95=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit self.select_cls() --- app/utils/create_module_service.py | 26 +++++++++++++++----------- 1 file changed, 15 insertions(+), 11 deletions(-) diff --git a/app/utils/create_module_service.py b/app/utils/create_module_service.py index 7d30c86..d572842 100644 --- a/app/utils/create_module_service.py +++ b/app/utils/create_module_service.py @@ -179,6 +179,13 @@ def create_model_service( # 将路由添加到FastAPI应用 app.include_router(router) + def select_cls(self): + if custom_query is not None: + query = custom_query() + else: + query = select(Cls) + return query + # 分页查询工具方法:执行带过滤和分页的查询 async def query_list(self, query_param: PageQueryParams, session: AsyncSessionDep): @@ -186,10 +193,7 @@ def create_model_service( await before_query_list(query_param, session) # 创建基础查询:查询当前模型类的所有记录 - if custom_query is not None: - query = custom_query() - else: - query = select(Cls) + query = self.select_cls() count_query = select(func.count()).select_from(Cls) # 若有过滤条件,验证并应用过滤 @@ -244,7 +248,7 @@ def create_model_service( await before_query_item(row_dict, session) # 创建基础查询:查询当前模型类的所有记录 - query = select(Cls) + query = self.select_cls() # 验证查询条件中的键是否有效 self.check_invalid_keys(row_dict) @@ -325,7 +329,7 @@ def create_model_service( await session.commit() # 查询并返回所有插入的实例(刷新数据,确保获取最新状态) - refresh_cls_list = (await session.execute(select(Cls).where(Cls.id.in_([obj.id for obj in insert_cls_list])))).scalars().all() + refresh_cls_list = (await session.execute(self.select_cls().where(Cls.id.in_([obj.id for obj in insert_cls_list])))).scalars().all() if after_batch_insert is not None: await after_batch_insert(refresh_cls_list, row_dict_list, session) @@ -343,7 +347,7 @@ def create_model_service( if not row_dict.get('id'): raise HTTPException(status_code=400, detail="ID不能为空") # 根据id查询要更新的记录 - update_cls = (await session.exec(select(Cls).where(Cls.id == row_dict.get('id')))).first() + update_cls = (await session.exec(self.select_cls().where(Cls.id == row_dict.get('id')))).first() if not update_cls: # 若记录不存在,抛出异常 raise HTTPException(status_code=500, detail="Update row not found") @@ -374,7 +378,7 @@ def create_model_service( # 提取所有待更新记录的id update_id_list = [row_dict['id'] for row_dict in row_dict_list] # 根据id查询所有待更新的记录 - update_cls_list = (await session.exec(select(Cls).where(Cls.id.in_(update_id_list)))).all() + update_cls_list = (await session.exec(self.select_cls().where(Cls.id.in_(update_id_list)))).all() # 若查询到的记录数量与待更新数量不一致,说明部分id不存在 if len(update_cls_list) != len(row_dict_list): # 抛出异常并提示不存在的id @@ -397,7 +401,7 @@ def create_model_service( # 提交事务 await session.commit() # 查询并返回所有更新后的实例(刷新数据) - refresh_cls_list = (await session.execute(select(Cls).where(Cls.id.in_([obj.id for obj in update_cls_list])))).scalars().all() + refresh_cls_list = (await session.execute(self.select_cls().where(Cls.id.in_([obj.id for obj in update_cls_list])))).scalars().all() if after_batch_update is not None: await after_batch_update(refresh_cls_list, row_dict_list, session) @@ -412,7 +416,7 @@ def create_model_service( await before_delete(row_dict, session) # 根据id查询要删除的记录 - delete_cls = (await session.exec(select(Cls).where(Cls.id == row_dict.get('id')))).first() + delete_cls = (await session.exec(self.select_cls().where(Cls.id == row_dict.get('id')))).first() if not delete_cls: # 若记录不存在,返回删除失败 return False @@ -441,7 +445,7 @@ def create_model_service( row_id_list = [row_dict.get("id") for row_dict in row_dict_list] # 根据id查询所有待删除的记录 - delete_cls_list = (await session.exec(select(Cls).where(Cls.id.in_(row_id_list)))).all() + delete_cls_list = (await session.exec(self.select_cls().where(Cls.id.in_(row_id_list)))).all() # 若查询到的记录数量与待删除数量不一致,说明部分id不存在 if len(delete_cls_list) != len(row_id_list): # 抛出异常并提示不存在的id