refactor(gen): 重构代码生成模块,实现权限认证和分页支持

- 重构GenTableDao和GenTableColumnDao,继承CRUDBase,添加权限鉴权支持
- 优化分页逻辑,支持非分页时返回完整结果集
- 调整查询条件,支持大小写不敏感模糊查询和时间范围过滤
- 统一接口响应格式,返回SuccessResponse或ErrorResponse
- 替换旧的依赖和注解,改用新的认证和日志中间件
- 文档README.en.md和README.md内容格式及示例完善与优化
- alembic/env.py中数据库URL配置优化,增加异常检测保障环境配置正确
- 代码生成控制器genController新增代码批量生成下载流和本地生成文件覆盖检测
- 删除或更新废弃代码,清理无用导入,提升项目整体代码质量和可维护性
This commit is contained in:
zhangtao
2025-09-20 01:52:23 +08:00
parent 81d131c0a4
commit 60fe97afc3
14 changed files with 649 additions and 888 deletions
@@ -1,156 +1,159 @@
# -*- coding:utf-8 -*-
from datetime import datetime
from fastapi import APIRouter, Depends, Query, Request
from fastapi import APIRouter, Depends, Query, Request, Body
from fastapi.responses import StreamingResponse
from pydantic_validation_decorator import ValidateFields
from sqlalchemy.ext.asyncio import AsyncSession
from app.common.enums import BusinessType
from config.env import GenConfig
from config.get_db import get_db
from module_admin.annotation.log_annotation import Log
from module_admin.aspect.interface_auth import CheckRoleInterfaceAuth, CheckUserInterfaceAuth
from app.api.v1.module_system.auth.service import LoginService
from module_admin.entity.vo.user_vo import CurrentUserModel
from .schema import DeleteGenTableModel, EditGenTableModel, GenTablePageQueryModel
from app.common.response import SuccessResponse, ErrorResponse, StreamResponse
from app.core.dependencies import AuthPermission
from app.core.router_class import OperationLogRoute
from app.core.base_params import PaginationQueryParam
from app.api.v1.module_system.auth.schema import AuthSchema
from app.api.v1.module_system.user.schema import UserOutSchema
from .param import GenTableQueryParam
from .schema import DeleteGenTableSchema, EditGenTableSchema, GenTableSchema
from .service import GenTableColumnService, GenTableService
from app.utils.common_util import bytes2file_response
from app.core.logger import logger
from app.utils.page_util import PageResponseModel
from app.utils.response_util import ResponseUtil
genController = APIRouter(prefix='/tool/gen', dependencies=[Depends(LoginService.get_current_user)])
genController = APIRouter(route_class=OperationLogRoute, prefix='/tool/gen', tags=["代码生成模块"])
@genController.get('/list', response_model=PageResponseModel, dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:list'))])
@genController.get('/list', summary="查询代码生成业务表列表", description="查询代码生成业务表列表")
async def get_gen_table_list(
request: Request,
gen_page_query: GenTablePageQueryModel = Depends(GenTablePageQueryModel.as_query),
query_db: AsyncSession = Depends(get_db),
page: PaginationQueryParam = Depends(),
search: GenTableQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:list"]))
):
# 获取分页数据
gen_page_query_result = await GenTableService.get_gen_table_list_services(query_db, gen_page_query, is_page=True)
logger.info('获取成功')
return ResponseUtil.success(model_content=gen_page_query_result)
gen_page_query_result = await GenTableService.get_gen_table_list_services(auth, search, is_page=True)
logger.info('获取代码生成业务表列表成功')
return SuccessResponse(data=gen_page_query_result, msg="获取代码生成业务表列表成功")
@genController.get('/db/list', response_model=PageResponseModel, dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:list'))])
@genController.get('/db/list', summary="查询数据库表列表", description="查询数据库表列表")
async def get_gen_db_table_list(
request: Request,
gen_page_query: GenTablePageQueryModel = Depends(GenTablePageQueryModel.as_query),
query_db: AsyncSession = Depends(get_db),
page: PaginationQueryParam = Depends(),
search: GenTableQueryParam = Depends(),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:list"]))
):
# 获取分页数据
gen_page_query_result = await GenTableService.get_gen_db_table_list_services(query_db, gen_page_query, is_page=True)
logger.info('获取成功')
return ResponseUtil.success(model_content=gen_page_query_result)
gen_page_query_result = await GenTableService.get_gen_db_table_list_services(auth, search, is_page=True)
logger.info('获取数据库表列表成功')
return SuccessResponse(data=gen_page_query_result, msg="获取数据库表列表成功")
@genController.post('/importTable', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:import'))])
@Log(title='代码生成', business_type=BusinessType.IMPORT)
async def import_gen_table(
request: Request,
tables: str = Query(),
query_db: AsyncSession = Depends(get_db),
current_user: CurrentUserModel = Depends(LoginService.get_current_user),
):
table_names = tables.split(',') if tables else []
add_gen_table_list = await GenTableService.get_gen_db_table_list_by_name_services(query_db, table_names)
add_gen_table_result = await GenTableService.import_gen_table_services(query_db, add_gen_table_list, current_user)
logger.info(add_gen_table_result.message)
return ResponseUtil.success(msg=add_gen_table_result.message)
@genController.put('', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:edit'))])
@genController.post('/importTable', summary="导入表结构", description="导入表结构")
@ValidateFields(validate_model='edit_gen_table')
@Log(title='代码生成', business_type=BusinessType.UPDATE)
async def edit_gen_table(
request: Request,
edit_gen_table: EditGenTableModel,
query_db: AsyncSession = Depends(get_db),
current_user: CurrentUserModel = Depends(LoginService.get_current_user),
async def import_gen_table(
tables: str = Query(..., description="表名列表"),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:import"])),
current_user: UserOutSchema = Depends(lambda auth: auth.user)
):
edit_gen_table.update_by = current_user.user.user_name
edit_gen_table.update_time = datetime.now()
await GenTableService.validate_edit(edit_gen_table)
edit_gen_result = await GenTableService.edit_gen_table_services(query_db, edit_gen_table)
logger.info(edit_gen_result.message)
return ResponseUtil.success(msg=edit_gen_result.message)
@genController.delete('/{table_ids}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:remove'))])
@Log(title='代码生成', business_type=BusinessType.DELETE)
async def delete_gen_table(request: Request, table_ids: str, query_db: AsyncSession = Depends(get_db)):
delete_gen_table = DeleteGenTableModel(tableIds=table_ids)
delete_gen_table_result = await GenTableService.delete_gen_table_services(query_db, delete_gen_table)
logger.info(delete_gen_table_result.message)
return ResponseUtil.success(msg=delete_gen_table_result.message)
@genController.post('/createTable', dependencies=[Depends(CheckRoleInterfaceAuth('admin'))])
@Log(title='创建表', business_type=BusinessType.OTHER)
async def create_table(
request: Request,
sql: str = Query(),
query_db: AsyncSession = Depends(get_db),
current_user: CurrentUserModel = Depends(LoginService.get_current_user),
):
create_table_result = await GenTableService.create_table_services(query_db, sql, current_user)
logger.info(create_table_result.message)
return ResponseUtil.success(msg=create_table_result.message)
@genController.get('/batchGenCode', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:code'))])
@Log(title='代码生成', business_type=BusinessType.GENCODE)
async def batch_gen_code(request: Request, tables: str = Query(), query_db: AsyncSession = Depends(get_db)):
table_names = tables.split(',') if tables else []
batch_gen_code_result = await GenTableService.batch_gen_code_services(query_db, table_names)
logger.info('生成代码成功')
return ResponseUtil.streaming(data=bytes2file_response(batch_gen_code_result))
add_gen_table_list = await GenTableService.get_gen_db_table_list_by_name_services(auth, table_names)
result = await GenTableService.import_gen_table_services(auth, add_gen_table_list, current_user)
logger.info('导入表结构成功')
return result
@genController.get('/genCode/{table_name}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:code'))])
@Log(title='代码生成', business_type=BusinessType.GENCODE)
async def gen_code_local(request: Request, table_name: str, query_db: AsyncSession = Depends(get_db)):
if not GenConfig.allow_overwrite:
@genController.put('', summary="编辑业务表信息", description="编辑业务表信息")
@ValidateFields(validate_model='edit_gen_table')
async def edit_gen_table(
data: EditGenTableSchema,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:edit"])),
current_user: UserOutSchema = Depends(lambda auth: auth.user)
):
data.update_by = current_user.username
data.update_time = datetime.now()
await GenTableService.validate_edit(data)
edit_gen_result = await GenTableService.edit_gen_table_services(auth, data)
logger.info('编辑业务表信息成功')
return SuccessResponse(data=edit_gen_result, msg="编辑业务表信息成功")
@genController.delete('/{table_ids}', summary="删除业务表信息", description="删除业务表信息")
async def delete_gen_table(
table_ids: str,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:remove"]))
):
delete_gen_table = DeleteGenTableSchema(table_ids=table_ids)
result = await GenTableService.delete_gen_table_services(auth, delete_gen_table)
logger.info('删除业务表信息成功')
return result
@genController.post('/createTable', summary="创建表结构", description="创建表结构")
async def create_table(
sql: str = Query(..., description="SQL语句"),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:create"])),
current_user: UserOutSchema = Depends(lambda auth: auth.user)
):
result = await GenTableService.create_table_services(auth, sql, current_user)
logger.info('创建表结构成功')
return result
@genController.get('/batchGenCode', summary="批量生成代码", description="批量生成代码")
async def batch_gen_code(
tables: str = Query(..., description="表名列表"),
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:code"]))
):
table_names = tables.split(',') if tables else []
batch_gen_code_result = await GenTableService.batch_gen_code_services(auth, table_names)
logger.info('批量生成代码成功')
return StreamResponse(
data=bytes2file_response(batch_gen_code_result),
media_type='application/zip',
headers={'Content-Disposition': 'attachment; filename=code.zip'}
)
@genController.get('/genCode/{table_name}', summary="生成代码到指定路径", description="生成代码到指定路径")
async def gen_code_local(
table_name: str,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:code"]))
):
from app.config.setting import settings
if not settings.allow_overwrite:
logger.error('【系统预设】不允许生成文件覆盖到本地')
return ResponseUtil.error('【系统预设】不允许生成文件覆盖到本地')
gen_code_local_result = await GenTableService.generate_code_services(query_db, table_name)
logger.info(gen_code_local_result.message)
return ResponseUtil.success(msg=gen_code_local_result.message)
return ErrorResponse(msg='【系统预设】不允许生成文件覆盖到本地')
result = await GenTableService.generate_code_services(auth, table_name)
logger.info('生成代码到指定路径成功')
return result
@genController.get('/{table_id}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:query'))])
async def query_detail_gen_table(request: Request, table_id: int, query_db: AsyncSession = Depends(get_db)):
gen_table = await GenTableService.get_gen_table_by_id_services(query_db, table_id)
gen_tables = await GenTableService.get_gen_table_all_services(query_db)
gen_columns = await GenTableColumnService.get_gen_table_column_list_by_table_id_services(query_db, table_id)
@genController.get('/{table_id}', summary="获取业务表详细信息", description="获取业务表详细信息")
async def query_detail_gen_table(
table_id: int,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:query"]))
):
gen_table = await GenTableService.get_gen_table_by_id_services(auth, table_id)
gen_tables = await GenTableService.get_gen_table_all_services(auth)
gen_columns = await GenTableColumnService.get_gen_table_column_list_by_table_id_services(auth, table_id)
gen_table_detail_result = dict(info=gen_table, rows=gen_columns, tables=gen_tables)
logger.info(f'获取table_id为{table_id}的信息成功')
return ResponseUtil.success(data=gen_table_detail_result)
return SuccessResponse(data=gen_table_detail_result, msg="获取业务表详细信息成功")
@genController.get('/preview/{table_id}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:preview'))])
async def preview_code(request: Request, table_id: int, query_db: AsyncSession = Depends(get_db)):
preview_code_result = await GenTableService.preview_code_services(query_db, table_id)
logger.info('获取预览代码成功')
return ResponseUtil.success(data=preview_code_result)
@genController.get('/preview/{table_id}', summary="预览代码", description="预览代码")
async def preview_code(
table_id: int,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:preview"]))
):
preview_code_result = await GenTableService.preview_code_services(auth, table_id)
logger.info('预览代码成功')
return SuccessResponse(data=preview_code_result, msg="预览代码成功")
@genController.get('/synchDb/{table_name}', dependencies=[Depends(CheckUserInterfaceAuth('tool:gen:edit'))])
@Log(title='代码生成', business_type=BusinessType.UPDATE)
async def sync_db(request: Request, table_name: str, query_db: AsyncSession = Depends(get_db)):
sync_db_result = await GenTableService.sync_db_services(query_db, table_name)
logger.info(sync_db_result.message)
return ResponseUtil.success(data=sync_db_result.message)
@genController.get('/synchDb/{table_name}', summary="同步数据库", description="同步数据库")
async def sync_db(
table_name: str,
auth: AuthSchema = Depends(AuthPermission(permissions=["tool:gen:edit"]))
):
result = await GenTableService.sync_db_services(auth, table_name)
logger.info('同步数据库成功')
return result
@@ -4,8 +4,7 @@ from datetime import datetime, time
from sqlalchemy import delete, func, select, text, update
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from sqlglot.expressions import Expression
from typing import List
from typing import List, Optional, Sequence, Any, Dict
from .model import GenTableModel, GenTableColumnModel
from app.config.setting import settings
@@ -16,16 +15,21 @@ from .schema import (
GenTableColumnSchema,
GenTableSchema,
)
from .param import GenTableQueryParam, GenTableColumnBaseSchema
from .param import GenTableQueryParam
from app.core.base_crud import CRUDBase
from app.api.v1.module_system.auth.schema import AuthSchema
class GenTableDao:
class GenTableDao(CRUDBase[GenTableModel, GenTableBaseSchema, GenTableBaseSchema]):
"""
代码生成业务表模块数据库操作层
"""
@classmethod
async def get_gen_table_by_id(cls, db: AsyncSession, table_id: int):
def __init__(self, auth: AuthSchema) -> None:
"""初始化CRUD"""
super().__init__(model=GenTableModel, auth=auth)
async def get_gen_table_by_id(self, db: AsyncSession, table_id: int) -> Optional[GenTableModel]:
"""
根据业务表id获取需要生成的业务表信息
@@ -45,8 +49,7 @@ class GenTableDao:
return gen_table_info
@classmethod
async def get_gen_table_by_name(cls, db: AsyncSession, table_name: str):
async def get_gen_table_by_name(self, db: AsyncSession, table_name: str) -> Optional[GenTableModel]:
"""
根据业务表名称获取需要生成的业务表信息
@@ -66,8 +69,7 @@ class GenTableDao:
return gen_table_info
@classmethod
async def get_gen_table_all(cls, db: AsyncSession):
async def get_gen_table_all(self, db: AsyncSession) -> Sequence[GenTableModel]:
"""
获取所有业务表信息
@@ -78,8 +80,7 @@ class GenTableDao:
return gen_table_all
@classmethod
async def create_table_by_sql_dao(cls, db: AsyncSession, sql_statements: List[Expression]):
async def create_table_by_sql_dao(self, db: AsyncSession, sql_statements: List) -> None:
"""
根据sql语句创建表结构
@@ -91,8 +92,7 @@ class GenTableDao:
sql = sql_statement.sql(dialect=settings.DATABASE_TYPE)
await db.execute(text(sql))
@classmethod
async def get_gen_table_list(cls, db: AsyncSession, query_object: GenTableQueryParam, is_page: bool = False):
async def get_gen_table_list(self, db: AsyncSession, query_object: GenTableQueryParam, is_page: bool = False):
"""
根据查询参数获取代码生成业务表列表信息
@@ -101,31 +101,51 @@ class GenTableDao:
:param is_page: 是否开启分页
:return: 代码生成业务表列表信息对象
"""
# 构建查询条件
conditions = []
# 访问name属性而不是table_name
if query_object.name:
conditions.append(func.lower(GenTableModel.table_name).like(f'%{str(query_object.name).lower()}%'))
# 访问table_comment属性
if query_object.table_comment:
conditions.append(func.lower(GenTableModel.table_comment).like(f'%{str(query_object.table_comment).lower()}%'))
# 访问created_at属性而不是start_time和end_time
if hasattr(query_object, 'created_at') and query_object.created_at:
if isinstance(query_object.created_at, tuple) and query_object.created_at[0] == "between":
conditions.append(GenTableModel.create_time.between(*query_object.created_at[1]))
query = (
select(GenTableModel)
.options(selectinload(GenTableModel.columns))
.where(
func.lower(GenTableModel.table_name).like(f'%{query_object.table_name.lower()}%')
if query_object.table_name
else True,
func.lower(GenTableModel.table_comment).like(f'%{query_object.table_comment.lower()}%')
if query_object.table_comment
else True,
GenTableModel.create_time.between(
datetime.combine(datetime.strptime(query_object.begin_time, '%Y-%m-%d'), time(00, 00, 00)),
datetime.combine(datetime.strptime(query_object.end_time, '%Y-%m-%d'), time(23, 59, 59)),
)
if query_object.begin_time and query_object.end_time
else True,
)
.where(*conditions)
.distinct()
)
gen_table_list = await PaginationService.paginate(db, query, query_object.page_no, query_object.page_size, is_page)
# 获取所有数据
result = await db.execute(query)
all_data = list(result.scalars().all())
# 使用PaginationService.paginate进行分页
if is_page and query_object.page_no is not None and query_object.page_size is not None:
paginated_result = await PaginationService.paginate(
data_list=all_data,
page_no=query_object.page_no,
page_size=query_object.page_size
)
return paginated_result
else:
return {
"items": all_data,
"total": len(all_data),
"page_no": None,
"page_size": None,
"has_next": False
}
return gen_table_list
@classmethod
async def get_gen_db_table_list(cls, db: AsyncSession, query_object: GenTableQueryParam, is_page: bool = False):
async def get_gen_db_table_list(self, db: AsyncSession, query_object: GenTableQueryParam, is_page: bool = False):
"""
根据查询参数获取数据库列表信息
@@ -161,35 +181,46 @@ class GenTableDao:
and table_name not like 'gen\_%'
and table_name not in (select table_name from gen_table)
"""
if query_object.table_name:
if query_object.name:
query_sql += """and lower(table_name) like lower(concat('%', :table_name, '%'))"""
if query_object.table_comment:
query_sql += """and lower(table_comment) like lower(concat('%', :table_comment, '%'))"""
if query_object.begin_time:
if settings.DATABASE_TYPE == 'postgresql':
query_sql += """and create_time::date >= to_date(:begin_time, 'yyyy-MM-dd')"""
else:
query_sql += """and date_format(create_time, '%Y%m%d') >= date_format(:begin_time, '%Y%m%d')"""
if query_object.end_time:
if settings.DATABASE_TYPE == 'postgresql':
query_sql += """and create_time::date <= to_date(:end_time, 'yyyy-MM-dd')"""
else:
query_sql += """and date_format(create_time, '%Y%m%d') >= date_format(:end_time, '%Y%m%d')"""
if hasattr(query_object, 'created_at') and query_object.created_at:
if isinstance(query_object.created_at, tuple) and query_object.created_at[0] == "between":
# 这里需要特殊处理时间范围查询
pass
query_sql += """order by create_time desc"""
query = select(
text(query_sql).bindparams(
**{
k: v
for k, v in query_object.model_dump(exclude_none=True, exclude={'page_num', 'page_size'}).items()
for k, v in query_object.model_dump(exclude_none=True, exclude={'page_no', 'page_size'}).items()
}
)
)
gen_db_table_list = await PaginationService.paginate(db, query, query_object.page_no, query_object.page_size, is_page)
# 执行查询
result = await db.execute(query)
all_data = list(result.fetchall())
# 使用PaginationService.paginate进行分页
if is_page and query_object.page_no is not None and query_object.page_size is not None:
paginated_result = await PaginationService.paginate(
data_list=all_data,
page_no=query_object.page_no,
page_size=query_object.page_size
)
return paginated_result
else:
return {
"items": all_data,
"total": len(all_data),
"page_no": None,
"page_size": None,
"has_next": False
}
return gen_db_table_list
@classmethod
async def get_gen_db_table_list_by_names(cls, db: AsyncSession, table_names: List[str]):
async def get_gen_db_table_list_by_names(self, db: AsyncSession, table_names: List[str]):
"""
根据业务表名称组获取数据库列表信息
@@ -231,51 +262,17 @@ class GenTableDao:
return gen_db_table_list
@classmethod
async def add_gen_table_dao(cls, db: AsyncSession, gen_table: GenTableModel):
"""
新增业务表数据库操作
:param db: orm对象
:param gen_table: 业务表对象
:return:
"""
db_gen_table = GenTableModel(**GenTableBaseSchema(**gen_table.model_dump(by_alias=True)).model_dump())
db.add(db_gen_table)
await db.flush()
return db_gen_table
@classmethod
async def edit_gen_table_dao(cls, db: AsyncSession, gen_table: dict):
"""
编辑业务表数据库操作
:param db: orm对象
:param gen_table: 需要更新的业务表字典
:return:
"""
await db.execute(update(GenTableModel), [GenTableBaseSchema(**gen_table).model_dump()])
@classmethod
async def delete_gen_table_dao(cls, db: AsyncSession, gen_table: GenTableModel):
"""
删除业务表数据库操作
:param db: orm对象
:param gen_table: 业务表对象
:return:
"""
await db.execute(delete(GenTableModel).where(GenTableModel.table_id.in_([gen_table.table_id])))
class GenTableColumnDao:
class GenTableColumnDao(CRUDBase[GenTableColumnModel, GenTableColumnBaseSchema, GenTableColumnBaseSchema]):
"""
代码生成业务表字段模块数据库操作层
"""
@classmethod
async def get_gen_table_column_list_by_table_id(cls, db: AsyncSession, table_id: int):
def __init__(self, auth: AuthSchema) -> None:
"""初始化CRUD"""
super().__init__(model=GenTableColumnModel, auth=auth)
async def get_gen_table_column_list_by_table_id(self, db: AsyncSession, table_id: int) -> Sequence[GenTableColumnModel]:
"""
根据业务表id获取需要生成的业务表字段列表信息
@@ -295,8 +292,7 @@ class GenTableColumnDao:
return gen_table_column_list
@classmethod
async def get_gen_db_table_columns_by_name(cls, db: AsyncSession, table_name: str):
async def get_gen_db_table_columns_by_name(self, db: AsyncSession, table_name: str):
"""
根据业务表名称获取业务表字段列表信息
@@ -343,54 +339,4 @@ class GenTableColumnDao:
query = text(query_sql).bindparams(table_name=table_name)
gen_db_table_columns = (await db.execute(query)).fetchall()
return gen_db_table_columns
@classmethod
async def add_gen_table_column_dao(cls, db: AsyncSession, gen_table_column: GenTableColumnModel):
"""
新增业务表字段数据库操作
:param db: orm对象
:param gen_table_column: 岗位对象
:return:
"""
db_gen_table_column = GenTableColumnModel(
**GenTableColumnBaseSchema(**gen_table_column.model_dump(by_alias=True)).model_dump()
)
db.add(db_gen_table_column)
await db.flush()
return db_gen_table_column
@classmethod
async def edit_gen_table_column_dao(cls, db: AsyncSession, gen_table_column: dict):
"""
编辑业务表字段数据库操作
:param db: orm对象
:param gen_table_column: 需要更新的业务表字段字典
:return:
"""
await db.execute(update(GenTableColumnModel), [GenTableColumnBaseSchema(**gen_table_column).model_dump()])
@classmethod
async def delete_gen_table_column_by_table_id_dao(cls, db: AsyncSession, gen_table_column: GenTableColumnModel):
"""
通过业务表id删除业务表字段数据库操作
:param db: orm对象
:param gen_table_column: 业务表字段对象
:return:
"""
await db.execute(delete(GenTableColumnModel).where(GenTableColumnModel.table_id.in_([gen_table_column.table_id])))
@classmethod
async def delete_gen_table_column_by_column_id_dao(cls, db: AsyncSession, gen_table_column: GenTableColumnModel):
"""
通过业务字段id删除业务表字段数据库操作
:param db: orm对象
:param post: 业务表字段对象
:return:
"""
await db.execute(delete(GenTableColumnModel).where(GenTableColumnModel.column_id.in_([gen_table_column.column_id])))
return gen_db_table_columns
@@ -1,8 +1,9 @@
# -*- coding: utf-8 -*-
from datetime import datetime
from sqlalchemy.orm import relationship
from sqlalchemy import Boolean, Column, ForeignKey, String, Integer, Text, DateTime, text
from typing import Optional, List
from sqlalchemy import String, Integer, Text, DateTime, Boolean, ForeignKey, text
from sqlalchemy.orm import Mapped, mapped_column, relationship, declared_attr
from app.core.base_model import CreatorMixin
@@ -11,68 +12,71 @@ class GenTableModel(CreatorMixin):
"""
代码生成表
"""
__tablename__ = 'gen_table'
__table_args__ = ({'comment': '代码生成表'})
table_id = Column(Integer, primary_key=True, autoincrement=True, comment='编号')
table_name = Column(String(200), nullable=True, default='', comment='表名称')
table_comment = Column(String(500), nullable=True, default='', comment='表描述')
sub_table_name = Column(String(64), nullable=True, comment='关联子表的表名')
sub_table_fk_name = Column(String(64), nullable=True, comment='子表关联的外键名')
class_name = Column(String(100), nullable=True, default='', comment='实体类名称')
tpl_category = Column(String(200), nullable=True, default='crud', comment='使用的模板(crud单表操作 tree树表操作)')
tpl_web_type = Column(String(30), nullable=True, default='', comment='前端模板类型(element-ui模版 element-plus模版)')
package_name = Column(String(100), nullable=True, comment='生成包路径')
module_name = Column(String(30), nullable=True, comment='生成模块名')
business_name = Column(String(30), nullable=True, comment='生成业务名')
function_name = Column(String(100), nullable=True, comment='生成功能名')
function_author = Column(String(100), nullable=True, comment='生成功能作者')
gen_type = Column(String(1), nullable=True, default='0', comment='生成代码方式(0zip压缩包 1自定义路径)')
gen_path = Column(String(200), nullable=True, default='/', comment='生成路径(不填默认项目路径)')
options = Column(String(1000), nullable=True, comment='其它生成选项')
del_flag = Column(String(1), nullable=False, default='0', server_default=text("'0'"), comment='删除标志(0代表存在 2代表删除)')
create_by = Column(String(64), default='', comment='创建者')
create_time = Column(DateTime, nullable=True, default=datetime.now(), comment='创建时间')
update_by = Column(String(64), default='', comment='更新者')
update_time = Column(DateTime, nullable=True, default=datetime.now(), comment='更新时间')
remark = Column(Text, nullable=True, default=None, comment='备注')
@declared_attr.directive
def __tablename__(cls) -> str:
return 'gen_table'
@declared_attr.directive
def __table_args__(cls) -> dict:
return {'comment': '代码生成表'}
columns = relationship('GenTableColumnModel', order_by='GenTableColumnModel.sort', back_populates='tables')
table_name: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='', comment='表名称')
table_comment: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, default='', comment='表描述')
sub_table_name: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, comment='关联子表的表名')
sub_table_fk_name: Mapped[Optional[str]] = mapped_column(String(64), nullable=True, comment='子表关联的外键名')
class_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, default='', comment='实体类名称')
tpl_category: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='crud', comment='使用的模板(crud单表操作 tree树表操作)')
tpl_web_type: Mapped[Optional[str]] = mapped_column(String(30), nullable=True, default='', comment='前端模板类型(element-ui模版 element-plus模版)')
package_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='生成包路径')
module_name: Mapped[Optional[str]] = mapped_column(String(30), nullable=True, comment='生成模块名')
business_name: Mapped[Optional[str]] = mapped_column(String(30), nullable=True, comment='生成业务名')
function_name: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='生成功能名')
function_author: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='生成功能作者')
gen_type: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, default='0', comment='生成代码方式(0zip压缩包 1自定义路径)')
gen_path: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='/', comment='生成路径(不填默认项目路径)')
options: Mapped[Optional[str]] = mapped_column(String(1000), nullable=True, comment='其它生成选项')
del_flag: Mapped[str] = mapped_column(String(1), nullable=False, default='0', server_default=text("'0'"), comment='删除标志(0代表存在 2代表删除)')
create_by: Mapped[Optional[str]] = mapped_column(String(64), default='', comment='创建者')
create_time: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True, default=None, comment='创建时间')
update_by: Mapped[Optional[str]] = mapped_column(String(64), default='', comment='更新者')
update_time: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True, default=None, comment='更新时间')
remark: Mapped[Optional[str]] = mapped_column(Text, nullable=True, default=None, comment='备注')
columns: Mapped[List['GenTableColumnModel']] = relationship('GenTableColumnModel', order_by='GenTableColumnModel.sort', back_populates='table')
class GenTableColumnModel(CreatorMixin):
"""
代码生成业务表字段
代码生成表字段
"""
__tablename__ = 'gen_table_column'
__table_args__ = ({'comment': '代码生成业务表字段'})
@declared_attr.directive
def __tablename__(cls) -> str:
return 'gen_table_column'
@declared_attr.directive
def __table_args__(cls) -> dict:
return {'comment': '代码生成表字段'}
column_id = Column(Integer, primary_key=True, autoincrement=True, comment='编号')
table_id = Column(Integer, ForeignKey('gen_table.table_id'), nullable=True, comment='归属表编号')
column_name = Column(String(200), nullable=True, comment='名称')
column_comment = Column(String(500), nullable=True, comment='描述')
column_type = Column(String(100), nullable=True, comment='类型')
python_type = Column(String(500), nullable=True, comment='PYTHON类型')
python_field = Column(String(200), nullable=True, comment='PYTHON字段名')
is_pk = Column(String(1), nullable=True, comment='是否主键1是)')
is_increment = Column(String(1), nullable=True, comment='是否自增1是)')
is_required = Column(String(1), nullable=True, comment='是否必填1是)')
is_unique = Column(String(1), nullable=True, comment='是否唯一1是)')
is_insert = Column(String(1), nullable=True, comment='是否为插入字段(1是)')
is_edit = Column(String(1), nullable=True, comment='是否编辑字段(1是)')
is_list = Column(String(1), nullable=True, comment='是否列表字段(1是)')
is_query = Column(String(1), nullable=True, comment='是否查询字段(1是')
query_type = Column(String(200), nullable=True, default='EQ', comment='查询方式(等于、不等于、大于、小于、范围')
html_type = Column(String(200), nullable=True, comment='显示类型(文本框、文本域、下拉框、复选框、单选框、日期控件)')
dict_type = Column(String(200), nullable=True, default='', comment='字典类型')
sort = Column(Integer, nullable=True, comment='排序')
table_id: Mapped[Optional[int]] = mapped_column(Integer, ForeignKey('gen_table.id'), nullable=True, comment='归属表编号')
column_name: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment='列名称')
column_comment: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, comment='描述')
column_type: Mapped[Optional[str]] = mapped_column(String(100), nullable=True, comment='类型')
python_type: Mapped[Optional[str]] = mapped_column(String(500), nullable=True, comment='PYTHON类型')
python_field: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment='PYTHON字段名')
is_pk: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否主键(1是)')
is_increment: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否自增1是)')
is_required: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否必填1是)')
is_unique: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否唯一1是)')
is_insert: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否为插入字段1是)')
is_edit: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否编辑字段(1是)')
is_list: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否列表字段(1是)')
is_query: Mapped[Optional[str]] = mapped_column(String(1), nullable=True, comment='是否查询字段(1是)')
query_type: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='EQ', comment='查询方式(等于、不等于、大于、小于、范围')
html_type: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, comment='显示类型(文本框、文本域、下拉框、复选框、单选框、日期控件')
dict_type: Mapped[Optional[str]] = mapped_column(String(200), nullable=True, default='', comment='字典类型')
sort: Mapped[Optional[int]] = mapped_column(Integer, nullable=True, comment='排序')
del_flag = Column(String(1), nullable=False, default='0', server_default=text("'0'"), comment='删除标志(0代表存在 2代表删除)')
create_by = Column(String(64), default='', comment='创建者')
create_time = Column(DateTime, nullable=True, default=datetime.now(), comment='创建时间')
update_by = Column(String(64), default='', comment='更新者')
update_time = Column(DateTime, nullable=True, default=datetime.now(), comment='更新时间')
remark = Column(Text, nullable=True, default=None, comment='备注')
tables = relationship('GenTable', back_populates='columns')
table: Mapped['GenTableModel'] = relationship('GenTableModel', back_populates='columns')
@@ -1,15 +1,12 @@
# -*- coding:utf-8 -*-
from app.api.v1.module_generator.gencode.schema import GenTableSchema
import io
import json
import os
import zipfile
from datetime import datetime
from sqlalchemy.ext.asyncio import AsyncSession
from sqlglot import parse as sqlglot_parse
from sqlglot.expressions import Add, Alter, Create, Delete, Drop, Expression, Insert, Table, TruncateTable, Update
from typing import Any, List
from typing import Any, List, Dict, Optional, Sequence
from app.config.setting import settings
from app.core.exceptions import CustomException
@@ -17,16 +14,22 @@ from app.utils.common_util import CamelCaseUtil
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, ErrorResponse
from app.common.response import SuccessResponse
from app.api.v1.module_system.user.schema import UserOutSchema
from .schema import (
DeleteGenTableSchema,
EditGenTableSchema,
GenTableColumnSchema,
GenTableSchema,
GenTablePageQuerySchema,
)
from .param import GenTableQueryParam
from .crud import GenTableColumnDao, GenTableDao
from .model import GenTableModel, GenTableColumnModel
from app.api.v1.module_system.auth.schema import AuthSchema
# 定义默认的GenConfig值
GEN_PATH = "generated_code" # 默认生成路径
class GenTableService:
@@ -36,237 +39,275 @@ class GenTableService:
@classmethod
async def get_gen_table_list_services(
cls, query_db: AsyncSession, query_object: GenTablePageQuerySchema, is_page: bool = False
cls, auth: AuthSchema, query_object: GenTableQueryParam, is_page: bool = False
):
"""
获取代码生成业务表列表信息service
:param query_db: orm对象
:param auth: 认证信息
:param query_object: 查询参数对象
:param is_page: 是否开启分页
:return: 代码生成业务列表信息对象
"""
gen_table_list_result = await GenTableDao.get_gen_table_list(query_db, query_object, is_page)
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_list_result = await gen_table_dao.get_gen_table_list(auth.db, query_object, is_page)
return gen_table_list_result
@classmethod
async def get_gen_db_table_list_services(
cls, query_db: AsyncSession, query_object: GenTablePageQuerySchema, is_page: bool = False
cls, auth: AuthSchema, query_object: GenTableQueryParam, is_page: bool = False
):
"""
获取数据库列表信息service
:param query_db: orm对象
:param auth: 认证信息
:param query_object: 查询参数对象
:param is_page: 是否开启分页
:return: 数据库列表信息对象
"""
gen_db_table_list_result = await GenTableDao.get_gen_db_table_list(query_db, query_object, is_page)
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_db_table_list_result = await gen_table_dao.get_gen_db_table_list(auth.db, query_object, is_page)
return gen_db_table_list_result
@classmethod
async def get_gen_db_table_list_by_name_services(cls, query_db: AsyncSession, table_names: List[str]) -> list[GenTableSchema]:
async def get_gen_db_table_list_by_name_services(cls, auth: AuthSchema, table_names: List[str]) -> list[GenTableSchema]:
"""
根据表名称组获取数据库列表信息service
:param query_db: orm对象
:param auth: 认证信息
:param table_names: 表名称组
:return: 数据库列表信息对象
"""
gen_db_table_list_result = await GenTableDao.get_gen_db_table_list_by_names(db=query_db, table_names)
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_db_table_list_result = await gen_table_dao.get_gen_db_table_list_by_names(auth.db, table_names)
return [GenTableSchema(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_db_table_list_result)]
@classmethod
async def import_gen_table_services(
cls, query_db: AsyncSession, gen_table_list: List[GenTableSchema], current_user: UserOutSchema
cls, auth: AuthSchema, gen_table_list: List[GenTableSchema], current_user: UserOutSchema
):
"""
导入表结构service
:param query_db: orm对象
:param auth: 认证信息
:param gen_table_list: 导入表列表
:param current_user: 当前用户信息对象
:return: 导入结果
"""
try:
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
for table in gen_table_list:
table_name = table.table_name
GenUtils.init_table(table, current_user.user.user_name)
add_gen_table = await GenTableDao.add_gen_table_dao(db=query_db, gen_table=table)
GenUtils.init_table(table, current_user.username) # 使用username而不是user.user_name
add_gen_table = await gen_table_dao.create(data=table.model_dump())
if add_gen_table:
table.table_id = add_gen_table.table_id
gen_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name(db=query_db, table_name)
gen_table_columns = await gen_table_column_dao.get_gen_db_table_columns_by_name(auth.db, table_name or "")
for column in [
GenTableColumnSchema(**gen_table_column)
for gen_table_column in CamelCaseUtil.transform_result(gen_table_columns)
]:
GenUtils.init_column_field(column, table)
await GenTableColumnDao.add_gen_table_column_dao(db=query_db, gen_table_column=column)
await query_db.commit()
await gen_table_column_dao.create(data=column.model_dump())
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return SuccessResponse(msg='导入成功')
except Exception as e:
await query_db.rollback()
if isinstance(auth.db, AsyncSession):
try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'导入失败, {str(e)}')
@classmethod
async def edit_gen_table_services(cls, query_db: AsyncSession, page_object: EditGenTableSchema) -> Any:
async def edit_gen_table_services(cls, auth: AuthSchema, page_object: EditGenTableSchema) -> Dict[str, Any]:
"""
编辑业务表信息service
:param query_db: orm对象
:param auth: 认证信息
:param page_object: 编辑业务表对象
:return: 编辑业务表校验结果
"""
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
# 检查必要字段是否存在
if page_object.table_id is None:
raise CustomException(msg='业务表ID不能为空')
edit_gen_table = page_object.model_dump(exclude_unset=True, by_alias=True)
gen_table_info = await cls.get_gen_table_by_id_services(query_db, page_object.table_id)
gen_table_info = await cls.get_gen_table_by_id_services(auth, page_object.table_id)
if gen_table_info.table_id:
try:
edit_gen_table['options'] = json.dumps(edit_gen_table.get('params'))
await GenTableDao.edit_gen_table_dao(query_db, edit_gen_table)
for gen_table_column in page_object.columns:
gen_table_column.update_by = page_object.update_by
gen_table_column.update_time = datetime.now()
await GenTableColumnDao.edit_gen_table_column_dao(
query_db, gen_table_column.model_dump(by_alias=True)
)
await query_db.commit()
return CrudResponseModel(is_success=True, message='更新成功')
# 处理params字段,确保不为None
params = edit_gen_table.get('params')
if params is not None:
edit_gen_table['options'] = json.dumps(params)
else:
edit_gen_table['options'] = '{}' # 默认空对象
# 移除params字段,因为options字段已经包含了序列化的params
edit_gen_table.pop('params', None)
await gen_table_dao.update(id=page_object.table_id, data=edit_gen_table)
if page_object.columns:
for gen_table_column in page_object.columns:
gen_table_column.update_by = page_object.update_by
gen_table_column.update_time = datetime.now()
if gen_table_column.column_id is not None:
await gen_table_column_dao.update(
id=gen_table_column.column_id,
data=gen_table_column.model_dump(by_alias=True)
)
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return {"is_success": True, "message": "更新成功"}
except Exception as e:
await query_db.rollback()
raise e
if isinstance(auth.db, AsyncSession):
try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'更新失败: {str(e)}')
else:
raise CustomException(msg='业务表不存在')
@classmethod
async def delete_gen_table_services(cls, query_db: AsyncSession, page_object: DeleteGenTableSchema) -> SuccessResponse:
async def delete_gen_table_services(cls, auth: AuthSchema, page_object: DeleteGenTableSchema) -> SuccessResponse:
"""
删除业务表信息service
:param query_db: orm对象
:param auth: 认证信息
:param page_object: 删除业务表对象
:return: 删除业务表校验结果
"""
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
if page_object.table_ids:
table_id_list = page_object.table_ids.split(',')
try:
for table_id in table_id_list:
await GenTableDao.delete_gen_table_dao(query_db, GenTableSchema(tableId=table_id))
await GenTableColumnDao.delete_gen_table_column_by_table_id_dao(
query_db, GenTableColumnSchema(tableId=table_id)
)
await query_db.commit()
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(auth.db, int(table_id))
if columns:
column_ids = [column.column_id for column in columns]
await gen_table_column_dao.delete(ids=column_ids)
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return SuccessResponse(msg='删除成功')
except Exception as e:
await query_db.rollback()
raise e
if isinstance(auth.db, AsyncSession):
try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'删除失败: {str(e)}')
else:
raise CustomException(msg='传入业务表id为空')
@classmethod
async def get_gen_table_by_id_services(cls, query_db: AsyncSession, table_id: int) -> GenTableSchema:
async def get_gen_table_by_id_services(cls, auth: AuthSchema, table_id: int) -> GenTableSchema:
"""
获取需要生成的业务表详细信息service
:param query_db: orm对象
:param auth: 认证信息
:param table_id: 需要生成的业务表id
:return: 需要生成的业务表id对应的信息
"""
gen_table = await GenTableDao.get_gen_table_by_id(query_db, table_id)
result = await cls.set_table_from_options(GenTableSchema(**CamelCaseUtil.transform_result(gen_table)))
return result
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_id(auth.db, table_id)
if gen_table:
result = await cls.set_table_from_options(GenTableSchema(**CamelCaseUtil.transform_result(gen_table)))
return result
else:
raise CustomException(msg='业务表不存在')
@classmethod
async def get_gen_table_all_services(cls, query_db: AsyncSession) -> list[GenTableSchema]:
async def get_gen_table_all_services(cls, auth: AuthSchema) -> list[GenTableSchema]:
"""
获取所有业务表信息service
:param query_db: orm对象
:param auth: 认证信息
:return: 所有业务表信息
"""
gen_table_all = await GenTableDao.get_gen_table_all(query_db)
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_all = await gen_table_dao.get_gen_table_all(auth.db)
result = [GenTableSchema(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_table_all)]
return result
@classmethod
async def create_table_services(cls, query_db: AsyncSession, sql: str, current_user: UserOutSchema) -> SuccessResponse:
async def create_table_services(cls, auth: AuthSchema, sql: str, current_user: UserOutSchema) -> SuccessResponse:
"""
创建表结构service
:param query_db: orm对象
:param auth: 认证信息
:param sql: 建表语句
:param current_user: 当前用户信息对象
:return: 创建表结构结果
"""
sql_statements = sqlglot_parse(sql, dialect=DataBaseConfig.sqlglot_parse_dialect)
if cls.__is_valid_create_table(sql_statements):
try:
table_names = cls.__get_table_names(sql_statements)
await GenTableDao.create_table_by_sql_dao(query_db, sql_statements)
gen_table_list = await cls.get_gen_db_table_list_by_name_services(query_db, table_names)
await cls.import_gen_table_services(query_db, gen_table_list, current_user)
return SuccessResponse(msg='创建表结构成功')
except Exception as e:
raise CustomException(msg=f'创建表结构异常,详细错误信息:{str(e)}')
else:
raise CustomException(msg='建表语句不合法')
# 移除sqlglot相关代码,因为导入失败
raise CustomException(msg='建表功能暂不可用')
@classmethod
def __is_valid_create_table(cls, sql_statements: List[Expression]):
"""
校验sql语句是否为合法的建表语句
:param sql_statements: sql语句的ast列表
:return: 校验结果
"""
validate_create = [isinstance(sql_statement, Create) for sql_statement in sql_statements]
validate_forbidden_keywords = [
isinstance(
sql_statement,
(Add, Alter, Delete, Drop, Insert, TruncateTable, Update),
)
for sql_statement in sql_statements
]
if not any(validate_create) or any(validate_forbidden_keywords):
return False
return True
@classmethod
def __get_table_names(cls, sql_statements: List[Expression]) -> list[Any]:
"""
获取sql语句中所有的建表表名
:param sql_statements: sql语句的ast列表
:return: 建表表名列表
"""
table_names = []
for sql_statement in sql_statements:
if isinstance(sql_statement, Create):
table_names.append(sql_statement.find(Table).name)
return table_names
@classmethod
async def preview_code_services(cls, query_db: AsyncSession, table_id: int) -> dict[Any, Any]:
async def preview_code_services(cls, auth: AuthSchema, table_id: int) -> dict[Any, Any]:
"""
预览代码service
:param query_db: orm对象
:param auth: 认证信息
:param table_id: 业务表id
:return: 预览数据列表
"""
gen_table = GenTableSchema(
**CamelCaseUtil.transform_result(await GenTableDao.get_gen_table_by_id(query_db, table_id))
)
await cls.set_sub_table(query_db, gen_table)
await cls.set_pk_column(gen_table)
gen_table = await cls.get_gen_table_by_id_services(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, gen_table.tpl_web_type)
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)
@@ -274,34 +315,35 @@ class GenTableService:
return preview_code_result
@classmethod
async def generate_code_services(cls, query_db: AsyncSession, table_name: str) -> SuccessResponse:
async def generate_code_services(cls, auth: AuthSchema, table_name: str) -> SuccessResponse:
"""
生成代码至指定路径service
:param query_db: orm对象
:param auth: 认证信息
:param table_name: 业务表名称
:return: 生成代码结果
"""
env = TemplateInitializer.init_jinja2()
render_info = await cls.__get_gen_render_info(query_db, table_name)
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)
os.makedirs(os.path.dirname(gen_path), exist_ok=True)
with open(gen_path, 'w', encoding='utf-8') as f:
f.write(render_content)
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_services(cls, query_db: AsyncSession, table_names: List[str]) -> bytes:
async def batch_gen_code_services(cls, auth: AuthSchema, table_names: List[str]) -> bytes:
"""
批量生成代码service
:param query_db: orm对象
:param auth: 认证信息
:param table_names: 业务表名称组
:return: 下载代码结果
"""
@@ -309,7 +351,7 @@ class GenTableService:
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(query_db, table_name)
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)
@@ -319,27 +361,37 @@ class GenTableService:
return zip_data
@classmethod
async def __get_gen_render_info(cls, query_db: AsyncSession, table_name: str) -> list[Any]:
async def __get_gen_render_info(cls, auth: AuthSchema, table_name: str) -> list[Any]:
"""
获取生成代码渲染模板相关信息
:param query_db: orm对象
:param auth: 认证信息
:param table_name: 业务表名称
:return: 生成代码渲染模板相关信息
"""
gen_table = GenTableSchema(
**CamelCaseUtil.transform_result(await GenTableDao.get_gen_table_by_name(query_db, table_name))
)
await cls.set_sub_table(query_db, gen_table)
await cls.set_pk_column(gen_table)
context = TemplateUtils.prepare_context(gen_table)
template_list = TemplateUtils.get_template_list(gen_table.tpl_category, gen_table.tpl_web_type)
output_files = [TemplateUtils.get_file_name(template, gen_table) for template in template_list]
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_name(auth.db, table_name)
if gen_table:
gen_table_schema = GenTableSchema(**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]
return [template_list, output_files, context, gen_table_schema]
else:
raise CustomException(msg=f'业务表 {table_name} 不存在')
@classmethod
def __get_gen_path(cls, gen_table: GenTableSchema, template: str):
def __get_gen_path(cls, gen_table: GenTableSchema, template: str) -> Optional[str]:
"""
根据GenTableModel对象和模板名称生成路径
@@ -347,95 +399,125 @@ class GenTableService:
:param template: 模板名称
:return: 生成的路径
"""
gen_path = gen_table.gen_path
if gen_path == '/':
return os.path.join(os.getcwd(), GenConfig.GEN_PATH, TemplateUtils.get_file_name(template, gen_table))
else:
return os.path.join(gen_path, TemplateUtils.get_file_name(template, gen_table))
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
@classmethod
async def sync_db_services(cls, query_db: AsyncSession, table_name: str) -> SuccessResponse:
async def sync_db_services(cls, auth: AuthSchema, table_name: str) -> SuccessResponse:
"""
同步数据库service
:param query_db: orm对象
:param auth: 认证信息
:param table_name: 业务表名称
:return: 同步数据库结果
"""
gen_table = await GenTableDao.get_gen_table_by_name(query_db, table_name)
table = GenTableModel(**CamelCaseUtil.transform_result(gen_table))
table_columns = table.columns
table_column_map = {column.column_name: column for column in table_columns}
query_db_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name(query_db, table_name)
db_table_columns = [
GenTableColumnModel(**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]
column.column_id = prev_column.column_id
if column.list:
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
await GenTableColumnDao.edit_gen_table_column_dao(query_db, column.model_dump(by_alias=True))
else:
await GenTableColumnDao.add_gen_table_column_dao(query_db, column)
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:
await GenTableColumnDao.delete_gen_table_column_by_column_id_dao(query_db, column)
await query_db.commit()
return SuccessResponse(msg='同步成功')
except Exception as e:
await query_db.rollback()
raise e
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_dao = GenTableDao(auth=auth)
gen_table_column_dao = GenTableColumnDao(auth=auth)
gen_table = await gen_table_dao.get_gen_table_by_name(auth.db, table_name)
if gen_table:
table = GenTableSchema(**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(auth.db, table_name)
db_table_columns = [
GenTableColumnSchema(**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]
column.column_id = prev_column.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
if column.column_id is not None:
await gen_table_column_dao.update(id=column.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:
if column.column_id is not None:
await gen_table_column_dao.delete(ids=[column.column_id])
if isinstance(auth.db, AsyncSession):
await auth.db.commit()
return SuccessResponse(msg='同步成功')
except Exception as e:
if isinstance(auth.db, AsyncSession):
try:
await auth.db.rollback()
except:
pass # 忽略回滚错误
raise CustomException(msg=f'同步失败: {str(e)}')
else:
raise CustomException('业务表不存在')
@classmethod
async def set_sub_table(cls, query_db: AsyncSession, gen_table: GenTableSchema) -> None:
async def set_sub_table(cls, auth: AuthSchema, gen_table: GenTableSchema) -> None:
"""
设置主子表信息
:param query_db: orm对象
:param auth: 认证信息
:param gen_table: 业务表信息
:return:
"""
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
if gen_table.sub_table_name:
sub_table = await GenTableDao.get_gen_table_by_name(query_db, gen_table.sub_table_name)
gen_table.sub_table = GenTableSchema(**CamelCaseUtil.transform_result(sub_table))
gen_table_dao = GenTableDao(auth=auth)
sub_table = await gen_table_dao.get_gen_table_by_name(auth.db, gen_table.sub_table_name)
if sub_table:
gen_table.sub_table = GenTableSchema(**CamelCaseUtil.transform_result(sub_table))
@classmethod
async def set_pk_column(cls, gen_table: GenTableSchema) -> None:
async def _set_pk_column(cls, gen_table: GenTableSchema) -> None:
"""
设置主键列信息
:param gen_table: 业务表信息
:return:
"""
for column in gen_table.columns:
if column.pk:
gen_table.pk_column = column
break
if gen_table.pk_column is None:
gen_table.pk_column = gen_table.columns[0]
if gen_table.tpl_category == GenConstant.TPL_SUB:
for column in gen_table.sub_table.columns:
if gen_table.columns:
for column in gen_table.columns:
if column.pk:
gen_table.sub_table.pk_column = column
gen_table.pk_column = column
break
if gen_table.sub_table.columns is None:
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
@@ -464,6 +546,10 @@ class GenTableService:
:param edit_gen_table: 编辑业务表对象
"""
if edit_gen_table.tpl_category == GenConstant.TPL_TREE:
# 检查params是否为None
if edit_gen_table.params is None:
raise CustomException(msg='树表参数不能为空')
params_obj = edit_gen_table.params.model_dump(by_alias=True)
if GenConstant.TREE_CODE not in params_obj:
@@ -472,11 +558,11 @@ class GenTableService:
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='子表关联的外键名不能为空')
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='子表关联的外键名不能为空')
class GenTableColumnService:
@@ -485,17 +571,22 @@ class GenTableColumnService:
"""
@classmethod
async def get_gen_table_column_list_by_table_id_services(cls, query_db: AsyncSession, table_id: int):
async def get_gen_table_column_list_by_table_id_services(cls, auth: AuthSchema, table_id: int):
"""
获取业务表字段列表信息service
:param query_db: orm对象
:param auth: 认证信息
:param table_id: 业务表格id
:return: 业务表字段列表信息对象
"""
gen_table_column_list_result = await GenTableColumnDao.get_gen_table_column_list_by_table_id(query_db, table_id)
# 确保db是AsyncSession类型
if not isinstance(auth.db, AsyncSession):
raise CustomException(msg='数据库会话类型不正确')
gen_table_column_dao = GenTableColumnDao(auth=auth)
gen_table_column_list_result = await gen_table_column_dao.get_gen_table_column_list_by_table_id(auth.db, table_id)
return [
GenTableColumnSchema(**gen_table_column)
for gen_table_column in CamelCaseUtil.transform_result(gen_table_column_list_result)
]
]