fix(gencode): 修复数据库类型映射和导入处理问题

refactor: 优化数据库类型处理逻辑,移除COLLATE和UNSIGNED标记
feat: 添加PostgreSQL数组类型支持
docs: 更新README添加全类型测试表SQL
style: 统一datetime导入方式
This commit is contained in:
zhangtao
2026-02-03 18:12:11 +08:00
parent 563e1b20b9
commit 9d97c5bb31
6 changed files with 455 additions and 101 deletions
+223
View File
@@ -170,3 +170,226 @@ uv run ruff check --watch
---
❤️ **感谢您的关注和支持!** 如果这个项目对您有帮助,请给我们一个 ⭐️ Star!
---
## mysql 全类型测试表
```sql
CREATE TABLE `gen_all_types_demo` (
`tinyint_field` TINYINT NOT NULL COMMENT 'TINYINT类型',
`tinyint_unsigned_field` TINYINT UNSIGNED NOT NULL COMMENT 'TINYINT UNSIGNED类型',
`smallint_field` SMALLINT NOT NULL COMMENT 'SMALLINT类型',
`smallint_unsigned_field` SMALLINT UNSIGNED NOT NULL COMMENT 'SMALLINT UNSIGNED类型',
`mediumint_field` MEDIUMINT NOT NULL COMMENT 'MEDIUMINT类型',
`mediumint_unsigned_field` MEDIUMINT UNSIGNED NOT NULL COMMENT 'MEDIUMINT UNSIGNED类型',
`int_field` INT NOT NULL COMMENT 'INT类型',
`int_unsigned_field` INT UNSIGNED NOT NULL COMMENT 'INT UNSIGNED类型',
`bigint_field` BIGINT NOT NULL COMMENT 'BIGINT类型',
`bigint_unsigned_field` BIGINT UNSIGNED NOT NULL COMMENT 'BIGINT UNSIGNED类型',
`float_field` FLOAT NOT NULL COMMENT 'FLOAT类型',
`double_field` DOUBLE NOT NULL COMMENT 'DOUBLE类型',
`decimal_field` DECIMAL(10,2) NOT NULL COMMENT 'DECIMAL类型',
`decimal_unsigned_field` DECIMAL(10,2) UNSIGNED NOT NULL COMMENT 'DECIMAL UNSIGNED类型',
`numeric_field` NUMERIC(10,2) NOT NULL COMMENT 'NUMERIC类型',
`bit_field` BIT(8) NOT NULL COMMENT 'BIT类型',
`char_field` CHAR(32) NOT NULL COMMENT 'CHAR类型',
`varchar_field` VARCHAR(255) NOT NULL COMMENT 'VARCHAR类型',
`binary_field` BINARY(32) NOT NULL COMMENT 'BINARY类型',
`varbinary_field` VARBINARY(255) NOT NULL COMMENT 'VARBINARY类型',
`tinyblob_field` TINYBLOB COMMENT 'TINYBLOB类型',
`blob_field` BLOB COMMENT 'BLOB类型',
`mediumblob_field` MEDIUMBLOB COMMENT 'MEDIUMBLOB类型',
`longblob_field` LONGBLOB COMMENT 'LONGBLOB类型',
`tinytext_field` TINYTEXT COMMENT 'TINYTEXT类型',
`text_field` TEXT COMMENT 'TEXT类型',
`mediumtext_field` MEDIUMTEXT COMMENT 'MEDIUMTEXT类型',
`longtext_field` LONGTEXT COMMENT 'LONGTEXT类型',
`enum_field` ENUM('active','inactive','pending') NOT NULL DEFAULT 'pending' COMMENT 'ENUM类型',
`set_field` SET('read','write','execute') NOT NULL DEFAULT '' COMMENT 'SET类型',
`date_field` DATE NOT NULL COMMENT 'DATE类型',
`time_field` TIME NOT NULL COMMENT 'TIME类型',
`datetime_field` DATETIME NOT NULL COMMENT 'DATETIME类型',
`timestamp_field` TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP COMMENT 'TIMESTAMP类型',
`year_field` YEAR NOT NULL COMMENT 'YEAR类型',
`json_field` JSON COMMENT 'JSON类型',
`id` BIGINT NOT NULL AUTO_INCREMENT COMMENT '主键ID',
`uuid` VARCHAR(64) NOT NULL COMMENT 'UUID全局唯一标识',
`status` VARCHAR(10) NOT NULL DEFAULT '0' COMMENT '是否启用(0:启用 1:禁用)',
`description` TEXT COMMENT '备注/描述',
`created_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间',
`updated_time` DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP COMMENT '更新时间',
`created_id` BIGINT DEFAULT NULL COMMENT '创建人ID',
`updated_id` BIGINT DEFAULT NULL COMMENT '更新人ID',
PRIMARY KEY (`id`),
UNIQUE KEY `uuid` (`uuid`),
KEY `ix_gen_all_types_demo_created_id` (`created_id`),
KEY `ix_gen_all_types_demo_updated_id` (`updated_id`),
KEY `ix_gen_all_types_demo_status` (`status`),
KEY `ix_gen_all_types_demo_created_time` (`created_time`)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci COMMENT='全类型测试表';
```
## postgresql 全类型测试表
```sql
-- PostgreSQL 全类型测试表
CREATE TABLE gen_all_types_demo (
-- 整数类型
smallint_field SMALLINT NOT NULL,
integer_field INTEGER NOT NULL,
bigint_field BIGINT NOT NULL,
-- 浮点类型
real_field REAL NOT NULL,
double_precision_field DOUBLE PRECISION NOT NULL,
numeric_field NUMERIC(10,2) NOT NULL,
decimal_field DECIMAL(10,2) NOT NULL,
-- 字符串类型
char_field CHAR(32) NOT NULL,
varchar_field VARCHAR(255) NOT NULL,
text_field TEXT NOT NULL,
-- 二进制类型
bytea_field BYTEA,
-- 日期时间类型
date_field DATE NOT NULL,
time_field TIME NOT NULL,
time_with_tz_field TIMESTAMP WITH TIME ZONE NOT NULL,
time_without_tz_field TIMESTAMP WITHOUT TIME ZONE NOT NULL,
timestamp_field TIMESTAMP NOT NULL,
timestamp_with_tz_field TIMESTAMP WITH TIME ZONE NOT NULL,
timestamp_without_tz_field TIMESTAMP WITHOUT TIME ZONE NOT NULL,
interval_field INTERVAL,
-- 布尔类型
boolean_field BOOLEAN NOT NULL,
-- JSON类型
json_field JSON,
jsonb_field JSONB,
-- 其他类型
uuid_field UUID,
inet_field INET,
cidr_field CIDR,
macaddr_field MACADDR,
-- 几何类型
point_field POINT,
line_field LINE,
lseg_field LSEG,
box_field BOX,
path_field PATH,
polygon_field POLYGON,
circle_field CIRCLE,
-- 位类型
bit_field BIT(8) NOT NULL,
bit_varying_field VARBIT(8) NOT NULL,
-- 文本搜索类型
tsvector_field TSVECTOR,
tsquery_field TSQUERY,
-- XML类型
xml_field XML,
-- 数组类型
array_field INTEGER[],
-- 范围类型
range_field INT4RANGE,
-- 货币类型
money_field MONEY,
-- 对象标识符类型
oid_field OID,
regproc_field REGPROC,
regclass_field REGCLASS,
regtype_field REGTYPE,
regrole_field REGROLE,
regnamespace_field REGNAMESPACE,
-- 常用字段
id BIGSERIAL PRIMARY KEY,
uuid VARCHAR(64) NOT NULL UNIQUE,
status VARCHAR(10) NOT NULL DEFAULT '0',
description TEXT,
created_time TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_time TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
created_id BIGINT,
updated_id BIGINT
);
```
## mysql类型
INT
VARCHAR
CHAR
DATETIME
TIMESTAMP
DATE
BIT
FLOAT
DOUBLE
DECIMAL
BIGINT
TEXT
JSON
BLOB
BINARY
ENUM
SET
TINYINT
SMALLINT
MEDIUMINT
TIME
YEAR
VARBINARY
TINYBLOB
MEDIUMBLOB
LONGBLOB
TINYTEXT
MEDIUMTEXT
LONGTEXT
GEOMETRY
POINT
LINESTRING
POLYGON
MULTIPOINT
MULTILINESTRING
MULTIPOLYGON
GEOMETRYCOLLECTION
## pg类型
INTEGER
VARCHAR
CHAR
TIMESTAMP
DATE
BOOLEAN
FLOAT
TEXT
JSON
BLOB
SMALLINT
BIGINT
REAL
DOUBLE PRECISION
BYTEA
XML
UUID
ARRAY
NUMERIC
MONEY
INTERVAL
CIDR
INET
MACADDR
+101 -70
View File
@@ -514,6 +514,7 @@ class GenConstant:
"real": "Float",
"double precision": "Float",
"numeric": "Numeric",
"decimal": "Numeric",
"character varying": "String",
"varchar": "String",
"character": "String",
@@ -619,33 +620,97 @@ class GenConstant:
# 数据库类型与python类型映射
DB_TO_PYTHON = (
{
"boolean": "bool",
# MySQL 整数类型
"tinyint": "int",
"smallint": "int",
"mediumint": "int",
"int": "int",
"integer": "int",
"int4": "int",
"bigint": "int",
"real": "float",
"double precision": "float",
# MySQL 浮点类型
"float": "float",
"double": "float",
"decimal": "Decimal",
"numeric": "Decimal",
"character varying": "str",
# MySQL 字符串类型
"char": "str",
"varchar": "str",
"character": "str",
"tinytext": "str",
"text": "str",
"bytea": "bytes",
"mediumtext": "str",
"longtext": "str",
# MySQL 二进制类型
"binary": "bytes",
"varbinary": "bytes",
"tinyblob": "bytes",
"blob": "bytes",
"mediumblob": "bytes",
"longblob": "bytes",
# MySQL 日期时间类型
"date": "date",
"time": "time",
"time with time zone": "time",
"time without time zone": "time",
"datetime": "datetime",
"timestamp": "datetime",
"year": "int",
# MySQL 其他类型
"json": "dict",
"enum": "str",
"set": "str",
"bit": "int",
# MySQL 空间数据类型
"geometry": "bytes",
"linestring": "bytes",
"multipoint": "bytes",
"multilinestring": "bytes",
"multipolygon": "bytes",
"geometrycollection": "bytes",
# PostgreSQL 整数类型
"int2": "int",
"int4": "int",
"int8": "int",
# PostgreSQL 浮点类型
"real": "float",
"double precision": "float",
"float8": "float",
# PostgreSQL 字符串类型
"character": "str",
"character varying": "str",
# PostgreSQL 二进制类型
"bytea": "bytes",
# PostgreSQL 日期时间类型
"time with time zone": "time",
"timetz": "time",
"time without time zone": "time",
"timestamptz": "datetime",
"timestamp with time zone": "datetime",
"timestamp without time zone": "datetime",
"interval": "timedelta",
"json": "dict",
# PostgreSQL 布尔类型
"boolean": "bool",
"bool": "bool",
# PostgreSQL JSON类型
"jsonb": "dict",
# PostgreSQL 其他类型
"uuid": "str",
"inet": "str",
"cidr": "str",
"macaddr": "str",
# PostgreSQL 几何类型(覆盖MySQL的映射)
"point": "list",
"line": "list",
"lseg": "list",
@@ -653,83 +718,49 @@ class GenConstant:
"path": "list",
"polygon": "list",
"circle": "list",
"bit": "int",
# PostgreSQL 位类型
"bit varying": "int",
"varbit": "int",
# PostgreSQL 文本搜索类型
"tsvector": "str",
"tsquery": "str",
# PostgreSQL XML类型
"xml": "str",
# PostgreSQL 数组类型
"array": "list",
"composite": "dict",
"enum": "str",
# PostgreSQL 范围类型
"range": "list",
"int4range": "list",
"int8range": "list",
"tsrange": "list",
"tstzrange": "list",
"daterange": "list",
# PostgreSQL 货币类型
"money": "Decimal",
"pg_lsn": "int",
"txid_snapshot": "str",
# PostgreSQL 对象标识符类型
"oid": "int",
"regproc": "str",
"regclass": "str",
"regtype": "str",
"regrole": "str",
"regnamespace": "str",
# PostgreSQL 向量类型
"int2vector": "list",
"oidvector": "list",
# PostgreSQL 其他内部类型
"pg_lsn": "int",
"txid_snapshot": "str",
"pg_node_tree": "str",
}
if settings.DATABASE_TYPE == "postgres"
else {
# 布尔类型(特殊处理tinyint(1))
"TINYINT": "bool",
# 数值类型
"SMALLINT": "int",
"MEDIUMINT": "int",
"INT": "int",
"INTEGER": "int",
"BIGINT": "int",
"FLOAT": "float",
"DOUBLE": "float",
"NUMERIC": "float",
"DECIMAL": "Decimal",
"BIT": "int",
# 日期和时间类型
"DATE": "datetime.date",
"TIME": "datetime.time",
"DATETIME": "datetime.datetime",
"TIMESTAMP": "datetime.datetime",
"YEAR": "int",
"TINYINT UNSIGNED": "int", # 无符号小整数类型
# 布尔类型
"BOOLEAN": "bool",
"BOOL": "bool", # 布尔类型,通常与 BOOLEAN 相同
# UUID
"UUID": "str", # UUID 一般作为字符串
# 字符串类型
"CHAR": "str",
"VARCHAR": "str",
"TINYTEXT": "str",
"TEXT": "str",
"MEDIUMTEXT": "str",
"LONGTEXT": "str",
"BINARY": "bytes",
"VARBINARY": "bytes",
"TINYBLOB": "bytes",
"BLOB": "bytes",
"MEDIUMBLOB": "bytes",
"LONGBLOB": "bytes",
# 枚举和集合类型
"ENUM": "str",
"SET": "list",
# JSON 类型
"JSON": "dict",
# 空间数据类型(通常需要特殊处理)
"GEOMETRY": "bytes", # 空间数据类型,通常存储为字节流
"POINT": "bytes", # 点数据类型
"LINESTRING": "bytes", # 线数据类型
"POLYGON": "bytes", # 多边形数据类型
"MULTIPOINT": "bytes", # 多点数据类型
"MULTILINESTRING": "bytes", # 多线数据类型
"MULTIPOLYGON": "bytes", # 多多边形数据类型
"GEOMETRYCOLLECTION": "bytes", # 几何集合类型
}
)
@@ -584,7 +584,7 @@ class GenTableService:
await anyio.Path(gen_path).write_text(render_content, encoding="utf-8")
module_init_path = BASE_DIR.parent.joinpath(
f"backend/app/api/v1/{gen_table_schema.module_name}/__init__.py"
f"backend/app/plugin/{gen_table_schema.module_name}/__init__.py"
)
if not module_init_path.exists():
# 创建module_name目录的__init__.py文件
@@ -152,9 +152,24 @@ class GenUtils:
# 因为现在我们确保传入的arr是GenConstant中定义的列表常量
# 并且target_value在调用前已经被处理过不会是None
# 移除 COLLATE 子句和 UNSIGNED 标记(不区分大小写)
target_str = str(target_value)
# 移除 COLLATE 子句
collate_pattern = re.compile(r'\s+COLLATE\s+', re.IGNORECASE)
if collate_pattern.search(target_str):
target_str = collate_pattern.split(target_str)[0].strip()
# 移除 UNSIGNED 标记
unsigned_pattern = re.compile(r'\s+UNSIGNED', re.IGNORECASE)
if unsigned_pattern.search(target_str):
target_str = unsigned_pattern.sub('', target_str).strip()
# 转换为小写进行比较
target_str = target_str.lower()
# 对于包含括号的类型(如TINYINT(1)),需要特殊处理
# 先获取基本类型名称(不含括号)用于比较
target_str = str(target_value).lower()
target_base_type = target_str.split("(")[0] if "(" in target_str else target_str
for item in arr:
@@ -205,9 +220,25 @@ class GenUtils:
返回:
- str: 数据库类型。
"""
# 特殊处理tinyint(1),保留括号和长度信息以便识别为布尔类型
# 移除 COLLATE 子句(处理带引号和不带引号的情况,不区分大小写)
collate_pattern = re.compile(r'\s+COLLATE\s+', re.IGNORECASE)
if collate_pattern.search(column_type):
column_type = collate_pattern.split(column_type)[0].strip()
# 移除 UNSIGNED 标记(不区分大小写)
unsigned_pattern = re.compile(r'\s+UNSIGNED', re.IGNORECASE)
if unsigned_pattern.search(column_type):
column_type = unsigned_pattern.sub('', column_type).strip()
# 特殊处理tinyint(1),映射为boolean
if column_type.lower().startswith("tinyint(1)"):
return column_type
return "boolean"
# 处理PostgreSQL数组类型(如 integer[], text[]
if "[]" in column_type:
return "array"
# 提取基本类型
if "(" in column_type:
return column_type.split("(")[0]
return column_type
@@ -1,3 +1,4 @@
import re
from datetime import datetime
from typing import Any
@@ -207,32 +208,58 @@ class Jinja2TemplateUtil:
"""
columns = gen_table.columns or []
import_list = set()
has_datetime_type = False
has_datetime_import = False
has_date_import = False
has_time_import = False
has_datetime_str = False
has_date_str = False
has_time_str = False
for column in columns:
# 处理嵌套的datetime类型,如datetime.date、datetime.time、datetime.datetime
if (
column.python_type.startswith("datetime.")
or column.python_type in GenConstant.TYPE_DATE
):
has_datetime_type = True
# 处理datetime类型的导入
if column.python_type and column.python_type in GenConstant.TYPE_DATE:
if column.python_type == "datetime":
has_datetime_import = True
elif column.python_type == "date":
has_date_import = True
elif column.python_type == "time":
has_time_import = True
elif column.python_type == GenConstant.TYPE_DECIMAL:
import_list.add("from decimal import Decimal")
# 检查是否需要DateTimeStr、DateStr、TimeStr
if column.column_name == "created_time" or column.column_name == "updated_time":
has_datetime_str = True
if gen_table.sub and gen_table.sub_table and gen_table.sub_table.columns:
sub_columns = gen_table.sub_table.columns or []
for sub_column in sub_columns:
# 处理嵌套的datetime类型,如datetime.date、datetime.time、datetime.datetime
if (
sub_column.python_type.startswith("datetime.")
or sub_column.python_type in GenConstant.TYPE_DATE
):
has_datetime_type = True
# 处理datetime类型的导入
if sub_column.python_type and sub_column.python_type in GenConstant.TYPE_DATE:
if sub_column.python_type == "datetime":
has_datetime_import = True
elif sub_column.python_type == "date":
has_date_import = True
elif sub_column.python_type == "time":
has_time_import = True
elif sub_column.python_type == GenConstant.TYPE_DECIMAL:
import_list.add("from decimal import Decimal")
if has_datetime_type:
import_list.add("import datetime")
# 添加datetime导入
if has_datetime_import:
import_list.add("from datetime import datetime")
if has_date_import:
import_list.add("from datetime import date")
if has_time_import:
import_list.add("from datetime import time")
# 添加validator导入
if has_datetime_str:
import_list.add("from app.core.validator import DateTimeStr")
if has_date_str:
import_list.add("from app.core.validator import DateStr")
if has_time_str:
import_list.add("from app.core.validator import TimeStr")
return import_list
@@ -246,6 +273,9 @@ class Jinja2TemplateUtil:
"""
columns = gen_table.columns or []
import_list = set()
has_datetime_import = False
has_date_import = False
has_time_import = False
for column in columns:
if column.column_type:
@@ -256,10 +286,16 @@ class Jinja2TemplateUtil:
f"from sqlalchemy import {StringUtil.get_mapping_value_by_key_ignore_case(GenConstant.DB_TO_SQLALCHEMY, data_type)}"
)
# 处理datetime类型的导入
if column.python_type and "." in column.python_type:
datetime_type = column.python_type.split(".")[0]
if datetime_type == "datetime":
import_list.add("import datetime")
if column.python_type and column.python_type in GenConstant.TYPE_DATE:
if column.python_type == "datetime":
has_datetime_import = True
elif column.python_type == "date":
has_date_import = True
elif column.python_type == "time":
has_time_import = True
# 处理Decimal类型的导入
elif column.python_type == GenConstant.TYPE_DECIMAL:
import_list.add("from decimal import Decimal")
if gen_table.sub:
import_list.add("from sqlalchemy import ForeignKey")
if gen_table.sub_table and gen_table.sub_table.columns:
@@ -271,10 +307,25 @@ class Jinja2TemplateUtil:
f"from sqlalchemy import {StringUtil.get_mapping_value_by_key_ignore_case(GenConstant.DB_TO_SQLALCHEMY, data_type)}"
)
# 处理datetime类型的导入
if sub_column.python_type and "." in sub_column.python_type:
datetime_type = sub_column.python_type.split(".")[0]
if datetime_type == "datetime":
import_list.add("import datetime")
if sub_column.python_type and sub_column.python_type in GenConstant.TYPE_DATE:
if sub_column.python_type == "datetime":
has_datetime_import = True
elif sub_column.python_type == "date":
has_date_import = True
elif sub_column.python_type == "time":
has_time_import = True
# 处理Decimal类型的导入
elif sub_column.python_type == GenConstant.TYPE_DECIMAL:
import_list.add("from decimal import Decimal")
# 添加datetime导入
if has_datetime_import:
import_list.add("from datetime import datetime")
if has_date_import:
import_list.add("from datetime import date")
if has_time_import:
import_list.add("from datetime import time")
return cls.merge_same_imports(list(import_list), "from sqlalchemy import")
@classmethod
@@ -288,6 +339,21 @@ class Jinja2TemplateUtil:
返回:
- str: 数据库类型(去除长度等修饰)。
"""
# 移除 COLLATE 子句(处理带引号和不带引号的情况,不区分大小写)
collate_pattern = re.compile(r'\s+COLLATE\s+', re.IGNORECASE)
if collate_pattern.search(column_type):
column_type = collate_pattern.split(column_type)[0].strip()
# 移除 UNSIGNED 标记(不区分大小写)
unsigned_pattern = re.compile(r'\s+UNSIGNED', re.IGNORECASE)
if unsigned_pattern.search(column_type):
column_type = unsigned_pattern.sub('', column_type).strip()
# 处理PostgreSQL数组类型(如 integer[], text[]
if "[]" in column_type:
return "array"
# 提取基本类型
if "(" in column_type:
return column_type.split("(")[0]
return column_type
@@ -426,6 +492,9 @@ class Jinja2TemplateUtil:
# 如果是字符串类型且包含括号参数,保持原参数
if sqlalchemy_type in ["String", "CHAR"]:
sqlalchemy_type += "(" + column_type_list[1]
# 如果是Numeric类型且包含括号参数,保持原参数
elif sqlalchemy_type == "Numeric":
sqlalchemy_type += "(" + column_type_list[1]
elif sqlalchemy_type is None:
# 处理没有括号的类型
col_type = column_type