You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 
 
 

655 lines
21 KiB

#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
Python API脚手架生成器
基于FastAPI + SQLAlchemy + Pydantic的三层架构代码生成
复用Java脚手架的config.yml配置文件
@author: Auto Generated
@date: 2024
"""
import os
import sys
import yaml
from urllib.parse import quote_plus
import argparse
from datetime import datetime
from typing import List, Dict, Any, Optional
from jinja2 import Environment, FileSystemLoader
# 添加当前目录到路径
sys.path.insert(0, os.path.dirname(__file__))
# 字段类型映射 (MySQL -> Python/SQLAlchemy)
TYPE_MAPPING = {
"int": {"python": "int", "sqlalchemy": "Integer"},
"bigint": {"python": "int", "sqlalchemy": "BigInteger"},
"smallint": {"python": "int", "sqlalchemy": "SmallInteger"},
"tinyint": {"python": "int", "sqlalchemy": "SmallInteger"},
"float": {"python": "float", "sqlalchemy": "Float"},
"double": {"python": "float", "sqlalchemy": "Float"},
"decimal": {"python": "float", "sqlalchemy": "Numeric"},
"varchar": {"python": "str", "sqlalchemy": "String"},
"char": {"python": "str", "sqlalchemy": "String"},
"text": {"python": "str", "sqlalchemy": "Text"},
"longtext": {"python": "str", "sqlalchemy": "Text"},
"date": {"python": "datetime", "sqlalchemy": "Date"},
"datetime": {"python": "datetime", "sqlalchemy": "DateTime"},
"timestamp": {"python": "datetime", "sqlalchemy": "DateTime"},
"boolean": {"python": "bool", "sqlalchemy": "Boolean"},
"bool": {"python": "bool", "sqlalchemy": "Boolean"},
"json": {"python": "dict", "sqlalchemy": "JSON"},
}
# 需要排除的系统字段
EXCLUDE_FIELDS = [
"id",
"created_at",
"created_by",
"updated_at",
"updated_by",
"deleted_flag",
"is_deleted",
]
class DatabaseInspector:
"""数据库结构检查器"""
def __init__(self, host: str, port: int, user: str, password: str, database: str):
self.host = host
self.port = port
self.user = user
self.password = password
self.database = database
self.connection = None
def connect(self):
"""连接数据库"""
try:
import pymysql
self.connection = pymysql.connect(
host=self.host,
port=self.port,
user=self.user,
password=self.password,
database=self.database,
charset="utf8mb4",
)
print(f"成功连接到数据库: {self.host}:{self.port}/{self.database}")
return True
except ImportError:
print("错误: 请安装 pymysql: pip install pymysql")
return False
except Exception as e:
print(f"数据库连接失败: {e}")
return False
def get_tables(self, table_names: List[str] = None) -> List[Dict[str, Any]]:
"""获取表信息"""
if not self.connection:
return []
tables = []
cursor = self.connection.cursor()
try:
# 获取所有表
if table_names:
placeholders = ",".join(["%s"] * len(table_names))
cursor.execute(
f"""
SELECT TABLE_NAME, TABLE_COMMENT
FROM information_schema.TABLES
WHERE TABLE_SCHEMA = %s AND TABLE_NAME IN ({placeholders})
""",
[self.database] + table_names,
)
else:
cursor.execute(
"""
SELECT TABLE_NAME, TABLE_COMMENT
FROM information_schema.TABLES
WHERE TABLE_SCHEMA = %s
""",
[self.database],
)
for table_name, table_comment in cursor.fetchall():
fields = self._get_table_fields(table_name)
tables.append(
{
"name": table_name,
"comment": table_comment or table_name,
"fields": fields,
}
)
except Exception as e:
print(f"获取表信息失败: {e}")
finally:
cursor.close()
return tables
def _get_table_fields(self, table_name: str) -> List[Dict[str, Any]]:
"""获取表字段信息"""
cursor = self.connection.cursor()
fields = []
try:
cursor.execute(
"""
SELECT COLUMN_NAME, COLUMN_TYPE, COLUMN_COMMENT, IS_NULLABLE, COLUMN_KEY
FROM information_schema.COLUMNS
WHERE TABLE_SCHEMA = %s AND TABLE_NAME = %s
ORDER BY ORDINAL_POSITION
""",
[self.database, table_name],
)
for (
col_name,
col_type,
col_comment,
is_nullable,
col_key,
) in cursor.fetchall():
# 排除系统字段
if col_name.lower() in EXCLUDE_FIELDS:
continue
# 解析类型
base_type = col_type.split("(")[0].lower()
type_info = TYPE_MAPPING.get(
base_type, {"python": "str", "sqlalchemy": "String"}
)
fields.append(
{
"name": col_name,
"type": col_type,
"comment": col_comment or col_name,
"nullable": is_nullable == "YES",
"is_primary": col_key == "PRI",
"searchable": col_key != "PRI", # 非主键默认可搜索
"python_type": type_info["python"],
"sqlalchemy_type": self._get_sqlalchemy_type(col_type),
}
)
except Exception as e:
print(f"获取字段信息失败 {table_name}: {e}")
finally:
cursor.close()
return fields
@staticmethod
def _get_sqlalchemy_type(db_type: str) -> str:
"""获取SQLAlchemy类型字符串"""
base_type = db_type.lower().split("(")[0]
if "(" in db_type and base_type in ["varchar", "char"]:
length = db_type.split("(")[1].rstrip(")")
return f"String({length})"
elif base_type in ["decimal", "numeric"]:
return "Numeric(precision=10, scale=2)"
else:
type_info = TYPE_MAPPING.get(
base_type, {"python": "str", "sqlalchemy": "String"}
)
return type_info["sqlalchemy"]
def close(self):
"""关闭连接"""
if self.connection:
self.connection.close()
class PythonScaffoldingGenerator:
"""Python API脚手架生成器"""
def __init__(self, template_dir: str = None):
"""初始化生成器"""
if template_dir is None:
template_dir = os.path.join(os.path.dirname(__file__), "templates", "py")
self.template_dir = template_dir
self.env = Environment(
loader=FileSystemLoader(template_dir), trim_blocks=True, lstrip_blocks=True
)
# 添加自定义过滤器
self.env.filters["snake_case"] = self._to_snake_case
self.env.filters["camel_case"] = self._to_camel_case
self.env.filters["pascal_case"] = self._to_pascal_case
@staticmethod
def _to_snake_case(name: str) -> str:
"""转换为snake_case"""
result = []
for i, char in enumerate(name):
if char.isupper() and i > 0:
result.append("_")
result.append(char.lower())
return "".join(result)
@staticmethod
def _to_camel_case(name: str) -> str:
"""转换为camelCase"""
parts = name.split("_")
return parts[0].lower() + "".join(p.capitalize() for p in parts[1:])
@staticmethod
def _to_pascal_case(name: str) -> str:
"""转换为PascalCase"""
return "".join(p.capitalize() for p in name.split("_"))
def _process_table(self, table: Dict[str, Any]) -> Dict[str, Any]:
"""处理表信息"""
name = table.get("name", "")
return {
**table,
"entity": self._to_pascal_case(name),
"name_lower": name.lower(),
"comment": table.get("comment", name),
}
def generate_project(
self,
project_name: str,
module_name: str,
tables: List[Dict[str, Any]],
output_dir: str,
author: str = "Auto Generated",
db_config: Dict[str, Any] = None,
**kwargs,
):
"""生成完整项目"""
print(f"开始生成Python API项目: {project_name}/{module_name}")
# 创建项目目录
project_dir = os.path.join(output_dir, project_name, module_name)
os.makedirs(project_dir, exist_ok=True)
# 创建子目录
dirs = [
"app",
"app/models",
"app/schemas",
"app/services",
"app/routers",
"app/utils",
"tests",
]
for d in dirs:
os.makedirs(os.path.join(project_dir, d), exist_ok=True)
# 处理表信息
processed_tables = [self._process_table(t) for t in tables]
# 构建数据库URL
if db_config:
db_host = db_config.get("host", "localhost")
db_port = db_config.get("port", 3306)
db_user = db_config.get("user", "root")
db_password = db_config.get("password", "")
db_database = db_config.get("database", "")
database_url = f"mysql+pymysql://{db_user}:{quote_plus(db_password)}@{db_host}:{db_port}/{db_database}?charset=utf8mb4"
else:
database_url = "sqlite:///./sql_app.db"
# 通用上下文
context = {
"project_name": project_name,
"module_name": module_name,
"author": author,
"date": datetime.now().strftime("%Y-%m-%d"),
"database_url": database_url,
"database_type": "mysql" if db_config else "sqlite",
"tables": processed_tables,
**kwargs,
}
# 生成主文件
self._render_template(
"main.py.j2", context, os.path.join(project_dir, "main.py")
)
self._render_template(
"config.py.j2", context, os.path.join(project_dir, "app", "config.py")
)
self._render_template(
"database.py.j2", context, os.path.join(project_dir, "app", "database.py")
)
self._render_template(
"requirements.txt.j2",
context,
os.path.join(project_dir, "requirements.txt"),
)
self._render_template(
"result.py.j2",
context,
os.path.join(project_dir, "app", "utils", "result.py"),
)
# 生成__init__.py文件
init_dirs = [
"app",
"app/models",
"app/schemas",
"app/services",
"app/routers",
"app/utils",
"tests",
]
for d in init_dirs:
self._render_template(
"__init__.py.j2",
{"package_name": d.replace("/", ".")},
os.path.join(project_dir, d, "__init__.py"),
)
# 为每个表生成代码
for table in processed_tables:
table_context = {
**context,
"table": table,
"fields": table.get("fields", []),
}
# 生成Model
self._render_template(
"model.py.j2",
table_context,
os.path.join(project_dir, "app", "models", f"{table['name_lower']}.py"),
)
# 生成Schema
self._render_template(
"schema.py.j2",
table_context,
os.path.join(
project_dir, "app", "schemas", f"{table['name_lower']}.py"
),
)
# 生成Service
self._render_template(
"service.py.j2",
table_context,
os.path.join(
project_dir, "app", "services", f"{table['name_lower']}_service.py"
),
)
# 生成Router
self._render_template(
"router.py.j2",
table_context,
os.path.join(
project_dir, "app", "routers", f"{table['name_lower']}.py"
),
)
print(f"项目生成完成: {project_dir}")
print(f"\n项目结构:")
self._print_tree(project_dir)
return project_dir
def _render_template(
self, template_name: str, context: Dict[str, Any], output_path: str
):
"""渲染模板并写入文件"""
try:
template = self.env.get_template(template_name)
content = template.render(**context)
with open(output_path, "w", encoding="utf-8") as f:
f.write(content)
print(
f" 生成: {os.path.relpath(output_path, os.path.dirname(output_path))}"
)
except Exception as e:
print(f" 生成失败 {template_name}: {str(e)}")
@staticmethod
def _print_tree(directory: str, prefix: str = ""):
"""打印目录树"""
entries = sorted(os.listdir(directory))
entries = [e for e in entries if not e.startswith("__pycache__")]
for i, entry in enumerate(entries):
path = os.path.join(directory, entry)
is_last = i == len(entries) - 1
connector = "└── " if is_last else "├── "
print(f"{prefix}{connector}{entry}")
if os.path.isdir(path):
extension = " " if is_last else ""
PythonScaffoldingGenerator._print_tree(path, prefix + extension)
def load_config(config_path: str) -> Dict[str, Any]:
"""加载YAML配置文件"""
with open(config_path, "r", encoding="utf-8") as f:
config = yaml.safe_load(f)
return config
def resolve_variables(config: Dict[str, Any]) -> Dict[str, Any]:
"""解析配置中的变量引用"""
config_str = yaml.dump(config)
# 简单的变量替换
replacements = {
"${mainModule}": config.get("mainModule", ""),
"${moduleName}": config.get("moduleName", ""),
"${package.Base}": config.get("package", {}).get("Base", ""),
"${package.Models}": config.get("package", {}).get("Models", ""),
}
for key, value in replacements.items():
config_str = config_str.replace(key, value)
return yaml.safe_load(config_str)
def main():
"""主函数"""
parser = argparse.ArgumentParser(
description="Python API脚手架生成器 - 复用Java配置"
)
parser.add_argument(
"-c", "--config", default="config.yml", help="配置文件路径(默认: config.yml)"
)
parser.add_argument(
"-o", "--output", default="./output", help="输出目录(默认: ./output)"
)
parser.add_argument(
"-t", "--tables", nargs="+", help="指定要生成的表名(默认: 全部)"
)
parser.add_argument(
"--no-db", action="store_true", help="不连接数据库,使用示例数据"
)
args = parser.parse_args()
# 检查配置文件
if not os.path.exists(args.config):
print(f"错误: 配置文件不存在: {args.config}")
return
# 加载配置
print(f"加载配置文件: {args.config}")
config = load_config(args.config)
config = resolve_variables(config)
# 获取基础信息
main_module = config.get("mainModule", "project")
module_name = config.get("moduleName", "api")
author = config.get("author", "Auto Generated")
db_config = config.get("db", {})
app_config = config.get("application", {})
print(f"项目: {main_module}/{module_name}")
print(f"作者: {author}")
# 获取表结构
tables = []
if not args.no_db and db_config:
print(f"\n连接数据库: {db_config.get('host')}:{db_config.get('port')}")
inspector = DatabaseInspector(
host=db_config.get("host", "localhost"),
port=db_config.get("port", 3306),
user=db_config.get("user", "root"),
password=db_config.get("password", ""),
database=db_config.get("database", ""),
)
if inspector.connect():
tables = inspector.get_tables(args.tables)
inspector.close()
print(f"获取到 {len(tables)} 个表")
else:
print("数据库连接失败,使用示例数据")
tables = _get_sample_tables()
else:
print("\n使用示例数据")
tables = _get_sample_tables()
if not tables:
print("没有要生成的表")
return
# 初始化生成器
generator = PythonScaffoldingGenerator()
# 生成项目
generator.generate_project(
project_name=main_module,
module_name=module_name,
tables=tables,
output_dir=args.output,
author=author,
db_config=db_config if not args.no_db else None,
redis_config=app_config.get("redis"),
minio_config=app_config.get("minio"),
)
print("\n生成完成!")
print(f"\n下一步:")
print(f" 1. cd {args.output}/{main_module}/{module_name}")
print(f" 2. pip install -r requirements.txt")
print(f" 3. python main.py")
def _get_sample_tables() -> List[Dict[str, Any]]:
"""获取示例表数据"""
return [
{
"name": "user",
"comment": "用户表",
"fields": [
{
"name": "username",
"type": "varchar(50)",
"comment": "用户名",
"nullable": False,
"searchable": True,
"python_type": "str",
"sqlalchemy_type": "String(50)",
},
{
"name": "password",
"type": "varchar(100)",
"comment": "密码",
"nullable": False,
"searchable": False,
"python_type": "str",
"sqlalchemy_type": "String(100)",
},
{
"name": "email",
"type": "varchar(100)",
"comment": "邮箱",
"nullable": True,
"searchable": True,
"python_type": "str",
"sqlalchemy_type": "String(100)",
},
{
"name": "phone",
"type": "varchar(20)",
"comment": "手机号",
"nullable": True,
"searchable": True,
"python_type": "str",
"sqlalchemy_type": "String(20)",
},
{
"name": "status",
"type": "int",
"comment": "状态",
"nullable": True,
"searchable": True,
"python_type": "int",
"sqlalchemy_type": "Integer",
},
],
},
{
"name": "product",
"comment": "商品表",
"fields": [
{
"name": "name",
"type": "varchar(100)",
"comment": "商品名称",
"nullable": False,
"searchable": True,
"python_type": "str",
"sqlalchemy_type": "String(100)",
},
{
"name": "description",
"type": "text",
"comment": "商品描述",
"nullable": True,
"searchable": False,
"python_type": "str",
"sqlalchemy_type": "Text",
},
{
"name": "price",
"type": "decimal(10,2)",
"comment": "价格",
"nullable": False,
"searchable": True,
"python_type": "float",
"sqlalchemy_type": "Numeric(precision=10, scale=2)",
},
{
"name": "stock",
"type": "int",
"comment": "库存",
"nullable": True,
"searchable": True,
"python_type": "int",
"sqlalchemy_type": "Integer",
},
{
"name": "category_id",
"type": "int",
"comment": "分类ID",
"nullable": True,
"searchable": True,
"python_type": "int",
"sqlalchemy_type": "Integer",
},
],
},
]
if __name__ == "__main__":
main()