feat: 封装函数
self.select_cls()
This commit is contained in:
@@ -179,6 +179,13 @@ def create_model_service(
|
|||||||
# 将路由添加到FastAPI应用
|
# 将路由添加到FastAPI应用
|
||||||
app.include_router(router)
|
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):
|
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)
|
await before_query_list(query_param, session)
|
||||||
|
|
||||||
# 创建基础查询:查询当前模型类的所有记录
|
# 创建基础查询:查询当前模型类的所有记录
|
||||||
if custom_query is not None:
|
query = self.select_cls()
|
||||||
query = custom_query()
|
|
||||||
else:
|
|
||||||
query = select(Cls)
|
|
||||||
count_query = select(func.count()).select_from(Cls)
|
count_query = select(func.count()).select_from(Cls)
|
||||||
|
|
||||||
# 若有过滤条件,验证并应用过滤
|
# 若有过滤条件,验证并应用过滤
|
||||||
@@ -244,7 +248,7 @@ def create_model_service(
|
|||||||
await before_query_item(row_dict, session)
|
await before_query_item(row_dict, session)
|
||||||
|
|
||||||
# 创建基础查询:查询当前模型类的所有记录
|
# 创建基础查询:查询当前模型类的所有记录
|
||||||
query = select(Cls)
|
query = self.select_cls()
|
||||||
|
|
||||||
# 验证查询条件中的键是否有效
|
# 验证查询条件中的键是否有效
|
||||||
self.check_invalid_keys(row_dict)
|
self.check_invalid_keys(row_dict)
|
||||||
@@ -325,7 +329,7 @@ def create_model_service(
|
|||||||
await session.commit()
|
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:
|
if after_batch_insert is not None:
|
||||||
await after_batch_insert(refresh_cls_list, row_dict_list, session)
|
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'):
|
if not row_dict.get('id'):
|
||||||
raise HTTPException(status_code=400, detail="ID不能为空")
|
raise HTTPException(status_code=400, detail="ID不能为空")
|
||||||
# 根据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:
|
if not update_cls:
|
||||||
# 若记录不存在,抛出异常
|
# 若记录不存在,抛出异常
|
||||||
raise HTTPException(status_code=500, detail="Update row not found")
|
raise HTTPException(status_code=500, detail="Update row not found")
|
||||||
@@ -374,7 +378,7 @@ def create_model_service(
|
|||||||
# 提取所有待更新记录的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查询所有待更新的记录
|
||||||
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不存在
|
# 若查询到的记录数量与待更新数量不一致,说明部分id不存在
|
||||||
if len(update_cls_list) != len(row_dict_list):
|
if len(update_cls_list) != len(row_dict_list):
|
||||||
# 抛出异常并提示不存在的id
|
# 抛出异常并提示不存在的id
|
||||||
@@ -397,7 +401,7 @@ def create_model_service(
|
|||||||
# 提交事务
|
# 提交事务
|
||||||
await session.commit()
|
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:
|
if after_batch_update is not None:
|
||||||
await after_batch_update(refresh_cls_list, row_dict_list, session)
|
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)
|
await before_delete(row_dict, session)
|
||||||
|
|
||||||
# 根据id查询要删除的记录
|
# 根据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:
|
if not delete_cls:
|
||||||
# 若记录不存在,返回删除失败
|
# 若记录不存在,返回删除失败
|
||||||
return False
|
return False
|
||||||
@@ -441,7 +445,7 @@ def create_model_service(
|
|||||||
row_id_list = [row_dict.get("id") for row_dict in row_dict_list]
|
row_id_list = [row_dict.get("id") for row_dict in row_dict_list]
|
||||||
|
|
||||||
# 根据id查询所有待删除的记录
|
# 根据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不存在
|
# 若查询到的记录数量与待删除数量不一致,说明部分id不存在
|
||||||
if len(delete_cls_list) != len(row_id_list):
|
if len(delete_cls_list) != len(row_id_list):
|
||||||
# 抛出异常并提示不存在的id
|
# 抛出异常并提示不存在的id
|
||||||
|
|||||||
Reference in New Issue
Block a user