From 812b8a0fb76636182e55197e688d3f4ffc5e8011 Mon Sep 17 00:00:00 2001 From: Wu Clan Date: Fri, 12 Jul 2024 21:21:01 +0800 Subject: [PATCH] Update code generation model column type storage (#352) --- backend/app/generator/crud/crud_gen_model.py | 10 +- backend/app/generator/model/gen_model.py | 3 +- backend/app/generator/schema/gen_model.py | 32 ++---- .../generator/service/gen_model_service.py | 15 +-- backend/app/generator/service/gen_service.py | 6 +- backend/common/enums.py | 99 +++++++++---------- backend/templates/py/model.jinja | 2 +- backend/utils/gen_template.py | 8 -- backend/utils/type_conversion.py | 95 ++++++++++++++++++ 9 files changed, 166 insertions(+), 104 deletions(-) create mode 100644 backend/utils/type_conversion.py diff --git a/backend/app/generator/crud/crud_gen_model.py b/backend/app/generator/crud/crud_gen_model.py index 7df2755b..c3b030d2 100644 --- a/backend/app/generator/crud/crud_gen_model.py +++ b/backend/app/generator/crud/crud_gen_model.py @@ -20,13 +20,13 @@ class CRUDGenModel(CRUDPlus[GenModel]): """ return await self.select_model_by_column(db, 'name', name) - async def get_by_business_id(self, db: AsyncSession, business_id: int) -> Sequence[GenModel]: + async def get_all_by_business_id(self, db: AsyncSession, business_id: int) -> Sequence[GenModel]: gen_model = await db.execute( select(self.model).where(self.model.gen_business_id == business_id).order_by(self.model.sort) ) return gen_model.scalars().all() - async def create(self, db: AsyncSession, obj_in: CreateGenModelParam) -> None: + async def create(self, db: AsyncSession, obj_in: CreateGenModelParam, **kwargs) -> None: """ 创建代码生成模型表 @@ -34,9 +34,9 @@ class CRUDGenModel(CRUDPlus[GenModel]): :param obj_in: :return: """ - return await self.create_model(db, obj_in) + return await self.create_model(db, obj_in, **kwargs) - async def update(self, db: AsyncSession, pk: int, obj_in: UpdateGenModelParam) -> int: + async def update(self, db: AsyncSession, pk: int, obj_in: UpdateGenModelParam, **kwargs) -> int: """ 更细代码生成模型表 @@ -45,7 +45,7 @@ class CRUDGenModel(CRUDPlus[GenModel]): :param obj_in: :return: """ - return await self.update_model(db, pk, obj_in) + return await self.update_model(db, pk, obj_in, **kwargs) async def delete(self, db: AsyncSession, pk: int) -> int: """ diff --git a/backend/app/generator/model/gen_model.py b/backend/app/generator/model/gen_model.py index d9d11570..92fbc2b2 100644 --- a/backend/app/generator/model/gen_model.py +++ b/backend/app/generator/model/gen_model.py @@ -16,7 +16,8 @@ class GenModel(DataClassBase): id: Mapped[id_key] = mapped_column(init=False) name: Mapped[str] = mapped_column(String(50), comment='列名称') comment: Mapped[str | None] = mapped_column(String(255), default=None, comment='列描述') - type: Mapped[str] = mapped_column(String(20), default='string', comment='列类型') + type: Mapped[str] = mapped_column(String(20), default='str', comment='SQLA 模型列类型') + pd_type: Mapped[str] = mapped_column(String(20), default='str', comment='列类型对应的 pydantic 类型') default: Mapped[str | None] = mapped_column(String(50), default=None, comment='列默认值') sort: Mapped[int | None] = mapped_column(default=1, comment='列排序') length: Mapped[int] = mapped_column(default=0, comment='列长度') diff --git a/backend/app/generator/schema/gen_model.py b/backend/app/generator/schema/gen_model.py index 813ef498..e29bd06b 100644 --- a/backend/app/generator/schema/gen_model.py +++ b/backend/app/generator/schema/gen_model.py @@ -2,14 +2,15 @@ # -*- coding: utf-8 -*- from pydantic import ConfigDict, Field, field_validator -from backend.common.enums import GenModelType +from backend.common.enums import GenModelColumnType from backend.common.schema import SchemaBase +from backend.utils.type_conversion import sql_type_to_sqlalchemy class GenModelSchemaBase(SchemaBase): name: str comment: str | None = None - type: GenModelType = Field(GenModelType.String, description='模型 column 类型') + type: GenModelColumnType = Field(GenModelColumnType.VARCHAR) default: str | None = None sort: int length: int @@ -19,30 +20,8 @@ class GenModelSchemaBase(SchemaBase): @field_validator('type') @classmethod - def sql_type_to_python(cls, v: GenModelType): - type_mapping = { - GenModelType.CHAR: 'str', - GenModelType.VARCHAR: 'str', - GenModelType.String: 'str', - GenModelType.TEXT: 'str', - GenModelType.Text: 'str', - GenModelType.LONGTEXT: 'str', - GenModelType.UnicodeText: 'str', - GenModelType.INT: 'int', - GenModelType.INTEGER: 'int', - GenModelType.Integer: 'int', - GenModelType.BigInteger: 'int', - GenModelType.SmallInteger: 'int', - GenModelType.BIGINT: 'int', - GenModelType.SMALLINT: 'int', - GenModelType.FLOAT: 'float', - GenModelType.Float: 'float', - GenModelType.Boolean: 'bool', - GenModelType.DECIMAL: 'decimal', - GenModelType.UUID: 'UUID', - GenModelType.Uuid: 'UUID', - } - return type_mapping.get(v) or v + def type_update(cls, v): + return sql_type_to_sqlalchemy(v) class CreateGenModelParam(GenModelSchemaBase): @@ -57,3 +36,4 @@ class GetGenModelListDetails(GenModelSchemaBase): model_config = ConfigDict(from_attributes=True) id: int + pd_type: str diff --git a/backend/app/generator/service/gen_model_service.py b/backend/app/generator/service/gen_model_service.py index 08df933f..e244674c 100644 --- a/backend/app/generator/service/gen_model_service.py +++ b/backend/app/generator/service/gen_model_service.py @@ -7,32 +7,33 @@ from backend.app.generator.model import GenModel from backend.app.generator.schema.gen_model import CreateGenModelParam, UpdateGenModelParam from backend.common.exception import errors from backend.database.db_mysql import async_db_session +from backend.utils.type_conversion import sql_type_to_pydantic class GenModelService: @staticmethod async def get_by_business(*, business_id: int) -> Sequence[GenModel]: async with async_db_session() as db: - gen_model = await gen_model_dao.get_by_business_id(db, business_id) + gen_model = await gen_model_dao.get_all_by_business_id(db, business_id) return gen_model @staticmethod async def create(*, obj: CreateGenModelParam) -> None: async with async_db_session.begin() as db: - gen_models = await gen_model_dao.get_by_business_id(db, obj.gen_business_id) + gen_models = await gen_model_dao.get_all_by_business_id(db, obj.gen_business_id) if gen_models: - if obj.name in [name.name for name in gen_models]: + if obj.name in [model.name for model in gen_models]: raise errors.ForbiddenError(msg='禁止添加相同列到模型表') - await gen_model_dao.create(db, obj) + await gen_model_dao.create(db, obj, {'pd_type': sql_type_to_pydantic(obj.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_by_business_id(obj.gen_business_id) + gen_models = await gen_model_dao.get_all_by_business_id(obj.gen_business_id) if gen_models: - if obj.name in [name.name for name in gen_models]: + if obj.name in [model.name for model in gen_models]: raise errors.ForbiddenError(msg='禁止添加相同列到模型表') - count = await gen_model_dao.update(db, pk, obj) + count = await gen_model_dao.update(db, pk, obj, {'pd_type': sql_type_to_pydantic(obj.type)}) return count @staticmethod diff --git a/backend/app/generator/service/gen_service.py b/backend/app/generator/service/gen_service.py index e82f3aef..97644a85 100644 --- a/backend/app/generator/service/gen_service.py +++ b/backend/app/generator/service/gen_service.py @@ -17,7 +17,7 @@ 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 GenModelType +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 @@ -63,14 +63,14 @@ class GenService: await db.flush() column_info = await gen_dao.get_all_columns(db, table_schema, table_name) for column in column_info: - column_type = column[-1].split('(')[0].lower() + column_type = column[-1].split('(')[0].upper() model_data = { 'name': column[0], 'comment': column[-2], 'type': column_type, 'sort': column[-3], 'length': column[-1].split('(')[1][:-1] - if column_type == GenModelType.CHAR or column_type == GenModelType.VARCHAR + if column_type == GenModelColumnType.CHAR or column_type == GenModelColumnType.VARCHAR else 0, 'is_pk': column[1], 'is_nullable': column[2], diff --git a/backend/common/enums.py b/backend/common/enums.py index 9ce64ef2..79ed5c21 100644 --- a/backend/common/enums.py +++ b/backend/common/enums.py @@ -90,57 +90,50 @@ class UserSocialType(StrEnum): linuxdo = 'LinuxDo' -class GenModelType(StrEnum): - """代码生成模型类型""" +class GenModelColumnType(StrEnum): + """代码生成模型列类型""" - # 待优化 - # https://github.com/zy7y/dfs-generate/blob/master/dfs_generate/types_map.py - ARRAY = 'array' - BIGINT = 'bigint' - BigInteger = 'biginteger' - BINARY = 'binary' - BLOB = 'blob' - BOOLEAN = 'boolean' - Boolean = 'boolean' - CHAR = 'char' - CLOB = 'clob' - DATE = 'date' - Date = 'date' - DATETIME = 'datetime' - DateTime = 'datetime' - DECIMAL = 'decimal' - DOUBLE = 'double' - Double = 'double' - DOUBLE_PRECISION = 'double_precision' - Enum = 'enum' - FLOAT = 'float' - Float = 'float' - INT = 'int' - INTEGER = 'integer' - Integer = 'integer' - Interval = 'interval' - JSON = 'json' - LargeBinary = 'largebinary' - LONGTEXT = 'longtext' - NCHAR = 'nchar' - NUMERIC = 'numeric' - Numeric = 'numeric' - NVARCHAR = 'nvarchar' - PickleType = 'pickletype' - REAL = 'real' - SMALLINT = 'smallint' - SmallInteger = 'smallinteger' - String = 'string' - TEXT = 'text' - Text = 'text' - TIME = 'time' - Time = 'time' - TIMESTAMP = 'timestamp' - TupleType = 'tupletype' - TypeDecorator = 'typedecorator' - Unicode = 'unicode' - UnicodeText = 'unicodetext' - UUID = 'uuid' - Uuid = 'uuid' - VARBINARY = 'varbinary' - VARCHAR = 'varchar' + BIGINT = 'BIGINT' + BINARY = 'BINARY' + BIT = 'BIT' + BLOB = 'BLOB' + BOOL = 'BOOL' + BOOLEAN = 'BOOLEAN' + CHAR = 'CHAR' + DATE = 'DATE' + DATETIME = 'DATETIME' + DECIMAL = 'DECIMAL' + DOUBLE = 'DOUBLE' + DOUBLE_PRECISION = 'DOUBLE PRECISION' + ENUM = 'ENUM' + FLOAT = 'FLOAT' + GEOMETRY = 'GEOMETRY' + GEOMETRYCOLLECTION = 'GEOMETRYCOLLECTION' + INT = 'INT' + INTEGER = 'INTEGER' + JSON = 'JSON' + LINESTRING = 'LINESTRING' + LONGBLOB = 'LONGBLOB' + LONGTEXT = 'LONGTEXT' + MEDIUMBLOB = 'MEDIUMBLOB' + MEDIUMINT = 'MEDIUMINT' + MEDIUMTEXT = 'MEDIUMTEXT' + MULTILINESTRING = 'MULTILINESTRING' + MULTIPOINT = 'MULTIPOINT' + MULTIPOLYGON = 'MULTIPOLYGON' + NUMERIC = 'NUMERIC' + POINT = 'POINT' + POLYGON = 'POLYGON' + REAL = 'REAL' + SERIAL = 'SERIAL' + SET = 'SET' + SMALLINT = 'SMALLINT' + TEXT = 'TEXT' + TIME = 'TIME' + TIMESTAMP = 'TIMESTAMP' + TINYBLOB = 'TINYBLOB' + TINYINT = 'TINYINT' + TINYTEXT = 'TINYTEXT' + VARBINARY = 'VARBINARY' + VARCHAR = 'VARCHAR' + YEAR = 'YEAR' diff --git a/backend/templates/py/model.jinja b/backend/templates/py/model.jinja index e4445f5a..53bae737 100644 --- a/backend/templates/py/model.jinja +++ b/backend/templates/py/model.jinja @@ -12,5 +12,5 @@ class {{ table_name_class }}({% if have_datetime_column %}Base{% else %}MappedBa id: Mapped[id_key] = mapped_column(init=False) {% for model in models %} - {{ model.name }}: {% if model.is_nullable %}Mapped[{{ model.type }} | None]{% else %}Mapped[{{ model.type }}]{% endif %} = mapped_column({% if model.type == 'str' and model.length != 0 %}{{ model_type_mapping.get(model.type) }}({{ model.length }}){% else %}{{ model_type_mapping.get(model.type) or model.type}}(){% endif %}, default={{ model.default }}, sort_order={{ model.sort }}, comment={{ model.comment }}) + {{ model.name }}: {% if model.is_nullable %}Mapped[{{ model.pd_type }} | None]{% else %}Mapped[{{ model.pd_type }}]{% endif %} = mapped_column({% if model.type == 'String' %}String({{ model.length }}){% else %}{{ model.type}}(){% endif %}, default={{ model.default }}, sort_order={{ model.sort }}, comment={{ model.comment }}) {% endfor %} diff --git a/backend/utils/gen_template.py b/backend/utils/gen_template.py index 567d2aef..90623fce 100644 --- a/backend/utils/gen_template.py +++ b/backend/utils/gen_template.py @@ -69,13 +69,6 @@ class GenTemplate: :param models: :return: """ - # python 类型对应的 sqlalchemy 类型 - model_type_mapping = { - 'str': 'String', - 'float': 'Float', - 'int': 'Integer', - 'bool': 'Boolean', - } return { 'app_name': business.app_name, 'table_name_en': to_snake(business.table_name_en), @@ -87,7 +80,6 @@ class GenTemplate: 'have_datetime_column': business.have_datetime_column, 'permission_sign': str(business.__tablename__.replace('_', ':')), 'models': models, - 'model_type_mapping': model_type_mapping, } diff --git a/backend/utils/type_conversion.py b/backend/utils/type_conversion.py new file mode 100644 index 00000000..3a5ed861 --- /dev/null +++ b/backend/utils/type_conversion.py @@ -0,0 +1,95 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from backend.common.enums import GenModelColumnType + + +def sql_type_to_sqlalchemy(typing: str) -> str: + """ + Converts a sql type to a SQLAlchemy type. + + :param typing: + :return: + """ + type_mapping = { + GenModelColumnType.BIGINT: 'BIGINT', + GenModelColumnType.BINARY: 'BINARY', + GenModelColumnType.BIT: 'BIT', + GenModelColumnType.BLOB: 'BLOG', + GenModelColumnType.BOOL: 'BOOLEAN', + GenModelColumnType.BOOLEAN: 'BOOLEAN', + GenModelColumnType.CHAR: 'CHAR', + GenModelColumnType.DATE: 'DATE', + GenModelColumnType.DATETIME: 'DATETIME', + GenModelColumnType.DECIMAL: 'DECIMAL', + GenModelColumnType.DOUBLE: 'DOUBLE', + GenModelColumnType.ENUM: 'ENUM', + GenModelColumnType.FLOAT: 'FLOAT', + GenModelColumnType.INT: 'INT', + GenModelColumnType.INTEGER: 'INTEGER', + GenModelColumnType.JSON: 'JSON', + GenModelColumnType.LONGBLOB: 'LONGBLOB', + GenModelColumnType.LONGTEXT: 'LONGTEXT', + GenModelColumnType.MEDIUMBLOB: 'MEDIUMBLOB', + GenModelColumnType.MEDIUMINT: 'MEDIUMINT', + GenModelColumnType.MEDIUMTEXT: 'MEDIUMTEXT', + GenModelColumnType.NUMERIC: 'NUMERIC', + GenModelColumnType.SET: 'SET', + GenModelColumnType.SMALLINT: 'SMALLINT', + GenModelColumnType.REAL: 'REAL', + GenModelColumnType.TEXT: 'TEXT', + GenModelColumnType.TIME: 'TIME', + GenModelColumnType.TIMESTAMP: 'TIMESTAMP', + GenModelColumnType.TINYBLOB: 'TINYBLOB', + GenModelColumnType.TINYINT: 'TINYINT', + GenModelColumnType.TINYTEXT: 'TINYTEXT', + GenModelColumnType.VARBINARY: 'VARBINARY', + GenModelColumnType.VARCHAR: 'String', + GenModelColumnType.YEAR: 'YEAR', + } + return type_mapping.get(typing, 'String') + + +def sql_type_to_pydantic(typing: str) -> str: + """ + Converts a sql type to a pydantic type. + + :param typing: + :return: + """ + type_mapping = { + GenModelColumnType.BIGINT: 'int', + GenModelColumnType.BINARY: 'bytes', + GenModelColumnType.BIT: 'bool', + GenModelColumnType.BLOB: 'bytes', + GenModelColumnType.BOOL: 'bool', + GenModelColumnType.BOOLEAN: 'bool', + GenModelColumnType.CHAR: 'str', + GenModelColumnType.DATE: 'date', + GenModelColumnType.DATETIME: 'datetime', + GenModelColumnType.DECIMAL: 'Decimal', + GenModelColumnType.DOUBLE: 'float', + GenModelColumnType.ENUM: 'Enum', + GenModelColumnType.FLOAT: 'float', + GenModelColumnType.INT: 'int', + GenModelColumnType.INTEGER: 'int', + GenModelColumnType.JSON: 'dict', + GenModelColumnType.LONGBLOB: 'bytes', + GenModelColumnType.LONGTEXT: 'str', + GenModelColumnType.MEDIUMBLOB: 'bytes', + GenModelColumnType.MEDIUMINT: 'int', + GenModelColumnType.MEDIUMTEXT: 'str', + GenModelColumnType.NUMERIC: 'NUMERIC', + GenModelColumnType.SET: 'List[str]', + GenModelColumnType.SMALLINT: 'int', + GenModelColumnType.REAL: 'float', + GenModelColumnType.TEXT: 'str', + GenModelColumnType.TIME: 'time', + GenModelColumnType.TIMESTAMP: 'datetime', + GenModelColumnType.TINYBLOB: 'bytes', + GenModelColumnType.TINYINT: 'int', + GenModelColumnType.TINYTEXT: 'str', + GenModelColumnType.VARBINARY: 'bytes', + GenModelColumnType.VARCHAR: 'str', + GenModelColumnType.YEAR: 'int', + } + return type_mapping.get(typing, 'str')