diff --git a/backend/app/generator/api/v1/gen.py b/backend/app/generator/api/v1/gen.py index 6a164797..f0e4880e 100644 --- a/backend/app/generator/api/v1/gen.py +++ b/backend/app/generator/api/v1/gen.py @@ -2,24 +2,29 @@ # -*- coding: utf-8 -*- from typing import Annotated -from fastapi import APIRouter, Path, Query +from fastapi import APIRouter, Body, Depends, Path, Query from fastapi.responses import StreamingResponse from backend.app.generator.conf import generator_settings -from backend.app.generator.schema.gen_business import CreateGenBusinessParam, UpdateGenBusinessParam -from backend.app.generator.schema.gen_model import CreateGenModelParam, UpdateGenModelParam +from backend.app.generator.schema.gen_business import ( + CreateGenBusinessParam, + GetGenBusinessListDetails, + UpdateGenBusinessParam, +) +from backend.app.generator.schema.gen_model import CreateGenModelParam, GetGenModelListDetails, UpdateGenModelParam from backend.app.generator.service.gen_business_service import gen_business_service from backend.app.generator.service.gen_model_service import gen_model_service from backend.app.generator.service.gen_service import gen_service from backend.common.response.response_schema import ResponseModel, response_base from backend.common.security.jwt import DependsJwtAuth +from backend.common.security.permission import RequestPermission from backend.common.security.rbac import DependsRBAC -from backend.utils.serializers import select_list_serialize +from backend.utils.serializers import select_as_dict, select_list_serialize router = APIRouter() -@router.get('/all', summary='获取所有代码生成业务', dependencies=[DependsJwtAuth]) +@router.get('/businesses/all', summary='获取所有代码生成业务', dependencies=[DependsJwtAuth]) async def get_all_businesses() -> ResponseModel: businesses = await gen_business_service.get_all() data = await select_list_serialize(businesses) @@ -28,17 +33,40 @@ async def get_all_businesses() -> ResponseModel: @router.get('/businesses/{pk}', summary='获取代码生成业务详情', dependencies=[DependsJwtAuth]) async def get_business(pk: Annotated[int, Path(...)]) -> ResponseModel: - data = await gen_service.get_business_and_model(pk=pk) + business = await gen_service.get_business_with_model(pk=pk) + data = GetGenBusinessListDetails(**await select_as_dict(business)) return await response_base.success(data=data) -@router.post('/businesses', summary='创建代码生成业务', deprecated=True, dependencies=[DependsRBAC]) +@router.get('/businesses/{pk}/models', summary='获取代码生成业务所有模型', dependencies=[DependsJwtAuth]) +async def get_business_models(pk: Annotated[int, Path(...)]) -> ResponseModel: + models = await gen_model_service.get_by_business(business_id=pk) + data = await select_list_serialize(models) + return await response_base.success(data=data) + + +@router.post( + '/businesses', + summary='创建代码生成业务', + deprecated=True, + dependencies=[ + Depends(RequestPermission('gen:code:business:add')), + DependsRBAC, + ], +) async def create_business(obj: CreateGenBusinessParam) -> ResponseModel: await gen_business_service.create(obj=obj) return await response_base.success() -@router.put('/businesses/{pk}', summary='更新代码生成业务', dependencies=[DependsRBAC]) +@router.put( + '/businesses/{pk}', + summary='更新代码生成业务', + dependencies=[ + Depends(RequestPermission('gen:code:business:edit')), + DependsRBAC, + ], +) async def update_business(pk: Annotated[int, Path(...)], obj: UpdateGenBusinessParam) -> ResponseModel: count = await gen_business_service.update(pk=pk, obj=obj) if count > 0: @@ -46,21 +74,49 @@ async def update_business(pk: Annotated[int, Path(...)], obj: UpdateGenBusinessP return await response_base.fail() -@router.delete('/businesses', summary='删除代码生成业务', dependencies=[DependsRBAC]) -async def delete_business(pk: Annotated[int, Query(...)]) -> ResponseModel: +@router.delete( + '/businesses/{pk}', + summary='删除代码生成业务', + dependencies=[ + Depends(RequestPermission('gen:code:business:del')), + DependsRBAC, + ], +) +async def delete_business(pk: Annotated[int, Path(...)]) -> ResponseModel: count = await gen_business_service.delete(pk=pk) if count > 0: return await response_base.success() return await response_base.fail() -@router.post('/models', summary='创建代码生成模型', dependencies=[DependsRBAC]) +@router.get('/models/{pk}', summary='获取代码生成模型详情', dependencies=[DependsJwtAuth]) +async def get_model(pk: Annotated[int, Path(...)]) -> ResponseModel: + model = await gen_model_service.get(pk=pk) + data = GetGenModelListDetails(**await select_as_dict(model)) + return await response_base.success(data=data) + + +@router.post( + '/models', + summary='创建代码生成模型', + dependencies=[ + Depends(RequestPermission('gen:code:model:add')), + DependsRBAC, + ], +) async def create_model(obj: CreateGenModelParam) -> ResponseModel: await gen_model_service.create(obj=obj) return await response_base.success() -@router.put('/models/{pk}', summary='更新代码生成模型', dependencies=[DependsRBAC]) +@router.put( + '/models/{pk}', + summary='更新代码生成模型', + dependencies=[ + Depends(RequestPermission('gen:code:model:edit')), + DependsRBAC, + ], +) async def update_model(pk: Annotated[int, Path(...)], obj: UpdateGenModelParam) -> ResponseModel: count = await gen_model_service.update(pk=pk, obj=obj) if count > 0: @@ -68,7 +124,14 @@ async def update_model(pk: Annotated[int, Path(...)], obj: UpdateGenModelParam) return await response_base.fail() -@router.delete('/models/{pk}', summary='删除代码生成模型', dependencies=[DependsRBAC]) +@router.delete( + '/models/{pk}', + summary='删除代码生成模型', + dependencies=[ + Depends(RequestPermission('gen:code:model:del')), + DependsRBAC, + ], +) async def delete_model(pk: Annotated[int, Path(...)]) -> ResponseModel: count = await gen_model_service.delete(pk=pk) if count > 0: @@ -82,11 +145,18 @@ async def get_all_tables(table_schema: Annotated[str, Query(..., description=' return await response_base.success(data=data) -@router.post('/import', summary='导入代码生成业务和模型列', dependencies=[DependsRBAC]) +@router.post( + '/import', + summary='导入代码生成业务和模型列', + dependencies=[ + Depends(RequestPermission('')), + DependsRBAC, + ], +) async def import_table( - app: Annotated[str, Query(..., description='应用名称,用于代码生成到指定 app')], - table_name: Annotated[str, Query(..., description='数据库表名')], - table_schema: Annotated[str, Query(..., description='数据库名')] = 'fba', + app: Annotated[str, Body(..., description='应用名称,用于代码生成到指定 app')], + table_name: Annotated[str, Body(..., description='数据库表名')], + table_schema: Annotated[str, Body(..., description='数据库名')] = 'fba', ) -> ResponseModel: await gen_service.import_business_and_model(app=app, table_schema=table_schema, table_name=table_name) return await response_base.success() @@ -98,13 +168,27 @@ async def preview_code(pk: Annotated[int, Path(..., description='业务ID')]) -> return await response_base.success(data=data) -@router.post('/generate/{pk}', summary='生成代码', description='文件磁盘写入,请谨慎操作', dependencies=[DependsRBAC]) +@router.get('/generate/{pk}/path', summary='获取代码生成路径', dependencies=[DependsJwtAuth]) +async def generate_path(pk: Annotated[int, Path(..., description='业务ID')]): + data = await gen_service.get_generate_path(pk=pk) + return await response_base.success(data=data) + + +@router.post( + '/generate/{pk}', + summary='代码生成', + description='文件磁盘写入,请谨慎操作', + dependencies=[ + Depends(RequestPermission('gen:code:generate')), + DependsRBAC, + ], +) async def generate_code(pk: Annotated[int, Path(..., description='业务ID')]) -> ResponseModel: await gen_service.generate(pk=pk) return await response_base.success() -@router.post('/download/{pk}', summary='下载代码', dependencies=[DependsRBAC]) +@router.get('/download/{pk}', summary='下载代码', dependencies=[DependsRBAC]) async def download_code(pk: Annotated[int, Path(..., description='业务ID')]): bio = await gen_service.download(pk=pk) return StreamingResponse( diff --git a/backend/app/generator/crud/crud_gen.py b/backend/app/generator/crud/crud_gen.py index ecef470e..27b06dbc 100644 --- a/backend/app/generator/crud/crud_gen.py +++ b/backend/app/generator/crud/crud_gen.py @@ -2,36 +2,47 @@ # -*- coding: utf-8 -*- from typing import Sequence -from sqlalchemy import Row, text +from sqlalchemy import Row, select, text from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import selectinload + +from backend.app.generator.model import GenBusiness class CRUDGen: + @staticmethod + async def get_business_with_model(db: AsyncSession, business_id: int) -> GenBusiness: + result = await db.execute( + select(GenBusiness).options(selectinload(GenBusiness.gen_model)).where(GenBusiness.id == business_id) + ) + data = result.scalars().first() + return data + @staticmethod async def get_all_tables(db: AsyncSession, table_schema: str) -> Sequence[str]: - t = text( + stmt = text( 'select table_name as table_name ' 'from information_schema.tables ' 'where table_name not like "sys_gen_%" ' 'and table_schema = :table_schema;' ).bindparams(table_schema=table_schema) - stmt = await db.execute(t) - return stmt.scalars().all() + result = await db.execute(stmt) + return result.scalars().all() @staticmethod async def get_table(db: AsyncSession, table_name: str) -> Row[tuple]: - t = text( + stmt = text( 'select table_name as table_name, table_comment as table_comment ' 'from information_schema.tables ' 'where table_name not like "sys_gen_%" ' 'and table_name = :table_name;' ).bindparams(table_name=table_name) - stmt = await db.execute(t) - return stmt.fetchone() + result = await db.execute(stmt) + return result.fetchone() @staticmethod async def get_all_columns(db: AsyncSession, table_schema: str, table_name: str) -> Sequence[Row[tuple]]: - t = text( + stmt = text( 'select column_name AS column_name, ' 'case when column_key = "PRI" then 1 else 0 end as is_pk, ' 'case when is_nullable = "NO" or column_key = "PRI" then 0 else 1 end as is_nullable, ' @@ -46,8 +57,8 @@ class CRUDGen: 'and column_name != "updated_time" ' 'order by sort;' ).bindparams(table_schema=table_schema, table_name=table_name) - stmt = await db.execute(t) - return stmt.fetchall() + result = await db.execute(stmt) + return result.fetchall() -gen_dao = CRUDGen() +gen_dao: CRUDGen = CRUDGen() diff --git a/backend/app/generator/crud/crud_gen_business.py b/backend/app/generator/crud/crud_gen_business.py index 6f539c08..1258c504 100644 --- a/backend/app/generator/crud/crud_gen_business.py +++ b/backend/app/generator/crud/crud_gen_business.py @@ -70,4 +70,4 @@ class CRUDGenBusiness(CRUDPlus[GenBusiness]): return await self.delete_model(db, pk) -gen_business_dao = CRUDGenBusiness(GenBusiness) +gen_business_dao: CRUDGenBusiness = CRUDGenBusiness(GenBusiness) diff --git a/backend/app/generator/crud/crud_gen_model.py b/backend/app/generator/crud/crud_gen_model.py index c3b030d2..3c9eb44b 100644 --- a/backend/app/generator/crud/crud_gen_model.py +++ b/backend/app/generator/crud/crud_gen_model.py @@ -11,9 +11,18 @@ from backend.app.generator.schema.gen_model import CreateGenModelParam, UpdateGe class CRUDGenModel(CRUDPlus[GenModel]): + async def get(self, db: AsyncSession, pk: int) -> GenModel | None: + """ + 获取代码生成模型列 + + :return: + """ + return await self.select_model_by_id(db, pk) + async def get_by_name(self, db: AsyncSession, name: str) -> GenModel | None: """ 通过 name 获取代码生成模型表 + :param db: :param name: :return: @@ -58,4 +67,4 @@ class CRUDGenModel(CRUDPlus[GenModel]): return await self.delete_model(db, pk) -gen_model_dao = CRUDGenModel(GenModel) +gen_model_dao: CRUDGenModel = CRUDGenModel(GenModel) diff --git a/backend/app/generator/model/gen_model.py b/backend/app/generator/model/gen_model.py index 75f36bab..a80631ea 100644 --- a/backend/app/generator/model/gen_model.py +++ b/backend/app/generator/model/gen_model.py @@ -15,7 +15,7 @@ class GenModel(DataClassBase): __tablename__ = 'sys_gen_model' id: Mapped[id_key] = mapped_column(init=False) - name: Mapped[str] = mapped_column(String(50), comment='列名称') + name: Mapped[str] = mapped_column(String(50), unique=True, comment='列名称') comment: Mapped[str | None] = mapped_column(String(255), default=None, comment='列描述') type: Mapped[str] = mapped_column(String(20), default='str', comment='SQLA 模型列类型') pd_type: Mapped[str] = mapped_column(String(20), default='str', comment='列类型对应的 pydantic 类型') diff --git a/backend/app/generator/schema/gen_business.py b/backend/app/generator/schema/gen_business.py index 35257d79..9c7c3ac9 100644 --- a/backend/app/generator/schema/gen_business.py +++ b/backend/app/generator/schema/gen_business.py @@ -4,6 +4,7 @@ from datetime import datetime from pydantic import ConfigDict, Field, model_validator +from backend.app.generator.schema.gen_model import GetGenModelListDetails from backend.common.schema import SchemaBase @@ -40,3 +41,4 @@ class GetGenBusinessListDetails(GenBusinessSchemaBase): id: int created_time: datetime updated_time: datetime | None = None + gen_model: list[GetGenModelListDetails] | None = None diff --git a/backend/app/generator/service/gen_model_service.py b/backend/app/generator/service/gen_model_service.py index f826fcea..8f157ca3 100644 --- a/backend/app/generator/service/gen_model_service.py +++ b/backend/app/generator/service/gen_model_service.py @@ -11,6 +11,12 @@ from backend.utils.type_conversion import sql_type_to_pydantic class GenModelService: + @staticmethod + async def get(*, pk: int) -> GenModel: + async with async_db_session() as db: + gen_model = await gen_model_dao.get(db, pk) + return gen_model + @staticmethod async def get_by_business(*, business_id: int) -> Sequence[GenModel]: async with async_db_session() as db: @@ -20,19 +26,19 @@ class GenModelService: @staticmethod async def create(*, obj: CreateGenModelParam) -> None: async with async_db_session.begin() as db: - gen_models = await gen_model_dao.get_all_by_business_id(db, obj.gen_business_id) - if gen_models: - if obj.name in [model.name for model in gen_models]: - raise errors.ForbiddenError(msg='禁止添加相同列到模型表') + gen_model = await gen_model_dao.get_by_name(db, obj.name) + if gen_model: + raise errors.ForbiddenError(msg='禁止添加相同列到模型表') pd_type = sql_type_to_pydantic(obj.type) await gen_model_dao.create(db, obj, pd_type=pd_type) @staticmethod async def update(*, pk: int, obj: UpdateGenModelParam) -> int: async with async_db_session.begin() as db: - gen_models = await gen_model_dao.get_all_by_business_id(obj.gen_business_id) - if gen_models: - if obj.name in [model.name for model in gen_models]: + model = await gen_model_dao.get(db, pk) + if obj.name != model.name: + model_check = await gen_model_dao.get_by_name(db, obj.name) + if model_check: raise errors.ForbiddenError(msg='禁止添加相同列到模型表') pd_type = sql_type_to_pydantic(obj.type) count = await gen_model_dao.update(db, pk, obj, pd_type=pd_type) diff --git a/backend/app/generator/service/gen_service.py b/backend/app/generator/service/gen_service.py index 94c08b55..3a3f5e7f 100644 --- a/backend/app/generator/service/gen_service.py +++ b/backend/app/generator/service/gen_service.py @@ -17,26 +17,20 @@ from backend.app.generator.crud.crud_gen_model import gen_model_dao from backend.app.generator.model import GenBusiness from backend.app.generator.schema.gen_business import CreateGenBusinessParam from backend.app.generator.schema.gen_model import CreateGenModelParam -from backend.app.generator.service.gen_business_service import gen_business_service from backend.app.generator.service.gen_model_service import gen_model_service from backend.common.enums import GenModelColumnType from backend.common.exception import errors from backend.core.path_conf import BasePath from backend.database.db_mysql import async_db_session from backend.utils.gen_template import gen_template -from backend.utils.serializers import select_as_dict, select_list_serialize class GenService: @staticmethod - async def get_business_and_model(*, pk: int) -> dict: - gen_business = await gen_business_service.get(pk=pk) - gen_models = await gen_model_service.get_by_business(business_id=pk) - business_data = await select_as_dict(gen_business) - if gen_models: - model_data = await select_list_serialize(gen_models) - business_data.update({'models': model_data}) - return business_data + async def get_business_with_model(*, pk: int) -> GenBusiness: + async with async_db_session() as db: + business = await gen_dao.get_business_with_model(db, pk) + return business @staticmethod async def get_tables(*, table_schema: str) -> Sequence[str]: @@ -102,6 +96,22 @@ class GenService: for tpl, code in tpl_code_map.items() } + @staticmethod + async def get_generate_path(*, pk: int) -> list: + async with async_db_session() as db: + business = await gen_business_dao.get(db, pk) + if not business: + raise errors.NotFoundError(msg='业务不存在') + gen_path = business.gen_path + if not gen_path: + # 伪加密路径 + gen_path = 'current-backend-app-path' + target_files = gen_template.get_code_gen_paths(business) + code_gen_paths = [] + for target_file in target_files: + code_gen_paths.append(os.path.join(gen_path, *target_file.split('/')[1:])) + return code_gen_paths + async def generate(self, *, pk: int) -> None: async with async_db_session() as db: business = await gen_business_dao.get(db, pk) diff --git a/backend/sql/init_test_data.sql b/backend/sql/init_test_data.sql index d13082c4..c3392ac8 100644 --- a/backend/sql/init_test_data.sql +++ b/backend/sql/init_test_data.sql @@ -33,7 +33,19 @@ VALUES (1, '测试', 'test', 0, 0, '', null, 0, null, null, 0, 0, 1, null, null (29, '删除', '', 0, 0, null, null, 2, null, 'sys:menu:del', 1, 1, 1, null, 26, '2024-01-07 12:01:48', null), (30, '系统监控', 'monitor', 0, 88, 'IconComputer', 'monitor', 0, null, null, 1, 1, 1, null, null, '2023-07-27 19:27:08', null), (31, 'Redis监控', 'Redis', 0, 0, null, 'redis', 1, '/monitor/redis/index.vue', 'sys:monitor:redis', 1, 1, 1, null, 30, '2023-07-27 19:28:03', null), - (32, '服务器监控', 'Server', 0, 0, null, 'server', 1, '/monitor/server/index.vue', 'sys:monitor:server', 1, 1, 1, null, 30, '2023-07-27 19:28:29', null); + (32, '服务器监控', 'Server', 0, 0, null, 'server', 1, '/monitor/server/index.vue', 'sys:monitor:server', 1, 1, 1, null, 30, '2023-07-27 19:28:29', null), + (33, '系统自动化', 'automation', 0, 777, 'IconCodeSquare', 'automation', 0, null, null, 1, 1, 1, null, null, '2024-07-27 02:06:20', '2024-07-27 02:18:52'), + (34, '代码生成', 'CodeGenerator', 0, 0, null, 'code-generator', 1, '/automation/generator/index.vue', null, 1, 1, 1, null, 33, '2024-07-27 12:24:54', null), + (35, '导入', '', 0, 0, null, null, 2, null, 'gen:code:import', 1, 1, 1, null, 34, '2024-08-04 12:49:58', null), + (36, '新增业务', '', 0, 0, null, null, 2, null, 'gen:code:business:add', 1, 1, 1, null, 34, '2024-08-04 12:51:29', null), + (37, '编辑业务', '', 0, 0, null, null, 2, null, 'gen:code:business:edit', 1, 1, 1, null, 34, '2024-08-04 12:51:45', null), + (48, '删除业务', '', 0, 0, null, null, 2, null, 'gen:code:business:del', 1, 1, 1, null, 34, '2024-08-04 12:52:05', null), + (49, '新增模型', '', 0, 0, null, null, 2, null, 'gen:code:model:add', 1, 1, 1, null, 34, '2024-08-04 12:52:28', null), + (50, '编辑模型', '', 0, 0, null, null, 2, null, 'gen:code:model:edit', 1, 1, 1, null, 34, '2024-08-04 12:52:45', null), + (51, '删除模型', '', 0, 0, null, null, 2, null, 'gen:code:model:del', 1, 1, 1, null, 34, '2024-08-04 12:52:59', null), + (52, '生成', '', 0, 0, null, null, 2, null, 'gen:code:generate', 1, 1, 1, null, 34, '2024-08-04 12:55:03', null), + (53, 'GitHub', 'github', 0, 8888, 'IconGithub', 'https://github.com/wu-clan', 0, null, null, 1, 1, 1, null, null, '2024-07-27 12:32:46', null), + (54, '赞助', 'sponsor', 0, 9999, 'IconFire', 'https://wu-clan.github.io/sponsor/', 0, null, null, 1, 1, 1, null, null, '2024-07-27 12:39:57', null); INSERT INTO fba.sys_role (id, name, data_scope, status, remark, created_time, updated_time) VALUES (1, 'test', 2, 1, null, '2023-06-26 17:13:45', null); diff --git a/backend/templates/py/model.jinja b/backend/templates/py/model.jinja index 4d6d8801..306d73de 100644 --- a/backend/templates/py/model.jinja +++ b/backend/templates/py/model.jinja @@ -1,6 +1,7 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- import sqlalchemy as sa +from sqlalchemy.dialects import mysql from sqlalchemy.orm import Mapped, mapped_column @@ -18,10 +19,10 @@ class {{ table_name_class }}({% if have_datetime_column %}Base{% else %}MappedBa {%- if model.is_nullable %} Mapped[{{ model.pd_type }} | None] {%- else %} Mapped[{{ model.pd_type }}] {%- endif %} = mapped_column( - {%- if model.type == 'String' -%} + {%- if model.type == 'VARCHAR' -%} sa.String({{ model.length }}) {%- else -%} - sa.{{ model.type }}() + mysql.{{ model.type }}() {%- endif -%}, default= {%- if model.is_nullable and model.default == None -%} None diff --git a/backend/utils/gen_template.py b/backend/utils/gen_template.py index 6f4f3540..6be3782f 100644 --- a/backend/utils/gen_template.py +++ b/backend/utils/gen_template.py @@ -45,10 +45,12 @@ class GenTemplate: f'{generator_settings.TEMPLATE_BACKEND_DIR_NAME}/service.jinja', ] - def get_code_gen_path(self, tpl_path: str, business: GenBusiness) -> str: + @staticmethod + def get_code_gen_paths(business: GenBusiness) -> list[str]: """ - 获取代码生成路径 + 获取代码生成路径列表 + :param business: :return: """ app_name = business.app_name @@ -60,6 +62,17 @@ class GenTemplate: f'{generator_settings.TEMPLATE_BACKEND_DIR_NAME}/{app_name}/schema/{module_name}.py', f'{generator_settings.TEMPLATE_BACKEND_DIR_NAME}/{app_name}/service/{module_name}_service.py', ] + return target_files + + def get_code_gen_path(self, tpl_path: str, business: GenBusiness) -> str: + """ + 获取代码生成路径 + + :param tpl_path: + :param business: + :return: + """ + target_files = self.get_code_gen_paths(business) code_gen_path_mapping = dict(zip(self.get_template_paths(), target_files)) return code_gen_path_mapping[tpl_path] diff --git a/backend/utils/type_conversion.py b/backend/utils/type_conversion.py index 3a5ed861..3ac6f397 100644 --- a/backend/utils/type_conversion.py +++ b/backend/utils/type_conversion.py @@ -43,10 +43,10 @@ def sql_type_to_sqlalchemy(typing: str) -> str: GenModelColumnType.TINYINT: 'TINYINT', GenModelColumnType.TINYTEXT: 'TINYTEXT', GenModelColumnType.VARBINARY: 'VARBINARY', - GenModelColumnType.VARCHAR: 'String', + GenModelColumnType.VARCHAR: 'VARCHAR', GenModelColumnType.YEAR: 'YEAR', } - return type_mapping.get(typing, 'String') + return type_mapping.get(typing, 'VARCHAR') def sql_type_to_pydantic(typing: str) -> str: