diff --git a/backend/cli.py b/backend/cli.py index b4d578eb..8680f4b4 100644 --- a/backend/cli.py +++ b/backend/cli.py @@ -9,7 +9,10 @@ from typing import Annotated, Literal import cappa import granian +from cappa import ValueFrom from rich.panel import Panel +from rich.prompt import IntPrompt +from rich.table import Table from rich.text import Text from sqlalchemy import text from watchfiles import PythonFilter @@ -19,7 +22,11 @@ from backend.common.enums import DataBaseType, PrimaryKeyType from backend.common.exception.errors import BaseExceptionMixin from backend.core.conf import settings from backend.database.db import async_db_session +from backend.plugin.code_generator.schema.code import ImportParam +from backend.plugin.code_generator.service.business_service import gen_business_service +from backend.plugin.code_generator.service.code_service import gen_service from backend.plugin.tools import get_plugin_sql +from backend.utils._await import run_await from backend.utils.file_ops import install_git_plugin, install_zip_plugin, parse_sql_script @@ -30,7 +37,7 @@ class CustomReloadFilter(PythonFilter): super().__init__(extra_extensions=['.json', '.yaml', '.yml']) -def run(host: str, port: int, reload: bool, workers: int | None) -> None: +def run(host: str, port: int, reload: bool, workers: int) -> None: url = f'http://{host}:{port}' docs_url = url + settings.FASTAPI_DOCS_URL redoc_url = url + settings.FASTAPI_REDOC_URL @@ -53,7 +60,7 @@ def run(host: str, port: int, reload: bool, workers: int | None) -> None: port=port, reload=not reload, reload_filter=CustomReloadFilter, - workers=workers or 1, + workers=workers, ).serve() @@ -125,13 +132,66 @@ async def execute_sql_scripts(sql_scripts: str) -> None: console.print(Text('SQL 脚本已执行完成', style='bold green')) -@cappa.command(help='运行 API 服务') +async def import_table( + app: str, + table_schema: str, + table_name: str, +) -> None: + try: + obj = ImportParam(app=app, table_schema=table_schema, table_name=table_name) + await gen_service.import_business_and_model(obj=obj) + except Exception as e: + raise cappa.Exit(e.msg if isinstance(e, BaseExceptionMixin) else str(e), code=1) + + +def get_business() -> int: + ids = [] + try: + results = run_await(gen_business_service.get_all)() + + if not results: + raise cappa.Exit('[red]暂无可用的代码生成业务!请先通过 import 命令导入![/]') + + table = Table(show_header=True, header_style='bold magenta') + table.add_column('业务编号', style='cyan', no_wrap=True, justify='center') + table.add_column('应用名称', style='green', no_wrap=True) + table.add_column('生成路径', style='yellow') + table.add_column('备注', style='blue') + + for result in results: + table.add_row( + str(result.id), + result.app_name, + result.gen_path or f'应用 {result.app_name} 根路径', + result.remark or '', + ) + ids.append(result.id) + + console.print(table) + except Exception as e: + raise cappa.Exit(e.msg if isinstance(e, BaseExceptionMixin) else str(e), code=1) + + business = IntPrompt.ask('请从中选择一个业务编号', choices=[str(_id) for _id in ids]) + + return business + + +def generate(pk: int) -> None: + try: + gen_path = run_await(gen_service.generate)(pk=pk) + except Exception as e: + raise cappa.Exit(e.msg if isinstance(e, BaseExceptionMixin) else str(e), code=1) + + console.print(Text('\n代码已生成完毕', style='bold green')) + console.print(Text('\n详情请查看:'), Text(gen_path, style='bold magenta')) + + +@cappa.command(help='运行 API 服务', default_long=True) @dataclass class Run: host: Annotated[ str, cappa.Arg( - long=True, default='127.0.0.1', help='提供服务的主机 IP 地址,对于本地开发,请使用 `127.0.0.1`。' '要启用公共访问,例如在局域网中,请使用 `0.0.0.0`', @@ -139,50 +199,56 @@ class Run: ] port: Annotated[ int, - cappa.Arg(long=True, default=8000, help='提供服务的主机端口号'), + cappa.Arg(default=8000, help='提供服务的主机端口号'), ] no_reload: Annotated[ bool, - cappa.Arg(long=True, default=False, help='禁用在(代码)文件更改时自动重新加载服务器'), + cappa.Arg(default=False, help='禁用在(代码)文件更改时自动重新加载服务器'), ] workers: Annotated[ - int | None, - cappa.Arg(long=True, default=None, help='使用多个工作进程,必须与 `--no-reload` 同时使用'), + int, + cappa.Arg(default=1, help='使用多个工作进程,必须与 `--no-reload` 同时使用'), ] def __call__(self): run(host=self.host, port=self.port, reload=self.no_reload, workers=self.workers) -@cappa.command(help='从当前主机启动 Celery worker 服务') +@cappa.command(help='从当前主机启动 Celery worker 服务', default_long=True) @dataclass class Worker: log_level: Annotated[ Literal['info', 'debug'], - cappa.Arg(long=True, short='-l', default='info', help='日志输出级别'), + cappa.Arg(short='-l', default='info', help='日志输出级别'), ] def __call__(self): run_celery_worker(log_level=self.log_level) -@cappa.command(help='从当前主机启动 Celery beat 服务') +@cappa.command(help='从当前主机启动 Celery beat 服务', default_long=True) @dataclass class Beat: log_level: Annotated[ Literal['info', 'debug'], - cappa.Arg(long=True, short='-l', default='info', help='日志输出级别'), + cappa.Arg(short='-l', default='info', help='日志输出级别'), ] def __call__(self): run_celery_beat(log_level=self.log_level) -@cappa.command(help='从当前主机启动 Celery flower 服务') +@cappa.command(help='从当前主机启动 Celery flower 服务', default_long=True) @dataclass class Flower: - port: Annotated[int, cappa.Arg(long=True, default=8555, help='提供服务的主机端口号')] - basic_auth: Annotated[str, cappa.Arg(long=True, default='admin:123456', help='页面登录的用户名和密码')] + port: Annotated[ + int, + cappa.Arg(default=8555, help='提供服务的主机端口号'), + ] + basic_auth: Annotated[ + str, + cappa.Arg(default='admin:123456', help='页面登录的用户名和密码'), + ] def __call__(self): run_celery_flower(port=self.port, basic_auth=self.basic_auth) @@ -194,46 +260,80 @@ class Celery: subcmd: cappa.Subcommands[Worker | Beat | Flower] -@cappa.command(help='新增插件') +@cappa.command(help='新增插件', default_long=True) @dataclass class Add: path: Annotated[ str | None, - cappa.Arg(long=True, help='ZIP 插件的本地完整路径'), + cappa.Arg(help='ZIP 插件的本地完整路径'), ] repo_url: Annotated[ str | None, - cappa.Arg(long=True, help='Git 插件的仓库地址'), + cappa.Arg(help='Git 插件的仓库地址'), ] no_sql: Annotated[ bool, - cappa.Arg(long=True, default=False, help='禁用插件 SQL 脚本自动执行'), + cappa.Arg(default=False, help='禁用插件 SQL 脚本自动执行'), ] db_type: Annotated[ DataBaseType, - cappa.Arg(long=True, default='mysql', help='执行插件 SQL 脚本的数据库类型'), + cappa.Arg(default='mysql', help='执行插件 SQL 脚本的数据库类型'), ] pk_type: Annotated[ PrimaryKeyType, - cappa.Arg(long=True, default='autoincrement', help='执行插件 SQL 脚本数据库主键类型'), + cappa.Arg(default='autoincrement', help='执行插件 SQL 脚本数据库主键类型'), ] async def __call__(self): await install_plugin(self.path, self.repo_url, self.no_sql, self.db_type, self.pk_type) -@cappa.command(help='一个高效的 fba 命令行界面') +@cappa.command(help='导入代码生成业务和模型列', default_long=True) +@dataclass +class Import: + app: Annotated[ + str, + cappa.Arg(help='应用名称,用于代码生成到指定 app'), + ] + table_schema: Annotated[ + str, + cappa.Arg(short='tc', default='fba', help='数据库名'), + ] + table_name: Annotated[ + str, + cappa.Arg(short='tn', help='数据库表名'), + ] + + async def __call__(self): + await import_table(self.app, self.table_schema, self.table_name) + + +@cappa.command(name='gen', help='代码生成,体验完成功能,请自行部署 fba vben 前端工程!', default_long=True) +@dataclass +class CodeGenerate: + business: Annotated[ + int, + cappa.Arg(default=ValueFrom(get_business), help='业务编号'), + ] + subcmd: cappa.Subcommands[Import | None] = None + + def __call__(self): + if self.business: + generate(self.business) + + +@cappa.command(help='一个高效的 fba 命令行界面', default_long=True) @dataclass class FbaCli: version: Annotated[ bool, - cappa.Arg(short='-V', long=True, default=False, show_default=False, help='打印当前版本号'), + cappa.Arg(short='-V', default=False, show_default=False, help='打印当前版本号'), ] sql: Annotated[ str, - cappa.Arg(value_name='PATH', long=True, default='', show_default=False, help='在事务中执行 SQL 脚本'), + cappa.Arg(value_name='PATH', default='', show_default=False, help='在事务中执行 SQL 脚本'), ] - subcmd: cappa.Subcommands[Run | Celery | Add | None] = None + subcmd: cappa.Subcommands[Run | Celery | Add | CodeGenerate | None] = None async def __call__(self): if self.version: diff --git a/backend/plugin/code_generator/service/code_service.py b/backend/plugin/code_generator/service/code_service.py index 2e5eb13a..0d6fa5d2 100644 --- a/backend/plugin/code_generator/service/code_service.py +++ b/backend/plugin/code_generator/service/code_service.py @@ -162,7 +162,7 @@ class GenService: return [os.path.join(gen_path, *target_file.split('/')) for target_file in target_files] - async def generate(self, *, pk: int) -> None: + async def generate(self, *, pk: int) -> str: """ 生成代码文件 @@ -217,6 +217,8 @@ class GenService: async with aiofiles.open(code_filepath, 'w', encoding='utf-8') as f: await f.write(code) + return gen_path + async def download(self, *, pk: int) -> io.BytesIO: """ 下载生成的代码 diff --git a/backend/utils/_await.py b/backend/utils/_await.py index e1950fde..1731cc8d 100644 --- a/backend/utils/_await.py +++ b/backend/utils/_await.py @@ -20,14 +20,15 @@ class _TaskRunner: def close(self): """关闭事件循环并清理""" - if self.__loop: - self.__loop.stop() + with self.__lock: + if self.__loop: + self.__loop.call_soon_threadsafe(self.__loop.stop) + if self.__thread and self.__thread.is_alive(): + self.__thread.join() self.__loop = None - if self.__thread: - self.__thread.join() self.__thread = None - name = f'TaskRunner-{threading.get_ident()}' - _runner_map.pop(name, None) + name = f'TaskRunner-{threading.get_ident()}' + _runner_map.pop(name, None) def _target(self): """后台线程的目标函数""" @@ -52,13 +53,14 @@ _runner_map = weakref.WeakValueDictionary() def run_await(coro: Callable[..., Awaitable[T]] | Callable[..., Coroutine[Any, Any, T]]) -> Callable[..., T]: - """将协程包装在函数中,该函数将在后台事件循环上运行,直到它执行完为止""" + """将协程包装在函数中,直到它执行完为止""" @wraps(coro) def wrapped(*args, **kwargs): inner = coro(*args, **kwargs) if not asyncio.iscoroutine(inner) and not asyncio.isfuture(inner): - raise TypeError(f'Expected coroutine, got {type(inner)}') + raise TypeError(f'Expected coroutine or future, got {type(inner)}') + try: # 如果事件循环正在运行,则使用任务调用 asyncio.get_running_loop() @@ -68,7 +70,11 @@ def run_await(coro: Callable[..., Awaitable[T]] | Callable[..., Coroutine[Any, A return _runner_map[name].run(inner) except RuntimeError: # 如果没有,则创建一个新的事件循环 - loop = asyncio.get_event_loop() + try: + loop = asyncio.get_event_loop() + except RuntimeError: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) return loop.run_until_complete(inner) wrapped.__doc__ = coro.__doc__