diff --git a/backend/app/generator/model/gen_business.py b/backend/app/generator/model/gen_business.py index e27e8d2a..3772d16e 100644 --- a/backend/app/generator/model/gen_business.py +++ b/backend/app/generator/model/gen_business.py @@ -19,7 +19,7 @@ class GenBusiness(Base): table_simple_name_zh: Mapped[str] = mapped_column(String(255), comment='表名称(中文简称)') table_comment: Mapped[str | None] = mapped_column(String(255), default=None, comment='表描述') # relate_model_fk: Mapped[int | None] = mapped_column(default=None, comment='关联表外键') - schema_name: Mapped[str | None] = mapped_column(String(255), default=None, comment='Schema 名称 (默认为英文表驼峰)') + schema_name: Mapped[str | None] = mapped_column(String(255), default=None, comment='Schema 名称 (默认为英文表名称)') have_datetime_column: Mapped[bool] = mapped_column(default=True, comment='是否存在默认时间列') api_version: Mapped[str] = mapped_column(String(20), default='v1', comment='代码生成 api 版本,默认为 v1') gen_path: Mapped[str | None] = mapped_column(String(255), default=None, comment='代码生成路径(默认为 app 根路径)') diff --git a/backend/app/generator/schema/gen_business.py b/backend/app/generator/schema/gen_business.py index 5596bd7c..9ffe9b0c 100644 --- a/backend/app/generator/schema/gen_business.py +++ b/backend/app/generator/schema/gen_business.py @@ -3,7 +3,6 @@ from datetime import datetime from pydantic import ConfigDict, Field, model_validator -from pydantic.alias_generators import to_pascal from backend.common.schema import SchemaBase @@ -23,7 +22,7 @@ class GenBusinessSchemaBase(SchemaBase): @model_validator(mode='after') def check_schema_name(self): if self.schema_name is None: - self.schema_name = to_pascal(self.table_name_en) + self.schema_name = self.table_name_en return self diff --git a/backend/app/generator/service/gen_model_service.py b/backend/app/generator/service/gen_model_service.py index e244674c..eee4d3cc 100644 --- a/backend/app/generator/service/gen_model_service.py +++ b/backend/app/generator/service/gen_model_service.py @@ -24,7 +24,8 @@ class GenModelService: if gen_models: if obj.name in [model.name for model in gen_models]: raise errors.ForbiddenError(msg='禁止添加相同列到模型表') - await gen_model_dao.create(db, obj, {'pd_type': sql_type_to_pydantic(obj.type)}) + 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: @@ -33,7 +34,8 @@ class GenModelService: if 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, {'pd_type': sql_type_to_pydantic(obj.type)}) + pd_type = sql_type_to_pydantic(obj.type) + count = await gen_model_dao.update(db, pk, obj, {'pd_type': pd_type}) return count @staticmethod diff --git a/backend/templates/py/model.jinja b/backend/templates/py/model.jinja index 53bae737..9eabe680 100644 --- a/backend/templates/py/model.jinja +++ b/backend/templates/py/model.jinja @@ -1,10 +1,12 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- -from sqlalchemy import String +import sqlalchemy as sa + from sqlalchemy.orm import Mapped, mapped_column from backend.common.model import {% if have_datetime_column %}Base{% else %}MappedBase{% endif %}, id_key + class {{ table_name_class }}({% if have_datetime_column %}Base{% else %}MappedBase{% endif %}): """{{ table_name_zh }}""" @@ -12,5 +14,40 @@ 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.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 }}) + {{ model.name }}: + {%- if model.is_nullable %} Mapped[{{ model.pd_type }} | None] + {%- else %} Mapped[{{ model.pd_type }}] + {%- endif %} = mapped_column( + {%- if model.type == 'String' -%} + sa.String({{ model.length }}) + {%- else -%} + sa.{{ model.type }}() + {%- endif -%}, default= + {%- if model.is_nullable and model.default == None -%} + None + {%- else -%} + {%- if model.default != None -%} + {{ model.default }} + {%- else -%} + {%- if model.pd_type == 'str' -%} + '' + {%- elif model.pd_type == 'int' -%} + 0 + {%- elif model.pd_type == 'bytes' -%} + b'' + {%- elif model.pd_type == 'bool' -%} + True + {%- elif model.pd_type == 'float' -%} + 0.0 + {%- elif model.pd_type == 'dict' -%} + {} + {%- elif model.pd_type == 'date' or model.pd_type == 'datetime' -%} + timezone.now() + {%- elif model.pd_type == 'List[str]' -%} + () + {%- else -%} + '' + {%- endif -%} + {%- endif -%} + {%- endif -%}, sort_order={{ model.sort }}, comment='{{ model.comment }}') {% endfor %} diff --git a/backend/templates/py/schema.jinja b/backend/templates/py/schema.jinja index 9897b71a..642806b3 100644 --- a/backend/templates/py/schema.jinja +++ b/backend/templates/py/schema.jinja @@ -1,17 +1,18 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- +{% if have_datetime_column %} from datetime import datetime - -from pydantic import ConfigDict, Field +{% endif %} +from pydantic import ConfigDict from backend.common.schema import SchemaBase class {{ schema_name }}SchemaBase(SchemaBase): {% for model in models %} - {{ model.name }}: {% if model.nullable %}{{ model.type }} | None = None{% else %}{{ model.type }}{% endif %} - {% endfor %} + {{ model.name }}: {% if model.nullable %}{{ model.pd_type }} | None = None{% else %}{{ model.pd_type }}{% endif %} + {% endfor %} class Create{{ schema_name }}Param({{ schema_name }}SchemaBase): diff --git a/backend/utils/gen_template.py b/backend/utils/gen_template.py index 90623fce..f71a9a91 100644 --- a/backend/utils/gen_template.py +++ b/backend/utils/gen_template.py @@ -11,9 +11,10 @@ class GenTemplate: def __init__(self): self.env = Environment( loader=FileSystemLoader(JINJA2_TEMPLATE_DIR), - autoescape=select_autoescape(['html', 'xml', 'jinja']), + autoescape=select_autoescape(enabled_extensions=['jinja']), trim_blocks=True, lstrip_blocks=True, + keep_trailing_newline=True, enable_async=True, )