diff --git a/backend/cli.py b/backend/cli.py index 071c7f12..53ce0680 100644 --- a/backend/cli.py +++ b/backend/cli.py @@ -25,7 +25,12 @@ from backend import __version__ from backend.common.enums import DataBaseType, PrimaryKeyType from backend.common.exception.errors import BaseExceptionError from backend.core.conf import settings -from backend.core.path_conf import BASE_PATH +from backend.core.path_conf import ( + ENV_EXAMPLE_FILE_PATH, + ENV_FILE_PATH, + MYSQL_SCRIPT_DIR, + POSTGRESQL_SCRIPT_DIR, +) from backend.database.db import async_db_session, create_tables, drop_tables from backend.database.redis import redis_client from backend.plugin.tools import get_plugin_sql, get_plugins @@ -44,15 +49,12 @@ class CustomReloadFilter(PythonFilter): def setup_env_file() -> bool: - env_path = BASE_PATH / '.env' - env_example_path = BASE_PATH / '.env.example' - - if not env_example_path.exists(): + if not ENV_EXAMPLE_FILE_PATH.exists(): console.print('.env.example 文件不存在', style='red') return False try: - env_content = Path(env_example_path).read_text(encoding='utf-8') + env_content = Path(ENV_EXAMPLE_FILE_PATH).read_text(encoding='utf-8') console.print('配置数据库连接信息...', style='white') db_type = Prompt.ask('数据库类型', choices=['mysql', 'postgresql'], default='postgresql') db_host = Prompt.ask('数据库主机', default='127.0.0.1') @@ -96,7 +98,7 @@ def setup_env_file() -> bool: ) settings.OPERA_LOG_ENCRYPT_SECRET_KEY = opera_log_secret - Path(env_path).write_text(env_content, encoding='utf-8') + Path(ENV_FILE_PATH).write_text(env_content, encoding='utf-8') console.print('.env 文件创建成功', style='green') except Exception as e: console.print(f'.env 文件创建失败: {e}', style='red') @@ -173,7 +175,7 @@ async def auto_init() -> None: panel_content.append('\n • Redis 连接信息') panel_content.append('\n • Token 密钥(自动生成)') - console.print(Panel(panel_content, title=f'fba v{__version__} 环境变量', border_style='cyan', padding=(1, 2))) + console.print(Panel(panel_content, title=f'fba (v{__version__}) - 环境变量', border_style='cyan', padding=(1, 2))) if not setup_env_file(): raise cappa.Exit('.env 文件配置失败', code=1) @@ -187,7 +189,7 @@ async def auto_init() -> None: panel_content.append('\n • 主机:') panel_content.append(f'{settings.DATABASE_HOST}:{settings.DATABASE_PORT}', style='yellow') - console.print(Panel(panel_content, title=f'fba v{__version__} 数据库', border_style='cyan', padding=(1, 2))) + console.print(Panel(panel_content, title=f'fba (v{__version__}) - 数据库', border_style='cyan', padding=(1, 2))) ok = Prompt.ask('即将[red]新建/重建数据库[/red],确认继续吗?', choices=['y', 'n'], default='n') if ok.lower() == 'y': @@ -229,7 +231,7 @@ async def init() -> None: else: panel_content.append('无', style='dim') - console.print(Panel(panel_content, title=f'fba v{__version__} 初始化', border_style='cyan', padding=(1, 2))) + console.print(Panel(panel_content, title=f'fba (v{__version__}) - 初始化', border_style='cyan', padding=(1, 2))) ok = Prompt.ask( '即将[red]新建/重建数据库表[/red]并[red]执行所有数据库脚本[/red],确认继续吗?', choices=['y', 'n'], default='n' ) @@ -290,7 +292,7 @@ def run(host: str, port: int, reload: bool, workers: int) -> None: # noqa: FBT0 panel_content.append('\n🌐 架构官方文档: ', style='bold magenta') panel_content.append('https://fastapi-practices.github.io/fastapi_best_architecture_docs/') - console.print(Panel(panel_content, title=f'fba v{__version__}', border_style='purple', padding=(1, 2))) + console.print(Panel(panel_content, title=f'fba (v{__version__})', border_style='purple', padding=(1, 2))) granian.Granian( target='backend.main:app', interface='asgi', @@ -364,15 +366,11 @@ async def install_plugin( async def get_sql_scripts() -> list[str]: sql_scripts = [] - db_dir = ( - BASE_PATH / 'sql' / 'mysql' - if DataBaseType.mysql == settings.DATABASE_TYPE - else BASE_PATH / 'sql' / 'postgresql' - ) + db_script_dir = MYSQL_SCRIPT_DIR if DataBaseType.mysql == settings.DATABASE_TYPE else POSTGRESQL_SCRIPT_DIR main_sql_file = ( - db_dir / 'init_test_data.sql' + db_script_dir / 'init_test_data.sql' if PrimaryKeyType.autoincrement == settings.DATABASE_PK_MODE - else db_dir / 'init_snowflake_test_data.sql' + else db_script_dir / 'init_snowflake_test_data.sql' ) main_sql_path = anyio.Path(main_sql_file) diff --git a/backend/core/conf.py b/backend/core/conf.py index 55942709..889e5a3e 100644 --- a/backend/core/conf.py +++ b/backend/core/conf.py @@ -1,3 +1,5 @@ +import shutil + from functools import lru_cache from re import Pattern from typing import Any, Literal @@ -5,14 +7,14 @@ from typing import Any, Literal from pydantic import model_validator from pydantic_settings import BaseSettings, SettingsConfigDict -from backend.core.path_conf import BASE_PATH +from backend.core.path_conf import ENV_EXAMPLE_FILE_PATH, ENV_FILE_PATH class Settings(BaseSettings): """全局配置""" model_config = SettingsConfigDict( - env_file=f'{BASE_PATH}/.env', + env_file=ENV_FILE_PATH, env_file_encoding='utf-8', extra='ignore', case_sensitive=True, @@ -296,6 +298,8 @@ class Settings(BaseSettings): @lru_cache def get_settings() -> Settings: """获取全局配置单例""" + if not ENV_FILE_PATH.exists(): + shutil.copy(ENV_EXAMPLE_FILE_PATH, ENV_FILE_PATH) return Settings() diff --git a/backend/core/path_conf.py b/backend/core/path_conf.py index 5d659ddf..2c6fd18e 100644 --- a/backend/core/path_conf.py +++ b/backend/core/path_conf.py @@ -3,6 +3,12 @@ from pathlib import Path # 项目根目录 BASE_PATH = Path(__file__).resolve().parent.parent +# 环境变量文件 +ENV_FILE_PATH = BASE_PATH / '.env' + +# 环境变量示例文件 +ENV_EXAMPLE_FILE_PATH = BASE_PATH / '.env.example' + # alembic 迁移文件存放路径 ALEMBIC_VERSION_DIR = BASE_PATH / 'alembic' / 'versions' @@ -20,3 +26,9 @@ PLUGIN_DIR = BASE_PATH / 'plugin' # 国际化文件目录 LOCALE_DIR = BASE_PATH / 'locale' + +# MySQL 脚本目录 +MYSQL_SCRIPT_DIR = BASE_PATH / 'sql' / 'mysql' + +# PostgreSQL 脚本目录 +POSTGRESQL_SCRIPT_DIR = BASE_PATH / 'sql' / 'postgresql'