Files
FastapiAdmin/backend/app/api/v1/module_generator/gencode/service.py
T
zhangtao ba7fddc34a refactor(generator): 优化代码生成模块接口和服务实现
- 修改接口定义,增加路径参数并完善请求描述,增强参数校验和依赖注入
- 优化CRUD层数据库操作,统一异步会话使用,删除多余db参数
- 增加业务表与字段模型关系级联删除配置,优化模型关联关系声明
- 精简pydantic模型,去除冗余校验装饰器,完善字段描述和必填约束
- 服务层增加类型检查和异常抛出,规范业务逻辑流程和错误提示
- 优化代码结构,调整模块导入顺序和注释,提升代码可读性和一致性
2025-10-01 16:58:49 +08:00

522 lines
25 KiB
Python

# -*- coding:utf-8 -*-
import io
import json
import os
import zipfile
from datetime import datetime
from typing import Any, List, Dict, Optional, Sequence
from sqlalchemy.ext.asyncio import AsyncSession
from app.config.setting import settings
from app.core.base_model import CamelCaseUtil
from app.core.exceptions import CustomException
from app.utils.gen_util import GenUtils
from app.utils.template_util import TemplateInitializer, TemplateUtils
from app.common.constant import GenConstant
from app.common.response import SuccessResponse
from app.api.v1.module_system.user.schema import UserOutSchema
from app.api.v1.module_system.auth.schema import AuthSchema
from .schema import GenTableCreateSchema, GenTableUpdateSchema, GenTableOutSchema, GenTableDeleteSchema, GenTableColumnCreateSchema, GenTableColumnUpdateSchema, GenTableColumnOutSchema, GenTableColumnDeleteSchema
from .param import GenTableQueryParam
from .crud import GenTableColumnCRUD, GenTableCRUD
from .model import GenTableModel, GenTableColumnModel
# 定义默认的GenConfig值
GEN_PATH = "generated_code" # 默认生成路径
class GenTableService:
"""代码生成业务表服务层"""
@classmethod
async def get_gen_table_detail_service(cls, auth: AuthSchema, table_id: int) -> Dict:
"""获取业务表详细信息"""
gen_table = await cls.get_gen_table_by_id_service(auth, table_id)
gen_tables = await cls.get_gen_table_all_service(auth)
gen_columns = await GenTableColumnService.get_gen_table_column_list_by_table_id_service(auth, table_id)
return dict(info=gen_table, rows=gen_columns, tables=gen_tables)
@classmethod
async def get_gen_table_list_service(
cls, auth: AuthSchema, query_object: GenTableQueryParam, is_page: bool = False
) -> Dict:
"""获取代码生成业务表列表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_table_list_result = await gen_table_dao.get_gen_table_list(db, query_object, is_page)
return gen_table_list_result
@classmethod
async def get_gen_db_table_list_service(
cls, auth: AuthSchema, query_object: GenTableQueryParam, is_page: bool = False
) -> Dict:
"""获取数据库列表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_db_table_list_result = await gen_table_dao.get_gen_db_table_list(db, query_object, is_page)
return gen_db_table_list_result
@classmethod
async def get_gen_db_table_list_by_name_service(cls, auth: AuthSchema, table_names: List[str]) -> List[GenTableOutSchema]:
"""根据表名称组获取数据库列表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_db_table_list_result = await gen_table_dao.get_gen_db_table_list_by_names(db, table_names)
return [GenTableOutSchema(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_db_table_list_result)]
@classmethod
async def import_gen_table_service(
cls, auth: AuthSchema, gen_table_list: List[GenTableOutSchema], current_user: UserOutSchema
) -> SuccessResponse:
"""导入表结构"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
try:
gen_table_dao = GenTableCRUD(auth=auth)
gen_table_column_dao = GenTableColumnCRUD(auth=auth)
for table in gen_table_list:
table_name = table.table_name
GenUtils.init_table(table, current_user.username)
add_gen_table = await gen_table_dao.create(data=table.model_dump())
if add_gen_table:
# 使用id而不是table_id
table.id = add_gen_table.id
gen_table_columns = await gen_table_column_dao.get_gen_db_table_columns_by_name(db, table_name or "")
for column in [
GenTableColumnOutSchema(**gen_table_column)
for gen_table_column in CamelCaseUtil.transform_result(gen_table_columns)
]:
GenUtils.init_column_field(column, table)
await gen_table_column_dao.create(data=column.model_dump())
if isinstance(db, AsyncSession):
await db.commit()
return SuccessResponse(msg='导入成功')
except Exception as e:
if isinstance(db, AsyncSession):
try:
await db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'导入失败, {str(e)}')
@classmethod
async def create_table_service(cls, auth: AuthSchema, sql: str, current_user: UserOutSchema) -> SuccessResponse:
"""创建表结构"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
try:
# 执行SQL语句创建表
await gen_table_dao.create_table_by_sql_dao(db, [sql])
if isinstance(db, AsyncSession):
await db.commit()
return SuccessResponse(msg='创建表结构成功')
except Exception as e:
if isinstance(db, AsyncSession):
try:
await db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'创建表结构失败: {str(e)}')
@classmethod
async def update_gen_table_service(cls, auth: AuthSchema, page_object: GenTableUpdateSchema, table_id: int) -> Dict[str, Any]:
"""编辑业务表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_table_column_dao = GenTableColumnCRUD(auth=auth)
edit_gen_table = page_object.model_dump(exclude_unset=True, by_alias=True)
gen_table_info = await cls.get_gen_table_by_id_service(auth, table_id)
if gen_table_info.id:
try:
# 确保options字段存在且为有效JSON
if 'options' not in edit_gen_table or edit_gen_table['options'] is None:
edit_gen_table['options'] = '{}' # 默认空对象
else:
# 验证options是否为有效的JSON
try:
json.loads(edit_gen_table['options'])
except json.JSONDecodeError:
edit_gen_table['options'] = '{}'
await gen_table_dao.update(id=table_id, data=edit_gen_table)
if hasattr(page_object, 'columns') and page_object.columns:
for gen_table_column in page_object.columns:
# 为列添加更新信息
gen_table_column_dict = gen_table_column.model_dump()
gen_table_column_dict['update_by'] = getattr(page_object, 'update_by', '')
gen_table_column_dict['update_time'] = datetime.now()
# 检查是否有id属性
column_id = getattr(gen_table_column, 'id', None)
if column_id is not None:
await gen_table_column_dao.update(
id=column_id,
data=gen_table_column_dict
)
if isinstance(db, AsyncSession):
await db.commit()
return {"is_success": True, "message": "更新成功"}
except Exception as e:
if isinstance(db, AsyncSession):
try:
await db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'更新失败: {str(e)}')
else:
raise CustomException(msg='业务表不存在')
@classmethod
async def delete_gen_table_service(cls, auth: AuthSchema, page_object: GenTableDeleteSchema) -> SuccessResponse:
"""删除业务表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_table_column_dao = GenTableColumnCRUD(auth=auth)
if page_object.table_ids:
table_id_list = page_object.table_ids.split(',')
try:
for table_id in table_id_list:
await gen_table_dao.delete(ids=[int(table_id)])
# 删除相关的字段信息
# 这里需要先查询出所有相关的column_id,然后删除
columns = await gen_table_column_dao.get_gen_table_column_list_by_table_id(db, int(table_id))
if columns:
column_ids = [column.id for column in columns]
await gen_table_column_dao.delete(ids=column_ids)
if isinstance(db, AsyncSession):
await db.commit()
return SuccessResponse(msg='删除成功')
except Exception as e:
if isinstance(db, AsyncSession):
try:
await db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'删除失败: {str(e)}')
else:
raise CustomException(msg='传入业务表id为空')
@classmethod
async def get_gen_table_by_id_service(cls, auth: AuthSchema, table_id: int) -> GenTableOutSchema:
"""获取需要生成的业务表详细信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_id(db, table_id)
if gen_table:
result = await cls.set_table_from_options(GenTableOutSchema(**CamelCaseUtil.transform_result(gen_table)))
return result
else:
raise CustomException(msg='业务表不存在')
@classmethod
async def get_gen_table_all_service(cls, auth: AuthSchema) -> List[GenTableOutSchema]:
"""获取所有业务表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_tables = await gen_table_dao.get_gen_table_all(db)
result = []
for table in gen_tables:
table_info = await cls.set_table_from_options(GenTableOutSchema(**CamelCaseUtil.transform_result(table)))
result.append(table_info)
return result
@classmethod
async def preview_code_service(cls, auth: AuthSchema, table_id: int) -> Dict[Any, Any]:
"""预览代码"""
gen_table = await cls.get_gen_table_by_id_service(auth, table_id)
await cls.set_sub_table(auth, gen_table)
await cls._set_pk_column(gen_table)
env = TemplateInitializer.init_jinja2()
context = TemplateUtils.prepare_context(gen_table)
template_list = TemplateUtils.get_template_list(
gen_table.tpl_category or "",
gen_table.tpl_web_type or ""
)
preview_code_result = {}
for template in template_list:
render_content = env.get_template(template).render(**context)
preview_code_result[template] = render_content
return preview_code_result
@classmethod
async def generate_code_service(cls, auth: AuthSchema, table_name: str) -> SuccessResponse:
"""生成代码至指定路径"""
env = TemplateInitializer.init_jinja2()
render_info = await cls.__get_gen_render_info(auth, table_name)
for template in render_info[0]:
try:
render_content = env.get_template(template).render(**render_info[2])
gen_path = cls.__get_gen_path(render_info[3], template)
if gen_path:
os.makedirs(os.path.dirname(gen_path), exist_ok=True)
with open(gen_path, 'w', encoding='utf-8') as f:
f.write(render_content)
except Exception as e:
raise CustomException(msg=f'渲染模板失败,表名:{render_info[3].table_name},详细错误信息:{str(e)}')
return SuccessResponse(msg='生成代码成功')
@classmethod
async def batch_gen_code_service(cls, auth: AuthSchema, table_names: List[str]) -> bytes:
"""批量生成代码"""
zip_buffer = io.BytesIO()
with zipfile.ZipFile(zip_buffer, 'w', zipfile.ZIP_DEFLATED) as zip_file:
for table_name in table_names:
env = TemplateInitializer.init_jinja2()
render_info = await cls.__get_gen_render_info(auth, table_name)
for template_file, output_file in zip(render_info[0], render_info[1]):
render_content = env.get_template(template_file).render(**render_info[2])
zip_file.writestr(output_file, render_content)
zip_data = zip_buffer.getvalue()
zip_buffer.close()
return zip_data
@classmethod
async def sync_db_service(cls, auth: AuthSchema, table_name: str) -> SuccessResponse:
"""同步数据库"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_table_column_dao = GenTableColumnCRUD(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_name(db, table_name)
if gen_table:
table = GenTableOutSchema(**CamelCaseUtil.transform_result(gen_table))
table_columns = table.columns or [] # 确保不为None
table_column_map = {column.column_name: column for column in table_columns}
query_db_table_columns = await gen_table_column_dao.get_gen_db_table_columns_by_name(db, table_name)
db_table_columns = [
GenTableColumnOutSchema(**column) for column in CamelCaseUtil.transform_result(query_db_table_columns)
]
if not db_table_columns:
raise CustomException('同步数据失败,原表结构不存在')
db_table_column_names = [column.column_name for column in db_table_columns]
try:
for column in db_table_columns:
GenUtils.init_column_field(column, table)
if column.column_name in table_column_map:
prev_column = table_column_map[column.column_name]
# 使用getattr安全访问id属性
column_id = getattr(prev_column, 'id', None)
if column_id is not None:
# 为column设置id属性
column.id = column_id
if getattr(column, 'list', False): # 使用getattr安全访问属性
column.dict_type = prev_column.dict_type
column.query_type = prev_column.query_type
if (
prev_column.is_required != ''
and not column.pk
and (column.insert or column.edit)
and (column.usable_column or column.super_column)
):
column.is_required = prev_column.is_required
column.html_type = prev_column.html_type
# 使用getattr安全访问id属性
column_id = getattr(column, 'id', None)
if column_id is not None:
await gen_table_column_dao.update(id=column_id, data=column.model_dump(by_alias=True))
else:
await gen_table_column_dao.create(data=column.model_dump(by_alias=True))
del_columns = [column for column in table_columns if column.column_name not in db_table_column_names]
if del_columns:
for column in del_columns:
# 使用getattr安全访问id属性
column_id = getattr(column, 'id', None)
if column_id is not None:
await gen_table_column_dao.delete(ids=[column_id])
if isinstance(db, AsyncSession):
await db.commit()
return SuccessResponse(msg='同步成功')
except Exception as e:
if isinstance(db, AsyncSession):
try:
await db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'同步失败: {str(e)}')
else:
raise CustomException('业务表不存在')
@classmethod
async def set_sub_table(cls, auth: AuthSchema, gen_table: GenTableOutSchema) -> None:
"""设置主子表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
if gen_table.sub_table_name:
gen_table_dao = GenTableCRUD(auth=auth)
sub_table = await gen_table_dao.get_gen_table_by_name(db, gen_table.sub_table_name)
if sub_table:
gen_table.sub_table = GenTableOutSchema(**CamelCaseUtil.transform_result(sub_table))
@classmethod
async def _set_pk_column(cls, gen_table: GenTableOutSchema) -> None:
"""设置主键列信息"""
if gen_table.columns:
for column in gen_table.columns:
if column.pk:
gen_table.pk_column = column
break
if gen_table.pk_column is None and gen_table.columns:
gen_table.pk_column = gen_table.columns[0]
if gen_table.tpl_category == GenConstant.TPL_SUB and gen_table.sub_table:
if gen_table.sub_table.columns:
for column in gen_table.sub_table.columns:
if column.pk:
gen_table.sub_table.pk_column = column
break
if gen_table.sub_table.pk_column is None and gen_table.sub_table.columns:
gen_table.sub_table.pk_column = gen_table.sub_table.columns[0]
@classmethod
async def set_table_from_options(cls, gen_table: GenTableOutSchema) -> GenTableOutSchema:
"""设置代码生成其他选项值"""
params_obj = json.loads(gen_table.options) if gen_table.options else None
if params_obj:
gen_table.tree_code = params_obj.get(GenConstant.TREE_CODE)
gen_table.tree_parent_code = params_obj.get(GenConstant.TREE_PARENT_CODE)
gen_table.tree_name = params_obj.get(GenConstant.TREE_NAME)
gen_table.parent_menu_id = params_obj.get(GenConstant.PARENT_MENU_ID)
gen_table.parent_menu_name = params_obj.get(GenConstant.PARENT_MENU_NAME)
return gen_table
@classmethod
async def validate_edit(cls, edit_gen_table: GenTableUpdateSchema) -> None:
"""编辑保存参数校验"""
if edit_gen_table.tpl_category == GenConstant.TPL_TREE:
# 从options字段获取参数,而不是params
if not edit_gen_table.options:
raise CustomException(msg='树表参数不能为空')
params_obj = json.loads(edit_gen_table.options)
if GenConstant.TREE_CODE not in params_obj:
raise CustomException(msg='树编码字段不能为空')
elif GenConstant.TREE_PARENT_CODE not in params_obj:
raise CustomException(msg='树父编码字段不能为空')
elif GenConstant.TREE_NAME not in params_obj:
raise CustomException(msg='树名称字段不能为空')
elif edit_gen_table.tpl_category == GenConstant.TPL_SUB:
if not edit_gen_table.sub_table_name:
raise CustomException(msg='关联子表的表名不能为空')
elif not edit_gen_table.sub_table_fk_name:
raise CustomException(msg='子表关联的外键名不能为空')
@classmethod
async def __get_gen_render_info(cls, auth: AuthSchema, table_name: str) -> List[Any]:
"""获取生成代码渲染模板相关信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_dao = GenTableCRUD(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_name(db, table_name)
if gen_table:
gen_table_schema = GenTableOutSchema(**CamelCaseUtil.transform_result(gen_table))
await cls.set_sub_table(auth, gen_table_schema)
await cls._set_pk_column(gen_table_schema)
context = TemplateUtils.prepare_context(gen_table_schema)
template_list = TemplateUtils.get_template_list(
gen_table_schema.tpl_category or "",
gen_table_schema.tpl_web_type or ""
)
output_files = [TemplateUtils.get_file_name([template], gen_table_schema)[0] for template in template_list]
return [template_list, output_files, context, gen_table_schema]
else:
raise CustomException(msg=f'业务表 {table_name} 不存在')
@classmethod
def __get_gen_path(cls, gen_table: GenTableOutSchema, template: str) -> Optional[str]:
"""根据GenTableModel对象和模板名称生成路径"""
try:
gen_path = gen_table.gen_path or ""
if gen_path == '/':
file_name = TemplateUtils.get_file_name([template], gen_table)[0]
return os.path.join(os.getcwd(), GEN_PATH, file_name)
else:
file_name = TemplateUtils.get_file_name([template], gen_table)[0]
return os.path.join(gen_path, file_name)
except Exception:
return None
class GenTableColumnService:
"""代码生成业务表字段服务层"""
@classmethod
async def get_gen_table_column_list_by_table_id_service(cls, auth: AuthSchema, table_id: int) -> List[GenTableColumnOutSchema]:
"""获取业务表字段列表信息"""
if not auth.db:
raise CustomException(msg='数据库连接不存在')
# 确保db是AsyncSession类型
db = auth.db
if not isinstance(db, AsyncSession):
raise CustomException(msg='数据库连接类型不正确')
gen_table_column_dao = GenTableColumnCRUD(auth=auth)
gen_table_column_list_result = await gen_table_column_dao.get_gen_table_column_list_by_table_id(db, table_id)
return [
GenTableColumnOutSchema(**gen_table_column)
for gen_table_column in CamelCaseUtil.transform_result(gen_table_column_list_result)
]