From 5c7da70df86bdd11fea9f98102116ffa63362338 Mon Sep 17 00:00:00 2001 From: insistence <3055204202@qq.com> Date: Thu, 20 Aug 2026 16:29:06 +0800 Subject: [PATCH 1/5] =?UTF-8?q?feat:=20=E6=96=B0=E5=A2=9E=E5=A4=9A?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E6=BA=90=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ruoyi-fastapi-backend/.env.dev | 55 +- ruoyi-fastapi-backend/.env.dockermy | 55 +- ruoyi-fastapi-backend/.env.dockerpg | 55 +- ruoyi-fastapi-backend/.env.prod | 55 +- ruoyi-fastapi-backend/alembic/env.py | 28 +- .../cli/core/context_factory.py | 7 +- .../cli/runtime/app/support.py | 2 +- .../cli/runtime/config/gateway.py | 2 +- .../cli/runtime/gen/gateway.py | 4 +- .../cli/runtime/job/gateway.py | 2 +- .../common/aspect/db_seesion.py | 40 +- .../common/aspect/pre_auth.py | 4 +- ruoyi-fastapi-backend/common/constant.py | 106 ++-- ruoyi-fastapi-backend/config/database.py | 533 ++++++++++++++++-- ruoyi-fastapi-backend/config/env.py | 128 ++++- ruoyi-fastapi-backend/config/get_redis.py | 6 +- ruoyi-fastapi-backend/config/get_scheduler.py | 54 +- .../config/{get_db.py => lifecycle.py} | 30 +- ruoyi-fastapi-backend/exceptions/exception.py | 39 ++ .../module_admin/entity/do/config_do.py | 2 +- .../module_admin/entity/do/dept_do.py | 6 +- .../module_admin/entity/do/dict_do.py | 8 +- .../module_admin/entity/do/menu_do.py | 6 +- .../module_admin/entity/do/notice_do.py | 6 +- .../module_admin/entity/do/post_do.py | 2 +- .../module_admin/entity/do/role_do.py | 6 +- .../module_admin/entity/do/user_do.py | 4 +- .../module_admin/service/file_service.py | 8 +- .../module_admin/service/log_service.py | 4 +- .../module_admin/service/login_service.py | 7 +- .../controller/gen_controller.py | 30 +- .../module_generator/dao/gen_dao.py | 107 +++- .../module_generator/entity/do/gen_do.py | 7 +- .../module_generator/entity/vo/gen_vo.py | 13 + .../module_generator/service/gen_service.py | 168 ++++-- .../templates/python/controller.py.jinja2 | 12 +- .../templates/python/do.py.jinja2 | 15 +- .../module_task/file_task.py | 6 +- .../plugins/ai/entity/do/ai_model_do.py | 10 +- .../plugins/ai/utils/ai_util.py | 10 +- .../plugins/core/lifecycle/migration.py | 2 +- .../plugins/core/lifecycle/seed.py | 2 +- .../core/management/entity/do/models.py | 4 +- .../core/management/service/gateway.py | 2 +- .../plugins/core/runtime/route_guard.py | 4 +- .../plugins/core/runtime/startup.py | 23 +- .../plugins/core/validation/manifest.py | 2 +- .../scripts/migrate_legacy_files.py | 4 +- ruoyi-fastapi-backend/server.py | 77 ++- .../sql/ruoyi-fastapi-pg.sql | 2 + ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql | 1 + .../tests/cli/runtime/test_app_runtime.py | 10 +- .../tests/cli/tui/adapters/conftest.py | 11 +- .../test_database_exception_handlers.py | 37 ++ .../tests/config/test_database_registry.py | 301 ++++++++++ .../tests/config/test_database_settings.py | 124 ++++ .../tests/config/test_lifecycle.py | 25 + .../config/test_multi_datasource_migration.py | 41 ++ .../tests/config/test_scheduler_database.py | 40 ++ .../config/test_scheduler_leader_lease.py | 5 - .../service/test_file_access_log.py | 9 +- .../module_generator/dao/test_gen_dao.py | 46 +- .../module_generator/test_gen_service.py | 57 ++ .../plugins/core/lifecycle/test_migration.py | 2 +- .../tests/plugins/core/lifecycle/test_seed.py | 2 +- .../plugins/core/runtime/test_startup.py | 20 +- .../plugins/core/validation/test_manifest.py | 4 +- .../scripts/test_migrate_legacy_files.py | 13 +- .../tests/server/test_plugin_runtime.py | 68 ++- .../tests/utils/test_template_util.py | 60 ++ ruoyi-fastapi-backend/utils/gen_util.py | 26 +- ruoyi-fastapi-backend/utils/template_util.py | 59 +- ruoyi-fastapi-frontend/src/api/tool/gen.js | 18 +- .../src/views/tool/gen/createTable.vue | 35 +- .../src/views/tool/gen/importTable.vue | 60 +- .../src/views/tool/gen/index.vue | 56 +- 76 files changed, 2318 insertions(+), 576 deletions(-) rename ruoyi-fastapi-backend/config/{get_db.py => lifecycle.py} (51%) create mode 100644 ruoyi-fastapi-backend/tests/config/test_database_exception_handlers.py create mode 100644 ruoyi-fastapi-backend/tests/config/test_database_registry.py create mode 100644 ruoyi-fastapi-backend/tests/config/test_database_settings.py create mode 100644 ruoyi-fastapi-backend/tests/config/test_lifecycle.py create mode 100644 ruoyi-fastapi-backend/tests/config/test_multi_datasource_migration.py create mode 100644 ruoyi-fastapi-backend/tests/config/test_scheduler_database.py create mode 100644 ruoyi-fastapi-backend/tests/module_generator/test_gen_service.py diff --git a/ruoyi-fastapi-backend/.env.dev b/ruoyi-fastapi-backend/.env.dev index ea95718..d68f083 100644 --- a/ruoyi-fastapi-backend/.env.dev +++ b/ruoyi-fastapi-backend/.env.dev @@ -46,28 +46,39 @@ JWT_REDIS_EXPIRE_MINUTES = 30 # -------- 数据库配置 -------- -# 数据库类型,可选的有'mysql'、'postgresql',默认为'mysql' -DB_TYPE = 'mysql' -# 数据库主机 -DB_HOST = '127.0.0.1' -# 数据库端口 -DB_PORT = 3306 -# 数据库用户名 -DB_USERNAME = 'root' -# 数据库密码 -DB_PASSWORD = 'mysqlroot' -# 数据库名称 -DB_DATABASE = 'ruoyi-fastapi' -# 是否开启sqlalchemy日志 -DB_ECHO = true -# 允许溢出连接池大小的最大连接数 -DB_MAX_OVERFLOW = 10 -# 连接池大小,0表示连接数无限制 -DB_POOL_SIZE = 50 -# 连接回收时间(单位:秒) -DB_POOL_RECYCLE = 3600 -# 连接池中没有线程可用时,最多等待的时间(单位:秒) -DB_POOL_TIMEOUT = 30 +# 默认数据源名称 +DB_DEFAULT_SOURCE = 'primary' +# 数据源配置(每个节点完整保留数据库连接字段) +# db_type:数据库类型,可选 mysql、postgresql +# db_host:数据库主机地址 +# db_port:数据库端口 +# db_username:数据库用户名 +# db_password:数据库密码 +# db_database:数据库名称 +# db_echo:是否输出SQLAlchemy SQL日志 +# db_connect_timeout:建立数据库连接的超时时间(秒) +# db_max_overflow:连接池允许的最大溢出连接数 +# db_pool_size:连接池常驻连接数 +# db_pool_recycle:连接回收间隔(秒),-1表示禁用回收 +# db_pool_timeout:获取连接的超时时间(秒) +# db_required:启动时连接失败是否阻止应用启动 +DB_SOURCES = '{ + "primary": { + "db_type": "mysql", + "db_host": "127.0.0.1", + "db_port": 3306, + "db_username": "root", + "db_password": "mysqlroot", + "db_database": "ruoyi-fastapi", + "db_echo": true, + "db_connect_timeout": 10, + "db_max_overflow": 10, + "db_pool_size": 50, + "db_pool_recycle": 3600, + "db_pool_timeout": 30, + "db_required": true + } +}' # -------- Redis配置 -------- # Redis主机 diff --git a/ruoyi-fastapi-backend/.env.dockermy b/ruoyi-fastapi-backend/.env.dockermy index fd2dcfa..895b717 100644 --- a/ruoyi-fastapi-backend/.env.dockermy +++ b/ruoyi-fastapi-backend/.env.dockermy @@ -46,28 +46,39 @@ JWT_REDIS_EXPIRE_MINUTES = 30 # -------- 数据库配置 -------- -# 数据库类型,可选的有'mysql'、'postgresql',默认为'mysql' -DB_TYPE = 'mysql' -# 数据库主机 -DB_HOST = 'ruoyi-mysql' -# 数据库端口 -DB_PORT = 3306 -# 数据库用户名 -DB_USERNAME = 'root' -# 数据库密码 -DB_PASSWORD = 'root' -# 数据库名称 -DB_DATABASE = 'ruoyi-fastapi' -# 是否开启sqlalchemy日志 -DB_ECHO = true -# 允许溢出连接池大小的最大连接数 -DB_MAX_OVERFLOW = 10 -# 连接池大小,0表示连接数无限制 -DB_POOL_SIZE = 50 -# 连接回收时间(单位:秒) -DB_POOL_RECYCLE = 3600 -# 连接池中没有线程可用时,最多等待的时间(单位:秒) -DB_POOL_TIMEOUT = 30 +# 默认数据源名称 +DB_DEFAULT_SOURCE = 'primary' +# 数据源配置(每个节点完整保留数据库连接字段) +# db_type:数据库类型,可选 mysql、postgresql +# db_host:数据库主机地址 +# db_port:数据库端口 +# db_username:数据库用户名 +# db_password:数据库密码 +# db_database:数据库名称 +# db_echo:是否输出SQLAlchemy SQL日志 +# db_connect_timeout:建立数据库连接的超时时间(秒) +# db_max_overflow:连接池允许的最大溢出连接数 +# db_pool_size:连接池常驻连接数 +# db_pool_recycle:连接回收间隔(秒),-1表示禁用回收 +# db_pool_timeout:获取连接的超时时间(秒) +# db_required:启动时连接失败是否阻止应用启动 +DB_SOURCES = '{ + "primary": { + "db_type": "mysql", + "db_host": "ruoyi-mysql", + "db_port": 3306, + "db_username": "root", + "db_password": "root", + "db_database": "ruoyi-fastapi", + "db_echo": true, + "db_connect_timeout": 10, + "db_max_overflow": 10, + "db_pool_size": 50, + "db_pool_recycle": 3600, + "db_pool_timeout": 30, + "db_required": true + } +}' # -------- Redis配置 -------- # Redis主机 diff --git a/ruoyi-fastapi-backend/.env.dockerpg b/ruoyi-fastapi-backend/.env.dockerpg index a75ef1e..cd52d0b 100644 --- a/ruoyi-fastapi-backend/.env.dockerpg +++ b/ruoyi-fastapi-backend/.env.dockerpg @@ -46,28 +46,39 @@ JWT_REDIS_EXPIRE_MINUTES = 30 # -------- 数据库配置 -------- -# 数据库类型,可选的有'mysql'、'postgresql',默认为'mysql' -DB_TYPE = 'postgresql' -# 数据库主机 -DB_HOST = 'ruoyi-pg' -# 数据库端口 -DB_PORT = 5432 -# 数据库用户名 -DB_USERNAME = 'postgres' -# 数据库密码 -DB_PASSWORD = 'root' -# 数据库名称 -DB_DATABASE = 'ruoyi-fastapi' -# 是否开启sqlalchemy日志 -DB_ECHO = true -# 允许溢出连接池大小的最大连接数 -DB_MAX_OVERFLOW = 10 -# 连接池大小,0表示连接数无限制 -DB_POOL_SIZE = 50 -# 连接回收时间(单位:秒) -DB_POOL_RECYCLE = 3600 -# 连接池中没有线程可用时,最多等待的时间(单位:秒) -DB_POOL_TIMEOUT = 30 +# 默认数据源名称 +DB_DEFAULT_SOURCE = 'primary' +# 数据源配置(每个节点完整保留数据库连接字段) +# db_type:数据库类型,可选 mysql、postgresql +# db_host:数据库主机地址 +# db_port:数据库端口 +# db_username:数据库用户名 +# db_password:数据库密码 +# db_database:数据库名称 +# db_echo:是否输出SQLAlchemy SQL日志 +# db_connect_timeout:建立数据库连接的超时时间(秒) +# db_max_overflow:连接池允许的最大溢出连接数 +# db_pool_size:连接池常驻连接数 +# db_pool_recycle:连接回收间隔(秒),-1表示禁用回收 +# db_pool_timeout:获取连接的超时时间(秒) +# db_required:启动时连接失败是否阻止应用启动 +DB_SOURCES = '{ + "primary": { + "db_type": "postgresql", + "db_host": "ruoyi-pg", + "db_port": 5432, + "db_username": "postgres", + "db_password": "root", + "db_database": "ruoyi-fastapi", + "db_echo": true, + "db_connect_timeout": 10, + "db_max_overflow": 10, + "db_pool_size": 50, + "db_pool_recycle": 3600, + "db_pool_timeout": 30, + "db_required": true + } +}' # -------- Redis配置 -------- # Redis主机 diff --git a/ruoyi-fastapi-backend/.env.prod b/ruoyi-fastapi-backend/.env.prod index b75f45e..542f3a9 100644 --- a/ruoyi-fastapi-backend/.env.prod +++ b/ruoyi-fastapi-backend/.env.prod @@ -46,28 +46,39 @@ JWT_REDIS_EXPIRE_MINUTES = 30 # -------- 数据库配置 -------- -# 数据库类型,可选的有'mysql'、'postgresql',默认为'mysql' -DB_TYPE = 'mysql' -# 数据库主机 -DB_HOST = '127.0.0.1' -# 数据库端口 -DB_PORT = 3306 -# 数据库用户名 -DB_USERNAME = 'root' -# 数据库密码 -DB_PASSWORD = 'root' -# 数据库名称 -DB_DATABASE = 'ruoyi-fastapi' -# 是否开启sqlalchemy日志 -DB_ECHO = true -# 允许溢出连接池大小的最大连接数 -DB_MAX_OVERFLOW = 10 -# 连接池大小,0表示连接数无限制 -DB_POOL_SIZE = 50 -# 连接回收时间(单位:秒) -DB_POOL_RECYCLE = 3600 -# 连接池中没有线程可用时,最多等待的时间(单位:秒) -DB_POOL_TIMEOUT = 30 +# 默认数据源名称 +DB_DEFAULT_SOURCE = 'primary' +# 数据源配置(每个节点完整保留数据库连接字段) +# db_type:数据库类型,可选 mysql、postgresql +# db_host:数据库主机地址 +# db_port:数据库端口 +# db_username:数据库用户名 +# db_password:数据库密码 +# db_database:数据库名称 +# db_echo:是否输出SQLAlchemy SQL日志 +# db_connect_timeout:建立数据库连接的超时时间(秒) +# db_max_overflow:连接池允许的最大溢出连接数 +# db_pool_size:连接池常驻连接数 +# db_pool_recycle:连接回收间隔(秒),-1表示禁用回收 +# db_pool_timeout:获取连接的超时时间(秒) +# db_required:启动时连接失败是否阻止应用启动 +DB_SOURCES = '{ + "primary": { + "db_type": "mysql", + "db_host": "127.0.0.1", + "db_port": 3306, + "db_username": "root", + "db_password": "root", + "db_database": "ruoyi-fastapi", + "db_echo": true, + "db_connect_timeout": 10, + "db_max_overflow": 10, + "db_pool_size": 50, + "db_pool_recycle": 3600, + "db_pool_timeout": 30, + "db_required": true + } +}' # -------- Redis配置 -------- # Redis主机 diff --git a/ruoyi-fastapi-backend/alembic/env.py b/ruoyi-fastapi-backend/alembic/env.py index 39c6cf1..5ddb691 100644 --- a/ruoyi-fastapi-backend/alembic/env.py +++ b/ruoyi-fastapi-backend/alembic/env.py @@ -7,11 +7,10 @@ from alembic import context from alembic.migration import MigrationContext from alembic.operations.ops import MigrationScript -from sqlalchemy import pool from sqlalchemy.engine import Connection -from sqlalchemy.ext.asyncio import async_engine_from_config -from config.database import ASYNC_SQLALCHEMY_DATABASE_URL, Base +from config.database import Base, DataSourceRegistry, build_async_sqlalchemy_database_url +from config.env import DataBaseConfig from utils.import_util import ImportUtil # 判断vesrions目录是否存在,如果不存在则创建 @@ -35,7 +34,20 @@ # add your model's MetaData object here # for 'autogenerate' support target_metadata = Base.metadata -# ASYNC_SQLALCHEMY_DATABASE_URL = 'mysql+asyncmy://root:mysqlroot@127.0.0.1:3306/ruoyi-fastapi' + + +def _default_source_url() -> str: + """ + 构建默认数据源的Alembic数据库连接URL + + :return: Alembic数据库连接URL + """ + database_url = build_async_sqlalchemy_database_url(DataBaseConfig.default_source) + return database_url.render_as_string(hide_password=False).replace('%', '%%') + + +ASYNC_SQLALCHEMY_DATABASE_URL = _default_source_url() + # other values from the config, defined by the needs of env.py, alembic_config.set_main_option('sqlalchemy.url', ASYNC_SQLALCHEMY_DATABASE_URL) @@ -111,16 +123,12 @@ async def run_async_migrations() -> None: """ - connectable = async_engine_from_config( - alembic_config.get_section(alembic_config.config_ini_section, {}), - prefix='sqlalchemy.', - poolclass=pool.NullPool, - ) + connectable = DataSourceRegistry.get_async_engine(DataBaseConfig.db_default_source) async with connectable.connect() as connection: await connection.run_sync(do_run_migrations) - await connectable.dispose() + await DataSourceRegistry.dispose_all() def run_migrations_online() -> None: diff --git a/ruoyi-fastapi-backend/cli/core/context_factory.py b/ruoyi-fastapi-backend/cli/core/context_factory.py index 23adf62..94d5331 100644 --- a/ruoyi-fastapi-backend/cli/core/context_factory.py +++ b/ruoyi-fastapi-backend/cli/core/context_factory.py @@ -49,7 +49,7 @@ def suppress_sqlalchemy_logs(self) -> None: 该逻辑同时处理两类来源: - 1. 将 `config.env.DataBaseConfig.db_echo` 强制关闭,避免后续新建 Engine 时 + 1. 将 `config.env.DataBaseConfig.default_source.db_echo` 强制关闭,避免后续新建 Engine 时 继续打开 SQLAlchemy echo。 2. 将已知 SQLAlchemy logger 级别提升到 WARNING,避免已有 logger 配置把 `INFO sqlalchemy.engine.Engine ...` 继续打到标准输出。 @@ -60,8 +60,9 @@ def suppress_sqlalchemy_logs(self) -> None: return env_module = import_module('config.env') database_config = getattr(env_module, 'DataBaseConfig', None) - if database_config is not None and hasattr(database_config, 'db_echo'): - database_config.db_echo = False + default_source = getattr(database_config, 'default_source', None) if database_config is not None else None + if default_source is not None and hasattr(default_source, 'db_echo'): + default_source.db_echo = False for logger_name in ( 'sqlalchemy', 'sqlalchemy.engine', diff --git a/ruoyi-fastapi-backend/cli/runtime/app/support.py b/ruoyi-fastapi-backend/cli/runtime/app/support.py index c7b3a6f..3af2563 100644 --- a/ruoyi-fastapi-backend/cli/runtime/app/support.py +++ b/ruoyi-fastapi-backend/cli/runtime/app/support.py @@ -42,7 +42,7 @@ def build_app_config_snapshot(self) -> dict[str, Any]: """ env_module = self.infrastructure_gateway.get_env_module() app_config = env_module.AppConfig - database_config = env_module.DataBaseConfig + database_config = env_module.DataBaseConfig.default_source log_config = env_module.LogConfig redis_config = env_module.RedisConfig transport_crypto_config = env_module.TransportCryptoConfig diff --git a/ruoyi-fastapi-backend/cli/runtime/config/gateway.py b/ruoyi-fastapi-backend/cli/runtime/config/gateway.py index bd4eb6c..7740e3f 100644 --- a/ruoyi-fastapi-backend/cli/runtime/config/gateway.py +++ b/ruoyi-fastapi-backend/cli/runtime/config/gateway.py @@ -53,7 +53,7 @@ def get_async_session_local() -> Any: :return: 异步数据库会话工厂 """ - return import_module('config.database').AsyncSessionLocal + return import_module('config.database').DataSourceRegistry.session @staticmethod def get_redis_util() -> Any: diff --git a/ruoyi-fastapi-backend/cli/runtime/gen/gateway.py b/ruoyi-fastapi-backend/cli/runtime/gen/gateway.py index e099a6a..3bfb95c 100644 --- a/ruoyi-fastapi-backend/cli/runtime/gen/gateway.py +++ b/ruoyi-fastapi-backend/cli/runtime/gen/gateway.py @@ -35,7 +35,7 @@ def get_async_session_local() -> Any: :return: 异步数据库会话工厂 """ - return import_module('config.database').AsyncSessionLocal + return import_module('config.database').DataSourceRegistry.session @staticmethod def get_page_model() -> Any: @@ -53,7 +53,7 @@ def get_database_config() -> Any: :return: 数据库配置对象 """ - return import_module('config.env').DataBaseConfig + return import_module('config.env').DataBaseConfig.default_source @staticmethod def get_gen_config() -> Any: diff --git a/ruoyi-fastapi-backend/cli/runtime/job/gateway.py b/ruoyi-fastapi-backend/cli/runtime/job/gateway.py index b2491d4..6c538a1 100644 --- a/ruoyi-fastapi-backend/cli/runtime/job/gateway.py +++ b/ruoyi-fastapi-backend/cli/runtime/job/gateway.py @@ -26,7 +26,7 @@ def get_async_session_local() -> Any: :return: 异步数据库会话工厂 """ - return import_module('config.database').AsyncSessionLocal + return import_module('config.database').DataSourceRegistry.session @staticmethod def get_redis_util() -> Any: diff --git a/ruoyi-fastapi-backend/common/aspect/db_seesion.py b/ruoyi-fastapi-backend/common/aspect/db_seesion.py index 3c88fb9..526777b 100644 --- a/ruoyi-fastapi-backend/common/aspect/db_seesion.py +++ b/ruoyi-fastapi-backend/common/aspect/db_seesion.py @@ -1,12 +1,46 @@ +from collections.abc import AsyncIterator +from functools import cache + from fastapi import Depends, params +from sqlalchemy.ext.asyncio import AsyncSession + +from config.database import DataSourceRegistry + + +class DBSessionProvider: + """ + 数据库会话依赖提供者 + """ + + def __init__(self, source_name: str | None = None) -> None: + self.source_name = source_name -from config.get_db import get_db + async def __call__(self) -> AsyncIterator[AsyncSession]: + """ + 创建指定数据源的数据库会话 + + :return: 异步数据库会话 + """ + async with DataSourceRegistry.session(self.source_name) as session: + yield session + + +@cache +def get_db_session_provider(source_name: str | None) -> DBSessionProvider: + """ + 获取指定数据源的数据库会话依赖提供者 + + :param source_name: 数据源名称 + :return: 数据库会话依赖提供者 + """ + return DBSessionProvider(source_name) -def DBSessionDependency() -> params.Depends: # noqa: N802 +def DBSessionDependency(source_name: str | None = None) -> params.Depends: # noqa: N802 """ 数据库会话依赖 + :param source_name: 数据源名称,为空时使用默认数据源 :return: 数据库会话依赖 """ - return Depends(get_db) + return Depends(get_db_session_provider(source_name)) diff --git a/ruoyi-fastapi-backend/common/aspect/pre_auth.py b/ruoyi-fastapi-backend/common/aspect/pre_auth.py index cdcc579..db3bceb 100644 --- a/ruoyi-fastapi-backend/common/aspect/pre_auth.py +++ b/ruoyi-fastapi-backend/common/aspect/pre_auth.py @@ -5,9 +5,9 @@ from fastapi.security import OAuth2PasswordBearer from sqlalchemy.ext.asyncio import AsyncSession +from common.aspect.db_seesion import DBSessionDependency from common.context import RequestContext from config.env import AppConfig -from config.get_db import get_db from exceptions.exception import AuthException from module_admin.entity.vo.user_vo import CurrentUserModel from module_admin.service.login_service import LoginService @@ -80,7 +80,7 @@ def _compile_path_pattern(self, path: str) -> re.Pattern: # 添加开始和结束锚点,确保精确匹配 return re.compile(f'^{pattern_str}$') - async def __call__(self, request: Request, db: AsyncSession = Depends(get_db)) -> CurrentUserModel | None: + async def __call__(self, request: Request, db: AsyncSession = DBSessionDependency()) -> CurrentUserModel | None: """ 执行登录认证校验 diff --git a/ruoyi-fastapi-backend/common/constant.py b/ruoyi-fastapi-backend/common/constant.py index 279132c..aee6467 100644 --- a/ruoyi-fastapi-backend/common/constant.py +++ b/ruoyi-fastapi-backend/common/constant.py @@ -1,6 +1,3 @@ -from config.env import DataBaseConfig - - class CommonConstant: """ 常用常量 @@ -614,16 +611,17 @@ class GenConstant: PARENT_MENU_ID = 'parentMenuId' PARENT_MENU_NAME = 'parentMenuName' GEN_VIEW = 'genView' - COLUMNTYPE_STR = ( - ['character varying', 'varchar', 'character', 'char'] - if DataBaseConfig.db_type == 'postgresql' - else ['char', 'varchar', 'nvarchar', 'varchar2'] - ) - COLUMNTYPE_TEXT = ( - ['text', 'citext'] if DataBaseConfig.db_type == 'postgresql' else ['tinytext', 'text', 'mediumtext', 'longtext'] - ) - COLUMNTYPE_TIME = ( - [ + COLUMNTYPE_STR = { + 'mysql': ['char', 'varchar', 'nvarchar', 'varchar2'], + 'postgresql': ['character varying', 'varchar', 'character', 'char'], + } + COLUMNTYPE_TEXT = { + 'mysql': ['tinytext', 'text', 'mediumtext', 'longtext'], + 'postgresql': ['text', 'citext'], + } + COLUMNTYPE_TIME = { + 'mysql': ['datetime', 'time', 'date', 'timestamp'], + 'postgresql': [ 'date', 'time', 'time with time zone', @@ -632,14 +630,10 @@ class GenConstant: 'timestamp with time zone', 'timestamp without time zone', 'interval', - ] - if DataBaseConfig.db_type == 'postgresql' - else ['datetime', 'time', 'date', 'timestamp'] - ) - COLUMNTYPE_GEOMETRY = ( - ['point', 'line', 'lseg', 'box', 'path', 'polygon', 'circle'] - if DataBaseConfig.db_type == 'postgresql' - else [ + ], + } + COLUMNTYPE_GEOMETRY = { + 'mysql': [ 'geometry', 'point', 'linestring', @@ -648,21 +642,35 @@ class GenConstant: 'multilinestring', 'multipolygon', 'geometrycollection', - ] - ) - COLUMNTYPE_NUMBER = [ - 'tinyint', - 'smallint', - 'mediumint', - 'int', - 'number', - 'integer', - 'bit', - 'bigint', - 'float', - 'double', - 'decimal', - ] + ], + 'postgresql': ['point', 'line', 'lseg', 'box', 'path', 'polygon', 'circle'], + } + COLUMNTYPE_NUMBER = { + 'mysql': [ + 'tinyint', + 'smallint', + 'mediumint', + 'int', + 'number', + 'integer', + 'bit', + 'bigint', + 'float', + 'double', + 'decimal', + ], + 'postgresql': [ + 'smallint', + 'integer', + 'bigint', + 'real', + 'double precision', + 'numeric', + 'decimal', + 'boolean', + 'bit', + ], + } COLUMNNAME_NOT_ADD_SHOW = ['create_by', 'create_time'] COLUMNNAME_NOT_EDIT_SHOW = ['update_by', 'update_time'] COLUMNNAME_NOT_EDIT = ['id', 'create_by', 'create_time', 'del_flag'] @@ -684,8 +692,8 @@ class GenConstant: QUERY_LIKE = 'LIKE' QUERY_EQ = 'EQ' REQUIRE = '1' - DB_TO_SQLALCHEMY_TYPE_MAPPING = ( - { + DB_TO_SQLALCHEMY_TYPE_MAPPING = { + 'postgresql': { 'boolean': 'Boolean', 'smallint': 'SmallInteger', 'integer': 'Integer', @@ -739,9 +747,8 @@ class GenConstant: 'int2vector': 'ARRAY', 'oidvector': 'ARRAY', 'pg_node_tree': 'Text', - } - if DataBaseConfig.db_type == 'postgresql' - else { + }, + 'mysql': { # 数值类型 'TINYINT': 'SmallInteger', 'SMALLINT': 'SmallInteger', @@ -786,10 +793,10 @@ class GenConstant: 'MULTILINESTRING': 'Geometry', 'MULTIPOLYGON': 'Geometry', 'GEOMETRYCOLLECTION': 'Geometry', - } - ) - DB_TO_PYTHON_TYPE_MAPPING = ( - { + }, + } + DB_TO_PYTHON_TYPE_MAPPING = { + 'postgresql': { 'boolean': 'bool', 'smallint': 'int', 'integer': 'int', @@ -843,9 +850,8 @@ class GenConstant: 'int2vector': 'list', 'oidvector': 'list', 'pg_node_tree': 'str', - } - if DataBaseConfig.db_type == 'postgresql' - else { + }, + 'mysql': { # 数值类型 'TINYINT': 'int', 'SMALLINT': 'int', @@ -890,5 +896,5 @@ class GenConstant: 'MULTILINESTRING': 'bytes', 'MULTIPOLYGON': 'bytes', 'GEOMETRYCOLLECTION': 'bytes', - } - ) + }, + } diff --git a/ruoyi-fastapi-backend/config/database.py b/ruoyi-fastapi-backend/config/database.py index e526b42..8846022 100644 --- a/ruoyi-fastapi-backend/config/database.py +++ b/ruoyi-fastapi-backend/config/database.py @@ -1,109 +1,520 @@ -from urllib.parse import quote_plus +from __future__ import annotations -from sqlalchemy import Engine, create_engine -from sqlalchemy.ext.asyncio import AsyncAttrs, AsyncEngine, async_sessionmaker, create_async_engine +import asyncio +from contextlib import asynccontextmanager +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from functools import cache +from typing import TYPE_CHECKING, Any + +from pydantic import SecretStr +from sqlalchemy import URL, Engine, create_engine, text +from sqlalchemy.exc import DBAPIError, InterfaceError, OperationalError +from sqlalchemy.ext.asyncio import ( + AsyncAttrs, + AsyncConnection, + AsyncEngine, + AsyncSession, + async_sessionmaker, + create_async_engine, +) from sqlalchemy.orm import DeclarativeBase, sessionmaker -from config.env import DataBaseConfig +from config.env import DataBaseConfig, DataBaseSettings, DataSourceSettings +from exceptions.exception import ( + DataSourceInitializationException, + DataSourceNotFoundException, + DataSourceUnavailableException, +) +from utils.log_util import logger +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + +_HEALTH_RETRY_COOLDOWN = timedelta(seconds=5) -def build_async_sqlalchemy_database_url() -> str: - """ - 构建异步 SQLAlchemy 数据库连接 URL - :return: 异步 SQLAlchemy 数据库连接 URL +@dataclass(frozen=True, slots=True) +class DatabaseDriverAdapter: """ - if DataBaseConfig.db_type == 'postgresql': - return ( - f'postgresql+asyncpg://{DataBaseConfig.db_username}:{quote_plus(DataBaseConfig.db_password)}@' - f'{DataBaseConfig.db_host}:{DataBaseConfig.db_port}/{DataBaseConfig.db_database}' + 数据库驱动适配器 + """ + + db_type: str + async_driver: str + sync_driver: str + async_connect_timeout_key: str + sync_connect_timeout_key: str + + def build_url(self, config: DataSourceSettings, *, sync: bool) -> URL: + """ + 根据数据源配置构建SQLAlchemy数据库连接URL + + :param config: 数据源配置 + :param sync: 是否构建同步数据库连接URL + :return: SQLAlchemy数据库连接URL + """ + return URL.create( + drivername=self.sync_driver if sync else self.async_driver, + username=config.db_username, + password=_secret_value(config.db_password), + host=config.db_host, + port=int(config.db_port), + database=config.db_database, ) - return ( - f'mysql+asyncmy://{DataBaseConfig.db_username}:{quote_plus(DataBaseConfig.db_password)}@' - f'{DataBaseConfig.db_host}:{DataBaseConfig.db_port}/{DataBaseConfig.db_database}' + + def build_connect_args(self, config: DataSourceSettings, *, sync: bool) -> dict[str, int]: + """ + 构建数据库驱动连接参数 + + :param config: 数据源配置 + :param sync: 是否构建同步数据库连接参数 + :return: 数据库驱动连接参数 + """ + timeout_key = self.sync_connect_timeout_key if sync else self.async_connect_timeout_key + return {timeout_key: config.db_connect_timeout} + + +_DATABASE_DRIVER_ADAPTERS = { + adapter.db_type: adapter + for adapter in ( + DatabaseDriverAdapter( + db_type='mysql', + async_driver='mysql+asyncmy', + sync_driver='mysql+pymysql', + async_connect_timeout_key='connect_timeout', + sync_connect_timeout_key='connect_timeout', + ), + DatabaseDriverAdapter( + db_type='postgresql', + async_driver='postgresql+asyncpg', + sync_driver='postgresql+psycopg2', + async_connect_timeout_key='timeout', + sync_connect_timeout_key='connect_timeout', + ), ) +} -ASYNC_SQLALCHEMY_DATABASE_URL = build_async_sqlalchemy_database_url() +def _secret_value(value: SecretStr | str) -> str: + """ + 获取密码配置的原始值 + + :param value: 密码配置 + :return: 密码原始值 + """ + return value.get_secret_value() if isinstance(value, SecretStr) else value -def build_sync_sqlalchemy_database_url() -> str: +def _database_source(config: DataBaseSettings | DataSourceSettings) -> DataSourceSettings: """ - 构建同步 SQLAlchemy 数据库连接 URL + 获取指定配置对应的数据源配置 - :return: 同步 SQLAlchemy 数据库连接 URL + :param config: 数据库集合配置或单个数据源配置 + :return: 单个数据源配置 """ - if DataBaseConfig.db_type == 'postgresql': - return ( - f'postgresql+psycopg2://{DataBaseConfig.db_username}:{quote_plus(DataBaseConfig.db_password)}@' - f'{DataBaseConfig.db_host}:{DataBaseConfig.db_port}/{DataBaseConfig.db_database}' - ) - return ( - f'mysql+pymysql://{DataBaseConfig.db_username}:{quote_plus(DataBaseConfig.db_password)}@' - f'{DataBaseConfig.db_host}:{DataBaseConfig.db_port}/{DataBaseConfig.db_database}' - ) + return config.get_source() if isinstance(config, DataBaseSettings) else config -SYNC_SQLALCHEMY_DATABASE_URL = build_sync_sqlalchemy_database_url() +def _driver_adapter(config: DataSourceSettings) -> DatabaseDriverAdapter: + """ + 获取数据库驱动适配器 + + :param config: 数据源配置 + :return: 数据库驱动适配器 + """ + adapter = _DATABASE_DRIVER_ADAPTERS.get(config.db_type) + if adapter is None: + raise ValueError(f'不支持的数据库类型:{config.db_type!r}') + return adapter -def create_async_db_engine(echo: bool | None = None) -> AsyncEngine: +def _build_url(config: DataSourceSettings, *, sync: bool) -> URL: """ - 创建异步 SQLAlchemy Engine + 使用对应的数据库驱动适配器构建连接URL - :param echo: 可选,是否输出 SQLAlchemy SQL 日志 - :return: 异步 SQLAlchemy Engine + :param config: 数据源配置 + :param sync: 是否构建同步数据库连接URL + :return: SQLAlchemy数据库连接URL """ + return _driver_adapter(config).build_url(config, sync=sync) + + +def build_async_sqlalchemy_database_url(config: DataBaseSettings | DataSourceSettings | None = None) -> URL: + """ + 构建异步SQLAlchemy数据库连接URL + + :param config: 数据库集合配置或单个数据源配置 + :return: 异步SQLAlchemy数据库连接URL + """ + return _build_url(_database_source(config or DataBaseConfig), sync=False) + + +def build_sync_sqlalchemy_database_url(config: DataBaseSettings | DataSourceSettings | None = None) -> URL: + """ + 构建同步SQLAlchemy数据库连接URL + + :param config: 数据库集合配置或单个数据源配置 + :return: 同步SQLAlchemy数据库连接URL + """ + return _build_url(_database_source(config or DataBaseConfig), sync=True) + + +def _engine_options(config: DataSourceSettings, echo: bool | None = None) -> dict[str, Any]: + """ + 构建数据库引擎连接池参数 + + :param config: 数据源配置 + :param echo: 是否输出SQLAlchemy SQL日志 + :return: 数据库引擎连接池参数 + """ + return { + 'echo': config.db_echo if echo is None else echo, + 'max_overflow': config.db_max_overflow, + 'pool_size': config.db_pool_size, + 'pool_recycle': config.db_pool_recycle, + 'pool_timeout': config.db_pool_timeout, + 'pool_pre_ping': True, + 'pool_use_lifo': True, + } + + +def create_async_db_engine( + echo: bool | None = None, config: DataBaseSettings | DataSourceSettings | None = None +) -> AsyncEngine: + """ + 创建异步SQLAlchemy Engine + + :param echo: 是否输出SQLAlchemy SQL日志 + :param config: 数据库集合配置或单个数据源配置 + :return: 异步SQLAlchemy Engine + """ + source = _database_source(config or DataBaseConfig) + adapter = _driver_adapter(source) return create_async_engine( - ASYNC_SQLALCHEMY_DATABASE_URL, - echo=DataBaseConfig.db_echo if echo is None else echo, - max_overflow=DataBaseConfig.db_max_overflow, - pool_size=DataBaseConfig.db_pool_size, - pool_recycle=DataBaseConfig.db_pool_recycle, - pool_timeout=DataBaseConfig.db_pool_timeout, + adapter.build_url(source, sync=False), + connect_args=adapter.build_connect_args(source, sync=False), + **_engine_options(source, echo), ) -def create_sync_db_engine(echo: bool | None = None) -> Engine: +def create_sync_db_engine( + echo: bool | None = None, config: DataBaseSettings | DataSourceSettings | None = None +) -> Engine: """ - 创建同步 SQLAlchemy Engine + 创建同步SQLAlchemy Engine - :param echo: 可选,是否输出 SQLAlchemy SQL 日志 - :return: 同步 SQLAlchemy Engine + :param echo: 是否输出SQLAlchemy SQL日志 + :param config: 数据库集合配置或单个数据源配置 + :return: 同步SQLAlchemy Engine """ + source = _database_source(config or DataBaseConfig) + adapter = _driver_adapter(source) return create_engine( - SYNC_SQLALCHEMY_DATABASE_URL, - echo=DataBaseConfig.db_echo if echo is None else echo, - max_overflow=DataBaseConfig.db_max_overflow, - pool_size=DataBaseConfig.db_pool_size, - pool_recycle=DataBaseConfig.db_pool_recycle, - pool_timeout=DataBaseConfig.db_pool_timeout, + adapter.build_url(source, sync=True), + connect_args=adapter.build_connect_args(source, sync=True), + **_engine_options(source, echo), ) -def create_async_session_local(engine: AsyncEngine) -> async_sessionmaker: +def create_async_session_factory(engine: AsyncEngine) -> async_sessionmaker[AsyncSession]: + """ + 创建异步Session工厂 + + :param engine: 异步SQLAlchemy Engine + :return: 异步Session工厂 + """ + return async_sessionmaker(bind=engine, autocommit=False, autoflush=False, expire_on_commit=False) + + +def create_sync_session_factory(engine: Engine) -> sessionmaker: """ - 创建异步 Session 工厂 + 创建同步Session工厂 - :param engine: 异步 SQLAlchemy Engine - :return: 异步 Session 工厂 + :param engine: 同步SQLAlchemy Engine + :return: 同步Session工厂 """ - return async_sessionmaker(autocommit=False, autoflush=False, bind=engine) + return sessionmaker(bind=engine, autocommit=False, autoflush=False) -def create_sync_session_local(engine: Engine) -> sessionmaker: +@dataclass(slots=True) +class DataSourceRuntime: """ - 创建同步 Session 工厂 + 数据源运行时状态 + """ + + name: str + config: DataSourceSettings + async_engine: AsyncEngine | None = None + async_session_factory: async_sessionmaker[AsyncSession] | None = None + sync_engine: Engine | None = None + available: bool = False + last_health_check_at: datetime | None = None + next_retry_at: datetime | None = None + health_lock: asyncio.Lock | None = None + - :param engine: 同步 SQLAlchemy Engine - :return: 同步 Session 工厂 +class _DataSourceRegistry: """ - return sessionmaker(autocommit=False, autoflush=False, bind=engine) + 数据源注册中心 + """ + + def __init__(self, settings: DataBaseSettings | None = None) -> None: + self.settings = settings or DataBaseConfig + self._runtimes: dict[str, DataSourceRuntime] = {} + self._initialized = False + self._configs = dict(self.settings.db_sources) + + def _resolve_name(self, name: str | None = None) -> str: + """ + 解析并校验数据源名称 + + :param name: 数据源名称 + :return: 已配置的数据源名称 + """ + source_name = name or self.settings.db_default_source + if source_name not in self._configs: + raise DataSourceNotFoundException(source_name) + return source_name + + def _runtime(self, name: str | None = None) -> DataSourceRuntime: + """ + 获取或创建数据源运行时状态 + + :param name: 数据源名称 + :return: 数据源运行时状态 + """ + source_name = self._resolve_name(name) + runtime = self._runtimes.get(source_name) + if runtime is None: + runtime = DataSourceRuntime( + name=source_name, + config=self._configs[source_name], + health_lock=asyncio.Lock(), + ) + self._runtimes[source_name] = runtime + return runtime + + @staticmethod + def _mark_unavailable(runtime: DataSourceRuntime) -> None: + """ + 标记数据源不可用并设置下次重试时间 + + :param runtime: 数据源运行时状态 + :return: None + """ + now = datetime.now(timezone.utc) + runtime.available = False + runtime.last_health_check_at = now + runtime.next_retry_at = now + _HEALTH_RETRY_COOLDOWN + + @classmethod + def _ensure_async_resources(cls, runtime: DataSourceRuntime) -> None: + """ + 确保数据源异步引擎和Session工厂已创建 + + :param runtime: 数据源运行时状态 + :return: None + """ + if runtime.async_engine is not None and runtime.async_session_factory is not None: + return + try: + engine = create_async_db_engine(config=runtime.config) + session_factory = create_async_session_factory(engine) + except Exception: + cls._mark_unavailable(runtime) + raise DataSourceInitializationException(runtime.name) from None + runtime.async_engine = engine + runtime.async_session_factory = session_factory + + async def initialize(self) -> None: + """ + 初始化并检查所有数据源的连接状态 + + :return: None + """ + if self._initialized: + return + names = tuple(self._configs) + results = await asyncio.gather(*(self._check_health(name) for name in names), return_exceptions=True) + default_name = self._resolve_name() + for name, result in zip(names, results, strict=True): + source_logger = logger.bind( + data_source=name, + database_type=self._configs[name].db_type, + required=self._configs[name].db_required or name == default_name, + ) + if not isinstance(result, BaseException): + source_logger.info('✅ 数据源初始化成功') + continue + config = self._configs[name] + required = config.db_required or name == default_name + if required: + source_logger.error('❌ 必需数据源连接检查失败') + await self.dispose_all() + raise DataSourceInitializationException(name) from None + source_logger.warning('⚠️ 非必需数据源连接检查失败,应用将降级启动') + self._initialized = True + + def get_async_engine(self, name: str | None = None) -> AsyncEngine: + """ + 获取数据源异步引擎 + + :param name: 数据源名称 + :return: 异步SQLAlchemy Engine + """ + runtime = self._runtime(name) + self._ensure_async_resources(runtime) + assert runtime.async_engine is not None + return runtime.async_engine + + def get_sync_engine(self, name: str | None = None) -> Engine: + """ + 获取数据源同步引擎 + :param name: 数据源名称 + :return: 同步SQLAlchemy Engine + """ + runtime = self._runtime(name) + if runtime.sync_engine is None: + try: + runtime.sync_engine = create_sync_db_engine(config=runtime.config) + except Exception: + raise DataSourceInitializationException(runtime.name) from None + return runtime.sync_engine -async_engine = create_async_db_engine() -AsyncSessionLocal = create_async_session_local(async_engine) + async def _check_health(self, name: str) -> None: + """ + 检查指定数据源的连接状态 + + :param name: 数据源名称 + :return: None + """ + runtime = self._runtime(name) + lock = runtime.health_lock + assert lock is not None + async with lock: + await self._check_health_locked(runtime) + + async def _check_health_locked(self, runtime: DataSourceRuntime) -> None: + """ + 在持有健康检查锁时检查数据源连接状态 + + :param runtime: 数据源运行时状态 + :return: None + """ + try: + self._ensure_async_resources(runtime) + assert runtime.async_engine is not None + async with runtime.async_engine.begin() as connection: + await connection.execute(text('SELECT 1')) + except Exception: + self._mark_unavailable(runtime) + raise DataSourceUnavailableException(runtime.name) from None + runtime.available = True + runtime.last_health_check_at = datetime.now(timezone.utc) + runtime.next_retry_at = None + + async def _ensure_available(self, runtime: DataSourceRuntime) -> None: + """ + 确保指定数据源当前可用 + + :param runtime: 数据源运行时状态 + :return: None + """ + lock = runtime.health_lock + assert lock is not None + async with lock: + if runtime.available: + return + now = datetime.now(timezone.utc) + if runtime.next_retry_at is not None and now < runtime.next_retry_at: + raise DataSourceUnavailableException(runtime.name) + await self._check_health_locked(runtime) + logger.bind(data_source=runtime.name).info('✅ 数据源连接已恢复') + + @asynccontextmanager + async def connection(self, name: str | None = None) -> AsyncGenerator[AsyncConnection, None]: + """ + 创建指定数据源的异步数据库连接事务 + + :param name: 数据源名称 + :return: 异步数据库连接 + """ + runtime = self._runtime(name) + await self._ensure_available(runtime) + assert runtime.async_engine is not None + try: + async with runtime.async_engine.begin() as connection: + yield connection + except (InterfaceError, OperationalError): + self._mark_unavailable(runtime) + raise DataSourceUnavailableException(runtime.name) from None + except DBAPIError as exc: + if not exc.connection_invalidated: + raise + self._mark_unavailable(runtime) + raise DataSourceUnavailableException(runtime.name) from None + + @asynccontextmanager + async def session(self, name: str | None = None) -> AsyncGenerator[AsyncSession, None]: + """ + 创建指定数据源的异步数据库会话 + + :param name: 数据源名称 + :return: 异步数据库会话 + """ + runtime = self._runtime(name) + await self._ensure_available(runtime) + factory = runtime.async_session_factory + assert factory is not None + try: + async with factory() as current_db: + yield current_db + except (InterfaceError, OperationalError): + self._mark_unavailable(runtime) + raise DataSourceUnavailableException(runtime.name) from None + except DBAPIError as exc: + if not exc.connection_invalidated: + raise + self._mark_unavailable(runtime) + raise DataSourceUnavailableException(runtime.name) from None + + async def dispose_all(self) -> None: + """ + 释放所有数据源的同步和异步引擎 + + :return: None + """ + runtimes = tuple(self._runtimes.values()) + self._runtimes.clear() + self._initialized = False + for runtime in runtimes: + if runtime.sync_engine is not None: + runtime.sync_engine.dispose() + async_engines = [runtime.async_engine for runtime in runtimes if runtime.async_engine is not None] + await asyncio.gather(*(engine.dispose() for engine in async_engines), return_exceptions=False) + + +DataSourceRegistry = _DataSourceRegistry() class Base(AsyncAttrs, DeclarativeBase): pass + + +@cache +def get_data_source_base(source_name: str) -> type[DeclarativeBase]: + """ + 获取指定数据源独立且可复用的ORM元数据基类 + + :param source_name: 数据源名称 + :return: ORM元数据基类 + """ + DataBaseConfig.get_source(source_name) + + class NamedDataSourceBase(AsyncAttrs, DeclarativeBase): + pass + + NamedDataSourceBase.__name__ = f'{source_name.title().replace("_", "").replace("-", "")}DataSourceBase' + return NamedDataSourceBase diff --git a/ruoyi-fastapi-backend/config/env.py b/ruoyi-fastapi-backend/config/env.py index ee8a299..9c6da80 100644 --- a/ruoyi-fastapi-backend/config/env.py +++ b/ruoyi-fastapi-backend/config/env.py @@ -1,13 +1,17 @@ import argparse import configparser +import json import os +import re import secrets import sys -from typing import Literal +from typing import Annotated, Any, Literal from dotenv import load_dotenv -from pydantic import Field, computed_field, field_validator -from pydantic_settings import BaseSettings +from pydantic import BaseModel, ConfigDict, Field, SecretStr, computed_field, field_validator, model_validator +from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict + +from exceptions.exception import DataSourceNotFoundException class AppSettings(BaseSettings): @@ -58,31 +62,125 @@ def generate_empty_secret_key(cls, value: object) -> object: return value -class DataBaseSettings(BaseSettings): +class DataSourceSettings(BaseModel): """ - 数据库配置 + 单个数据源配置 """ - db_type: Literal['mysql', 'postgresql'] = 'mysql' - db_host: str = '127.0.0.1' - db_port: int = 3306 - db_username: str = 'root' - db_password: str = 'mysqlroot' - db_database: str = 'ruoyi-fastapi' + model_config = ConfigDict(hide_input_in_errors=True) + + db_type: Literal['mysql', 'postgresql'] + db_host: str = Field(min_length=1) + db_port: int = Field(ge=1, le=65535) + db_username: str = Field(min_length=1) + db_password: SecretStr + db_database: str = Field(min_length=1) + db_echo: bool = True - db_max_overflow: int = 10 - db_pool_size: int = 50 - db_pool_recycle: int = 3600 - db_pool_timeout: int = 30 + db_connect_timeout: int = Field(default=10, gt=0) + db_max_overflow: int = Field(default=10, ge=0) + db_pool_size: int = Field(default=20, ge=1) + db_pool_recycle: int = Field(default=3600, ge=-1) + db_pool_timeout: int = Field(default=30, gt=0) + db_required: bool = True @computed_field @property def sqlglot_parse_dialect(self) -> str: + """ + 获取SQLGlot解析方言 + + :return: SQLGlot解析方言 + """ if self.db_type == 'postgresql': return 'postgres' return self.db_type +DATA_SOURCE_NAME_PATTERN = re.compile(r'^[a-z][a-z0-9_-]{0,63}$') + + +class DataBaseSettings(BaseSettings): + """ + 数据库集合配置 + """ + + model_config = SettingsConfigDict(hide_input_in_errors=True) + + db_default_source: str = 'primary' + db_sources: Annotated[dict[str, DataSourceSettings], NoDecode] = Field(default_factory=dict) + + def __init__(self, **values: Any) -> None: + """ + 预解析数据源JSON配置,避免校验异常泄露密码 + + :param values: 数据库配置项 + """ + raw_sources = values.get('db_sources') + if raw_sources is None: + raw_sources = os.environ.get('DB_SOURCES') + if isinstance(raw_sources, str): + try: + values['db_sources'] = json.loads(raw_sources) + except (json.JSONDecodeError, TypeError): + raise ValueError('DB_SOURCES JSON 格式错误') from None + super().__init__(**values) + + @field_validator('db_sources', mode='before') + @classmethod + def parse_sources_json(cls, value: object) -> object: + """ + 解析显式传入的数据源JSON字符串 + + :param value: 数据源配置原始值 + :return: 解析后的数据源配置 + """ + if not isinstance(value, str): + return value + try: + return json.loads(value) + except (json.JSONDecodeError, TypeError): + raise ValueError('DB_SOURCES JSON 格式错误') from None + + @model_validator(mode='after') + def validate_sources(self) -> 'DataBaseSettings': + """ + 校验数据源集合和默认数据源配置 + + :return: 数据库集合配置 + """ + if not self.db_sources: + raise ValueError('DB_SOURCES 不能为空') + if self.db_default_source not in self.db_sources: + raise ValueError(f'默认数据源不存在:{self.db_default_source}') + for name in self.db_sources: + if not DATA_SOURCE_NAME_PATTERN.fullmatch(name): + raise ValueError(f'数据源名称不合法:{name}') + return self + + def get_source(self, name: str | None = None) -> DataSourceSettings: + """ + 获取指定数据源配置 + + :param name: 数据源名称 + :return: 数据源配置 + """ + source_name = name or self.db_default_source + try: + return self.db_sources[source_name] + except KeyError as exc: + raise DataSourceNotFoundException(source_name) from exc + + @property + def default_source(self) -> DataSourceSettings: + """ + 获取默认数据源配置 + + :return: 默认数据源配置 + """ + return self.get_source() + + class RedisSettings(BaseSettings): """ Redis配置 diff --git a/ruoyi-fastapi-backend/config/get_redis.py b/ruoyi-fastapi-backend/config/get_redis.py index ef1e6ba..2b42461 100644 --- a/ruoyi-fastapi-backend/config/get_redis.py +++ b/ruoyi-fastapi-backend/config/get_redis.py @@ -3,7 +3,7 @@ from redis.exceptions import AuthenticationError, RedisError from redis.exceptions import TimeoutError as RedisTimeoutError -from config.database import AsyncSessionLocal +from config.database import DataSourceRegistry from config.env import RedisConfig from module_admin.service.config_service import ConfigService from module_admin.service.dict_service import DictDataService @@ -98,7 +98,7 @@ async def init_sys_dict(cls, redis: FastAPI) -> None: :param redis: redis对象 :return: """ - async with AsyncSessionLocal() as session: + async with DataSourceRegistry.session() as session: await DictDataService.init_cache_sys_dict_services(session, redis) @classmethod @@ -109,5 +109,5 @@ async def init_sys_config(cls, redis: aioredis.Redis) -> None: :param redis: redis对象 :return: """ - async with AsyncSessionLocal() as session: + async with DataSourceRegistry.session() as session: await ConfigService.init_cache_sys_config_services(session, redis) diff --git a/ruoyi-fastapi-backend/config/get_scheduler.py b/ruoyi-fastapi-backend/config/get_scheduler.py index 8837c33..5190ee5 100644 --- a/ruoyi-fastapi-backend/config/get_scheduler.py +++ b/ruoyi-fastapi-backend/config/get_scheduler.py @@ -21,18 +21,12 @@ from apscheduler.util import obj_to_ref from redis import asyncio as aioredis from sqlalchemy.engine import Engine -from sqlalchemy.ext.asyncio import AsyncEngine +from sqlalchemy.orm import sessionmaker import module_task # noqa: F401 from common.constant import LockConstant -from config.database import ( - SYNC_SQLALCHEMY_DATABASE_URL, - create_async_db_engine, - create_async_session_local, - create_sync_db_engine, - create_sync_session_local, -) -from config.env import AppConfig, LogConfig, RedisConfig +from config.database import DataSourceRegistry, create_sync_db_engine +from config.env import AppConfig, DataBaseConfig, LogConfig, RedisConfig from module_admin.dao.job_dao import JobDao from module_admin.entity.vo.job_vo import JobLogModel, JobModel from module_admin.service.job_log_service import JobLogService @@ -134,8 +128,6 @@ class SchedulerUtil: _reacquire_interval_seconds: float = 5.0 _reacquire_jitter_seconds: float = 1.0 _is_closing: bool = False - _sync_async_engine: AsyncEngine | None = None - _sync_async_sessionmaker: Any | None = None _disposed_sync_engines: bool = False # 懒加载的同步 Engine 和 SessionLocal @@ -183,7 +175,8 @@ def _get_jobstore_engine(cls) -> Engine: :return: 同步 Engine """ if cls._jobstore_engine is None: - cls._jobstore_engine = create_sync_db_engine(echo=False) + # JobStore 使用独立 Engine,避免 APScheduler 关闭时释放 Registry 共享的 Engine。 + cls._jobstore_engine = create_sync_db_engine(echo=False, config=DataBaseConfig.get_source()) return cls._jobstore_engine @classmethod @@ -194,7 +187,7 @@ def _get_listener_engine(cls) -> Engine: :return: 同步 Engine """ if cls._listener_engine is None: - cls._listener_engine = create_sync_db_engine() + cls._listener_engine = DataSourceRegistry.get_sync_engine(DataBaseConfig.db_default_source) return cls._listener_engine @classmethod @@ -205,7 +198,11 @@ def _get_session_local(cls) -> Any: :return: SessionLocal """ if cls._session_local is None: - cls._session_local = create_sync_session_local(cls._get_listener_engine()) + cls._session_local = sessionmaker( + autocommit=False, + autoflush=False, + bind=cls._get_listener_engine(), + ) return cls._session_local @classmethod @@ -219,7 +216,7 @@ def _configure_scheduler(cls) -> None: return job_stores = { 'default': MemoryJobStore(), - 'sqlalchemy': SQLAlchemyJobStore(url=SYNC_SQLALCHEMY_DATABASE_URL, engine=cls._get_jobstore_engine()), + 'sqlalchemy': SQLAlchemyJobStore(engine=cls._get_jobstore_engine()), 'redis': RedisJobStore(**redis_config), } executors = {'default': AsyncIOExecutor(), 'processpool': ProcessPoolExecutor(5)} @@ -425,7 +422,7 @@ async def _handle_lock_lost(cls) -> None: cls._sync_pending = False if getattr(scheduler, 'running', False): scheduler.shutdown() - await cls._dispose_sync_async_engine() + cls._scheduler_configured = False cls._dispose_sync_engines() cls._ensure_reacquire_task() @@ -603,22 +600,8 @@ def _get_sync_async_session(cls) -> Any: :return: 异步 Session """ - if not cls._sync_async_sessionmaker: - cls._sync_async_engine = create_async_db_engine(echo=False) - cls._sync_async_sessionmaker = create_async_session_local(cls._sync_async_engine) - return cls._sync_async_sessionmaker() - - @classmethod - async def _dispose_sync_async_engine(cls) -> None: - """ - 释放同步任务使用的异步 Engine - - :return: None - """ - if cls._sync_async_engine: - await cls._sync_async_engine.dispose() - cls._sync_async_engine = None - cls._sync_async_sessionmaker = None + # 每次同步都从注册中心创建新的异步上下文,避免复用已经退出的会话上下文。 + return DataSourceRegistry.session(DataBaseConfig.db_default_source) @classmethod def _dispose_sync_engines(cls) -> None: @@ -632,9 +615,8 @@ def _dispose_sync_engines(cls) -> None: if cls._jobstore_engine: cls._jobstore_engine.dispose() cls._jobstore_engine = None - if cls._listener_engine: - cls._listener_engine.dispose() - cls._listener_engine = None + # Listener 使用 Registry 共享的 Engine,此处只清理引用,避免影响其他服务。 + cls._listener_engine = None cls._session_local = None cls._disposed_sync_engines = True @@ -927,7 +909,6 @@ async def close_system_scheduler(cls) -> None: except asyncio.CancelledError: pass cls._reacquire_task = None - await cls._dispose_sync_async_engine() cls._dispose_sync_engines() if cls._lock_lost_task: cls._lock_lost_task.cancel() @@ -939,6 +920,7 @@ async def close_system_scheduler(cls) -> None: if getattr(scheduler, 'running', False): scheduler.shutdown() logger.info('✅️ 关闭定时任务成功') + cls._scheduler_configured = False # 必须在Redis连接池关闭前,原子释放当前进程持有的Application leader租约 redis = cls._redis cls._redis = None diff --git a/ruoyi-fastapi-backend/config/get_db.py b/ruoyi-fastapi-backend/config/lifecycle.py similarity index 51% rename from ruoyi-fastapi-backend/config/get_db.py rename to ruoyi-fastapi-backend/config/lifecycle.py index 547b2bd..9bd8b1f 100644 --- a/ruoyi-fastapi-backend/config/get_db.py +++ b/ruoyi-fastapi-backend/config/lifecycle.py @@ -1,29 +1,16 @@ -from collections.abc import AsyncGenerator from typing import Literal -from sqlalchemy.ext.asyncio import AsyncSession - -from config.database import AsyncSessionLocal, Base, async_engine +from config.database import Base, DataSourceRegistry from utils.log_util import logger -async def get_db() -> AsyncGenerator[AsyncSession, None]: - """ - 每一个请求处理完毕后会关闭当前连接,不同的请求使用不同的连接 - - :return: - """ - async with AsyncSessionLocal() as current_db: - yield current_db - - async def init_create_table( *, stage: Literal['platform', 'plugin_entities'] = 'platform', log_success_enabled: bool = True, ) -> None: """ - 应用启动时初始化数据库元数据。 + 在默认数据源中初始化平台数据库元数据 :param stage: 建表阶段 :param log_success_enabled: 是否输出阶段成功摘要 @@ -32,17 +19,8 @@ async def init_create_table( if log_success_enabled: message = '🔎 初始化平台数据库元数据...' if stage == 'platform' else '🔎 同步插件实体表...' logger.bind(database_init_stage=stage).info(message) - async with async_engine.begin() as conn: - await conn.run_sync(Base.metadata.create_all) + async with DataSourceRegistry.connection() as connection: + await connection.run_sync(Base.metadata.create_all) if log_success_enabled: message = '✅️ 平台数据库元数据初始化完成' if stage == 'platform' else '✅️ 插件实体表同步完成' logger.bind(database_init_stage=stage).info(message) - - -async def close_async_engine() -> None: - """ - 应用关闭时释放数据库连接池 - - :return: - """ - await async_engine.dispose() diff --git a/ruoyi-fastapi-backend/exceptions/exception.py b/ruoyi-fastapi-backend/exceptions/exception.py index 519df0e..a938dd2 100644 --- a/ruoyi-fastapi-backend/exceptions/exception.py +++ b/ruoyi-fastapi-backend/exceptions/exception.py @@ -65,3 +65,42 @@ class ModelValidatorException(Exception): def __init__(self, data: str | None = None, message: str | None = None) -> None: self.data = data self.message = message + + +class DataSourceException(ServiceException): + """ + 自定义数据源异常DataSourceException + + 异常信息仅包含数据源名称,不暴露驱动异常或可能包含密码的连接URL。 + """ + + def __init__(self, source_name: str, message: str | None = None) -> None: + self.source_name = source_name + super().__init__(message=message or f'数据源异常:{source_name}') + + +class DataSourceNotFoundException(DataSourceException): + """ + 自定义数据源未配置异常DataSourceNotFoundException + """ + + def __init__(self, source_name: str) -> None: + super().__init__(source_name, message=f'数据源未配置:{source_name}') + + +class DataSourceUnavailableException(DataSourceException): + """ + 自定义数据源不可用异常DataSourceUnavailableException + """ + + def __init__(self, source_name: str) -> None: + super().__init__(source_name, message=f'数据源暂不可用:{source_name}') + + +class DataSourceInitializationException(DataSourceException): + """ + 自定义数据源初始化异常DataSourceInitializationException + """ + + def __init__(self, source_name: str) -> None: + super().__init__(source_name, message=f'数据源初始化失败:{source_name}') diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/config_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/config_do.py index 0c40b8d..b10c7e0 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/config_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/config_do.py @@ -27,6 +27,6 @@ class SysConfig(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/dept_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/dept_do.py index c745b56..92d65de 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/dept_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/dept_do.py @@ -23,19 +23,19 @@ class SysDept(Base): leader = Column( String(20), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='负责人', ) phone = Column( String(11), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='联系电话', ) email = Column( String(50), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='邮箱', ) status = Column(CHAR(1), nullable=True, server_default='0', comment='部门状态(0正常 1停用)') diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/dict_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/dict_do.py index 37b2b7b..0a84b8f 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/dict_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/dict_do.py @@ -26,7 +26,7 @@ class SysDictType(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) @@ -47,13 +47,13 @@ class SysDictData(Base): css_class = Column( String(100), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='样式属性(其他样式扩展)', ) list_class = Column( String(100), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='表格回显样式', ) is_default = Column(CHAR(1), nullable=True, server_default='N', comment='是否默认(Y是 N否)') @@ -65,6 +65,6 @@ class SysDictData(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/menu_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/menu_do.py index b7aa4c7..94d31e1 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/menu_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/menu_do.py @@ -23,13 +23,13 @@ class SysMenu(Base): component = Column( String(255), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='组件路径', ) query = Column( String(255), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='路由参数', ) route_name = Column(String(50), nullable=True, server_default="''", comment='路由名称') @@ -41,7 +41,7 @@ class SysMenu(Base): perms = Column( String(100), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='权限标识', ) icon = Column(String(100), nullable=True, server_default='#', comment='菜单图标') diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/notice_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/notice_do.py index 00989e0..749004e 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/notice_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/notice_do.py @@ -20,9 +20,9 @@ class SysNotice(Base): notice_title = Column(String(50), nullable=False, comment='公告标题') notice_type = Column(CHAR(1), nullable=False, comment='公告类型(1通知 2公告)') notice_content = Column( - mysql.LONGBLOB if DataBaseConfig.db_type == 'mysql' else LargeBinary, + mysql.LONGBLOB if DataBaseConfig.default_source.db_type == 'mysql' else LargeBinary, nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type, False), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type, False), comment='公告内容', ) status = Column(CHAR(1), nullable=True, server_default='0', comment='公告状态(0正常 1关闭)') @@ -33,7 +33,7 @@ class SysNotice(Base): remark = Column( String(255), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/post_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/post_do.py index 5b18596..c42d057 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/post_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/post_do.py @@ -27,6 +27,6 @@ class SysPost(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/role_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/role_do.py index 39b0866..0d630cd 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/role_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/role_do.py @@ -27,13 +27,13 @@ class SysRole(Base): comment='数据范围(1:全部数据权限 2:自定数据权限 3:本部门数据权限 4:本部门及以下数据权限)', ) menu_check_strictly = Column( - mysql.TINYINT(display_width=1) if DataBaseConfig.db_type == 'mysql' else SmallInteger, + mysql.TINYINT(display_width=1) if DataBaseConfig.default_source.db_type == 'mysql' else SmallInteger, nullable=True, server_default='1', comment='菜单树选择项是否关联显示', ) dept_check_strictly = Column( - mysql.TINYINT(display_width=1) if DataBaseConfig.db_type == 'mysql' else SmallInteger, + mysql.TINYINT(display_width=1) if DataBaseConfig.default_source.db_type == 'mysql' else SmallInteger, nullable=True, server_default='1', comment='部门树选择项是否关联显示', @@ -47,7 +47,7 @@ class SysRole(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_admin/entity/do/user_do.py b/ruoyi-fastapi-backend/module_admin/entity/do/user_do.py index 4d1cfe1..500e0b7 100644 --- a/ruoyi-fastapi-backend/module_admin/entity/do/user_do.py +++ b/ruoyi-fastapi-backend/module_admin/entity/do/user_do.py @@ -19,7 +19,7 @@ class SysUser(Base): dept_id = Column( BigInteger, nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type, False), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type, False), comment='部门ID', ) user_name = Column(String(30), nullable=False, comment='用户账号') @@ -42,7 +42,7 @@ class SysUser(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_admin/service/file_service.py b/ruoyi-fastapi-backend/module_admin/service/file_service.py index 65bf281..d921cc7 100644 --- a/ruoyi-fastapi-backend/module_admin/service/file_service.py +++ b/ruoyi-fastapi-backend/module_admin/service/file_service.py @@ -11,7 +11,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from common.vo import CrudResponseModel, PageModel -from config.database import AsyncSessionLocal +from config.database import DataSourceRegistry from config.env import UploadConfig from exceptions.exception import ServiceException from module_admin.dao.file_access_dao import FileAccessLogDao @@ -1158,7 +1158,7 @@ async def execute_reconcile_run_services(cls, run_id: str) -> None: :return: None """ try: - async with AsyncSessionLocal() as query_db: + async with DataSourceRegistry.session() as query_db: reconcile_run = await FileInfoDao.get_reconcile_run_by_id(query_db, run_id) if reconcile_run is None or reconcile_run.status != 'running': return @@ -1199,7 +1199,7 @@ async def execute_reconcile_run_services(cls, run_id: str) -> None: ) except Exception as exc: logger.exception(f'文件存储对账任务{run_id}执行失败') - async with AsyncSessionLocal() as query_db: + async with DataSourceRegistry.session() as query_db: try: await FileInfoDao.finish_reconcile_run( query_db, @@ -1221,7 +1221,7 @@ async def run_scheduled_reconcile_services(cls, check_hash: bool = False) -> Non :param check_hash: 是否校验文件摘要 :return: None """ - async with AsyncSessionLocal() as query_db: + async with DataSourceRegistry.session() as query_db: try: reconcile_run = await cls.start_reconcile_run_services( query_db, diff --git a/ruoyi-fastapi-backend/module_admin/service/log_service.py b/ruoyi-fastapi-backend/module_admin/service/log_service.py index fdb334b..f810918 100644 --- a/ruoyi-fastapi-backend/module_admin/service/log_service.py +++ b/ruoyi-fastapi-backend/module_admin/service/log_service.py @@ -10,7 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from common.vo import CrudResponseModel, PageModel -from config.database import AsyncSessionLocal +from config.database import DataSourceRegistry from config.env import LogConfig from exceptions.exception import ServiceException from middlewares.trace_middleware.ctx import TraceCtx @@ -491,7 +491,7 @@ async def _process_messages(cls, redis: aioredis.Redis, stream_name: str, messag """ if not messages: return - async with AsyncSessionLocal() as session: + async with DataSourceRegistry.session() as session: ack_ids: list[str] = [] dedup_event_ids: list[str] = [] try: diff --git a/ruoyi-fastapi-backend/module_admin/service/login_service.py b/ruoyi-fastapi-backend/module_admin/service/login_service.py index 5a0e52b..5ce91cb 100644 --- a/ruoyi-fastapi-backend/module_admin/service/login_service.py +++ b/ruoyi-fastapi-backend/module_admin/service/login_service.py @@ -8,12 +8,12 @@ from sqlalchemy import Row from sqlalchemy.ext.asyncio import AsyncSession +from common.aspect.db_seesion import DBSessionDependency from common.constant import CommonConstant, MenuConstant from common.context import RequestContext from common.enums import PasswordCharacterType, RedisInitKeyConfig from common.vo import CrudResponseModel from config.env import AppConfig, JwtConfig -from config.get_db import get_db from exceptions.exception import AuthException, LoginException, ServiceException from module_admin.dao.login_dao import login_by_account from module_admin.dao.user_dao import UserDao @@ -211,7 +211,10 @@ async def create_access_token(cls, data: dict, expires_delta: timedelta | None = @classmethod async def get_current_user( - cls, request: Request = Request, token: str = Depends(oauth2_scheme), query_db: AsyncSession = Depends(get_db) + cls, + request: Request = Request, + token: str = Depends(oauth2_scheme), + query_db: AsyncSession = DBSessionDependency(), ) -> CurrentUserModel: """ 根据token获取当前用户信息 diff --git a/ruoyi-fastapi-backend/module_generator/controller/gen_controller.py b/ruoyi-fastapi-backend/module_generator/controller/gen_controller.py index a62e82c..d3db398 100644 --- a/ruoyi-fastapi-backend/module_generator/controller/gen_controller.py +++ b/ruoyi-fastapi-backend/module_generator/controller/gen_controller.py @@ -21,6 +21,7 @@ from module_generator.entity.vo.gen_vo import ( DeleteGenTableModel, EditGenTableModel, + GenDataSourceModel, GenTableDbRowModel, GenTableDetailModel, GenTablePageQueryModel, @@ -34,6 +35,16 @@ gen_controller = APIRouterPro(prefix='/tool/gen', order_num=17, tags=['代码生成'], dependencies=[PreAuthDependency()]) +@gen_controller.get( + '/dataSources', + summary='获取代码生成数据源选项接口', + response_model=DataResponseModel[list[GenDataSourceModel]], + dependencies=[UserInterfaceAuthDependency('tool:gen:list')], +) +async def get_gen_data_sources(request: Request) -> Response: + return ResponseUtil.success(data=GenTableService.get_data_source_list_services()) + + @gen_controller.get( '/list', summary='获取代码生成表分页列表接口', @@ -93,9 +104,12 @@ async def import_gen_table( tables: Annotated[str, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()], current_user: Annotated[CurrentUserModel, CurrentUserDependency()], + source_name: Annotated[str | None, Query(alias='dataSourceName')] = None, ) -> Response: table_names = tables.split(',') if tables else [] - add_gen_table_list = await GenTableService.get_gen_db_table_list_by_name_services(query_db, table_names) + add_gen_table_list = await GenTableService.get_gen_db_table_list_by_name_services( + query_db, table_names, source_name + ) add_gen_table_result = await GenTableService.import_gen_table_services(query_db, add_gen_table_list, current_user) logger.info(add_gen_table_result.message) @@ -167,8 +181,9 @@ async def create_table( sql: Annotated[str, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()], current_user: Annotated[CurrentUserModel, CurrentUserDependency()], + source_name: Annotated[str | None, Query(alias='dataSourceName')] = None, ) -> Response: - create_table_result = await GenTableService.create_table_services(query_db, sql, current_user) + create_table_result = await GenTableService.create_table_services(query_db, sql, current_user, source_name) logger.info(create_table_result.message) return ResponseUtil.success(msg=create_table_result.message) @@ -199,9 +214,10 @@ async def batch_gen_code( request: Request, tables: Annotated[str, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()], + source_name: Annotated[str | None, Query(alias='dataSourceName')] = None, ) -> Response: table_names = tables.split(',') if tables else [] - batch_gen_code_result = await GenTableService.batch_gen_code_services(query_db, table_names) + batch_gen_code_result = await GenTableService.batch_gen_code_services(query_db, table_names, source_name) logger.info('生成代码成功') return ResponseUtil.streaming(data=bytes2file_response(batch_gen_code_result)) @@ -225,11 +241,12 @@ async def gen_code_local( request: Request, table_name: Annotated[str, Path(description='表名称')], query_db: Annotated[AsyncSession, DBSessionDependency()], + source_name: Annotated[str | None, Query(alias='dataSourceName')] = None, ) -> Response: if not GenConfig.allow_overwrite: logger.error('【系统预设】不允许生成文件覆盖到本地') return ResponseUtil.error('【系统预设】不允许生成文件覆盖到本地') - gen_code_local_result = await GenTableService.generate_code_services(query_db, table_name) + gen_code_local_result = await GenTableService.generate_code_services(query_db, table_name, source_name) logger.info(gen_code_local_result.message) return ResponseUtil.success(msg=gen_code_local_result.message) @@ -249,7 +266,7 @@ async def query_detail_gen_table( query_db: Annotated[AsyncSession, DBSessionDependency()], ) -> Response: gen_table = await GenTableService.get_gen_table_by_id_services(query_db, table_id) - gen_tables = await GenTableService.get_gen_table_all_services(query_db) + gen_tables = await GenTableService.get_gen_table_all_services(query_db, gen_table.data_source_name) gen_columns = await GenTableColumnService.get_gen_table_column_list_by_table_id_services(query_db, table_id) gen_table_detail_result = {'info': gen_table, 'rows': gen_columns, 'tables': gen_tables} logger.info(f'获取table_id为{table_id}的信息成功') @@ -290,8 +307,9 @@ async def sync_db( request: Request, table_name: Annotated[str, Path(description='表名称')], query_db: Annotated[AsyncSession, DBSessionDependency()], + source_name: Annotated[str | None, Query(alias='dataSourceName')] = None, ) -> Response: - sync_db_result = await GenTableService.sync_db_services(query_db, table_name) + sync_db_result = await GenTableService.sync_db_services(query_db, table_name, source_name) logger.info(sync_db_result.message) return ResponseUtil.success(data=sync_db_result.message) diff --git a/ruoyi-fastapi-backend/module_generator/dao/gen_dao.py b/ruoyi-fastapi-backend/module_generator/dao/gen_dao.py index a855063..0bc8508 100644 --- a/ruoyi-fastapi-backend/module_generator/dao/gen_dao.py +++ b/ruoyi-fastapi-backend/module_generator/dao/gen_dao.py @@ -2,13 +2,13 @@ from datetime import datetime, time from typing import Any -from sqlalchemy import Row, delete, func, select, text, update +from sqlalchemy import Row, bindparam, delete, func, select, text, update from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import selectinload from sqlglot.expressions import Expression from common.vo import PageModel -from config.env import DataBaseConfig +from config.env import DataBaseConfig, DataSourceSettings from module_generator.entity.do.gen_do import GenTable, GenTableColumn from module_generator.entity.vo.gen_vo import ( GenTableBaseModel, @@ -47,18 +47,24 @@ async def get_gen_table_by_id(cls, db: AsyncSession, table_id: int) -> GenTable return gen_table_info @classmethod - async def get_gen_table_by_name(cls, db: AsyncSession, table_name: str) -> GenTable | None: + async def get_gen_table_by_name(cls, db: AsyncSession, table_name: str, source_name: str) -> GenTable | None: """ 根据业务表名称获取需要生成的业务表信息 :param db: orm对象 :param table_name: 业务表名称 + :param source_name: 数据源名称 :return: 需要生成的业务表信息对象 """ gen_table_info = ( ( await db.execute( - select(GenTable).options(selectinload(GenTable.columns)).where(GenTable.table_name == table_name) + select(GenTable) + .options(selectinload(GenTable.columns)) + .where( + GenTable.table_name == table_name, + GenTable.data_source_name == source_name, + ) ) ) .scalars() @@ -68,28 +74,35 @@ async def get_gen_table_by_name(cls, db: AsyncSession, table_name: str) -> GenTa return gen_table_info @classmethod - async def get_gen_table_all(cls, db: AsyncSession) -> Sequence[GenTable]: + async def get_gen_table_all(cls, db: AsyncSession, source_name: str | None = None) -> Sequence[GenTable]: """ 获取所有业务表信息 :param db: orm对象 + :param source_name: 数据源名称 :return: 所有业务表信息 """ - gen_table_all = (await db.execute(select(GenTable).options(selectinload(GenTable.columns)))).scalars().all() + query = select(GenTable).options(selectinload(GenTable.columns)) + if source_name: + query = query.where(GenTable.data_source_name == source_name) + gen_table_all = (await db.execute(query)).scalars().all() return gen_table_all @classmethod - async def create_table_by_sql_dao(cls, db: AsyncSession, sql_statements: list[Expression]) -> None: + async def create_table_by_sql_dao( + cls, db: AsyncSession, sql_statements: list[Expression], *, source_config: DataSourceSettings + ) -> None: """ 根据sql语句创建表结构 :param db: orm对象 :param sql_statements: sql语句的ast列表 + :param source_config: 目标数据源配置 :return: """ for sql_statement in sql_statements: - sql = sql_statement.sql(dialect=DataBaseConfig.sqlglot_parse_dialect) + sql = sql_statement.sql(dialect=source_config.sqlglot_parse_dialect) await db.execute(text(sql)) @classmethod @@ -120,6 +133,7 @@ async def get_gen_table_list( ) if query_object.begin_time and query_object.end_time else True, + GenTable.data_source_name == query_object.data_source_name if query_object.data_source_name else True, ) .distinct() ) @@ -129,9 +143,29 @@ async def get_gen_table_list( return gen_table_list + @classmethod + async def get_gen_table_names(cls, db: AsyncSession, source_name: str | None = None) -> set[str]: + """ + 获取控制库中指定数据源已导入的业务表名称 + + :param db: orm对象 + :param source_name: 数据源名称 + :return: 已导入的业务表名称集合 + """ + query = select(GenTable.table_name) + if source_name: + query = query.where(GenTable.data_source_name == source_name) + return {name for name in (await db.execute(query)).scalars().all() if name} + @classmethod async def get_gen_db_table_list( - cls, db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False + cls, + db: AsyncSession, + query_object: GenTablePageQueryModel, + is_page: bool = False, + *, + excluded_table_names: set[str] | None = None, + source_config: DataSourceSettings | None = None, ) -> PageModel | list[dict[str, Any]]: """ 根据查询参数获取数据库列表信息 @@ -139,10 +173,13 @@ async def get_gen_db_table_list( :param db: orm对象 :param query_object: 查询参数对象 :param is_page: 是否开启分页 + :param excluded_table_names: 需要排除的已导入表名称集合 + :param source_config: 目标数据源配置 :return: 数据库列表信息对象 """ - query_params: dict[str, str] = {} - if DataBaseConfig.db_type == 'postgresql': + source_config = source_config or DataBaseConfig.default_source + query_params: dict[str, Any] = {} + if source_config.db_type == 'postgresql': query_sql = """ table_name as table_name, table_comment as table_comment, @@ -153,7 +190,6 @@ async def get_gen_db_table_list( where table_name not like 'apscheduler_%' and table_name not like 'gen_%' - and table_name not in (select table_name from gen_table) """ else: query_sql = r""" @@ -167,28 +203,33 @@ async def get_gen_db_table_list( table_schema = (select database()) and table_name not like 'apscheduler\_%' and table_name not like 'gen\_%' - and table_name not in (select table_name from gen_table) """ + if excluded_table_names: + query_sql += ' and table_name not in :excluded_table_names' + query_params['excluded_table_names'] = tuple(excluded_table_names) if query_object.table_name: - query_sql += """ and lower(table_name) like lower(concat('%', :table_name, '%'))""" + query_sql += " and lower(table_name) like lower(concat('%', :table_name, '%'))" query_params['table_name'] = query_object.table_name if query_object.table_comment: - query_sql += """ and lower(table_comment) like lower(concat('%', :table_comment, '%'))""" + query_sql += " and lower(table_comment) like lower(concat('%', :table_comment, '%'))" query_params['table_comment'] = query_object.table_comment if query_object.begin_time: - if DataBaseConfig.db_type == 'postgresql': - query_sql += """ and create_time::date >= to_date(:begin_time, 'yyyy-MM-dd')""" + if source_config.db_type == 'postgresql': + query_sql += " and create_time::date >= to_date(:begin_time, 'yyyy-MM-dd')" else: - query_sql += """ and date_format(create_time, '%Y%m%d') >= date_format(:begin_time, '%Y%m%d')""" + query_sql += " and date_format(create_time, '%Y%m%d') >= date_format(:begin_time, '%Y%m%d')" query_params['begin_time'] = query_object.begin_time if query_object.end_time: - if DataBaseConfig.db_type == 'postgresql': - query_sql += """ and create_time::date <= to_date(:end_time, 'yyyy-MM-dd')""" + if source_config.db_type == 'postgresql': + query_sql += " and create_time::date <= to_date(:end_time, 'yyyy-MM-dd')" else: - query_sql += """ and date_format(create_time, '%Y%m%d') <= date_format(:end_time, '%Y%m%d')""" + query_sql += " and date_format(create_time, '%Y%m%d') <= date_format(:end_time, '%Y%m%d')" query_params['end_time'] = query_object.end_time - query_sql += """ order by create_time desc""" - query = select(text(query_sql).bindparams(**query_params)) + query_sql += ' order by create_time desc' + statement = text(query_sql) + if excluded_table_names: + statement = statement.bindparams(bindparam('excluded_table_names', expanding=True)) + query = select(statement.bindparams(**query_params)) gen_db_table_list: PageModel | list[dict[str, Any]] = await PageUtil.paginate( db, query, query_object.page_num, query_object.page_size, is_page ) @@ -196,15 +237,19 @@ async def get_gen_db_table_list( return gen_db_table_list @classmethod - async def get_gen_db_table_list_by_names(cls, db: AsyncSession, table_names: list[str]) -> Sequence[Row]: + async def get_gen_db_table_list_by_names( + cls, db: AsyncSession, table_names: list[str], source_config: DataSourceSettings | None = None + ) -> Sequence[Row]: """ 根据业务表名称组获取数据库列表信息 :param db: orm对象 :param table_names: 业务表名称组 + :param source_config: 目标数据源配置 :return: 数据库列表信息对象 """ - if DataBaseConfig.db_type == 'postgresql': + source_config = source_config or DataBaseConfig.default_source + if source_config.db_type == 'postgresql': query_sql = """ select table_name as table_name, @@ -216,7 +261,7 @@ async def get_gen_db_table_list_by_names(cls, db: AsyncSession, table_names: lis where table_name not like 'qrtz_%' and table_name not like 'gen_%' - and table_name = any(:table_names) + and table_name in :table_names """ else: query_sql = r""" @@ -233,7 +278,7 @@ async def get_gen_db_table_list_by_names(cls, db: AsyncSession, table_names: lis and table_schema = (select database()) and table_name in :table_names """ - query = text(query_sql).bindparams(table_names=tuple(table_names)) + query = text(query_sql).bindparams(bindparam('table_names', value=table_names, expanding=True)) gen_db_table_list = (await db.execute(query)).fetchall() return gen_db_table_list @@ -303,15 +348,19 @@ async def get_gen_table_column_list_by_table_id(cls, db: AsyncSession, table_id: return gen_table_column_list @classmethod - async def get_gen_db_table_columns_by_name(cls, db: AsyncSession, table_name: str) -> Sequence[Row]: + async def get_gen_db_table_columns_by_name( + cls, db: AsyncSession, table_name: str, source_config: DataSourceSettings | None = None + ) -> Sequence[Row]: """ 根据业务表名称获取业务表字段列表信息 :param db: orm对象 :param table_name: 业务表名称 + :param source_config: 目标数据源配置 :return: 业务表字段列表信息对象 """ - if DataBaseConfig.db_type == 'postgresql': + source_config = source_config or DataBaseConfig.default_source + if source_config.db_type == 'postgresql': query_sql = """ select column_name, is_required, is_pk, sort, column_comment, is_increment, column_type diff --git a/ruoyi-fastapi-backend/module_generator/entity/do/gen_do.py b/ruoyi-fastapi-backend/module_generator/entity/do/gen_do.py index f2163b6..a5b386b 100644 --- a/ruoyi-fastapi-backend/module_generator/entity/do/gen_do.py +++ b/ruoyi-fastapi-backend/module_generator/entity/do/gen_do.py @@ -19,16 +19,17 @@ class GenTable(Base): table_id = Column(BigInteger, primary_key=True, nullable=False, autoincrement=True, comment='编号') table_name = Column(String(200), nullable=True, server_default="''", comment='表名称') table_comment = Column(String(500), nullable=True, server_default="''", comment='表描述') + data_source_name = Column(String(64), nullable=False, server_default='primary', comment='目标数据源名称') sub_table_name = Column( String(64), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='关联子表的表名', ) sub_table_fk_name = Column( String(64), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='子表关联的外键名', ) class_name = Column(String(100), nullable=True, server_default="''", comment='实体类名称') @@ -54,7 +55,7 @@ class GenTable(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/module_generator/entity/vo/gen_vo.py b/ruoyi-fastapi-backend/module_generator/entity/vo/gen_vo.py index 4d30044..9486dcf 100644 --- a/ruoyi-fastapi-backend/module_generator/entity/vo/gen_vo.py +++ b/ruoyi-fastapi-backend/module_generator/entity/vo/gen_vo.py @@ -19,6 +19,7 @@ class GenTableBaseModel(BaseModel): table_id: int | None = Field(default=None, description='编号') table_name: str | None = Field(default=None, description='表名称') table_comment: str | None = Field(default=None, description='表描述') + data_source_name: str | None = Field(default=None, description='目标数据源名称') sub_table_name: str | None = Field(default=None, description='关联子表的表名') sub_table_fk_name: str | None = Field(default=None, description='子表关联的外键名') class_name: str | None = Field(default=None, description='实体类名称') @@ -103,6 +104,18 @@ class GenTableDbRowModel(BaseModel): update_time: datetime | None = Field(default=None, description='更新时间') +class GenDataSourceModel(BaseModel): + """ + 代码生成数据源选项模型 + """ + + model_config = ConfigDict(alias_generator=to_camel) + + name: str + db_type: str = Field(description='数据库类型') + is_default: bool = False + + class GenTableModel(GenTableBaseModel): """ 代码生成业务表模型 diff --git a/ruoyi-fastapi-backend/module_generator/service/gen_service.py b/ruoyi-fastapi-backend/module_generator/service/gen_service.py index e9ec609..8165a43 100644 --- a/ruoyi-fastapi-backend/module_generator/service/gen_service.py +++ b/ruoyi-fastapi-backend/module_generator/service/gen_service.py @@ -12,6 +12,7 @@ from common.constant import GenConstant from common.vo import CrudResponseModel, PageModel +from config.database import DataSourceRegistry from config.env import DataBaseConfig, GenConfig from exceptions.exception import ServiceException from module_admin.entity.vo.user_vo import CurrentUserModel @@ -19,6 +20,7 @@ from module_generator.entity.vo.gen_vo import ( DeleteGenTableModel, EditGenTableModel, + GenDataSourceModel, GenTableColumnModel, GenTableModel, GenTablePageQueryModel, @@ -33,6 +35,29 @@ class GenTableService: 代码生成业务表服务层 """ + @classmethod + def get_data_source_list_services(cls) -> list[GenDataSourceModel]: + """ + 获取代码生成可用的数据源列表 + + :return: 数据源选项列表 + """ + default_name = DataBaseConfig.db_default_source + return [ + GenDataSourceModel(name=name, dbType=config.db_type, isDefault=name == default_name) + for name, config in DataBaseConfig.db_sources.items() + ] + + @staticmethod + def _source_name(source_name: str | None) -> str: + """ + 获取代码生成操作使用的数据源名称 + + :param source_name: 数据源名称 + :return: 指定的数据源名称或默认数据源名称 + """ + return source_name or DataBaseConfig.db_default_source + @classmethod async def get_gen_table_list_services( cls, query_db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False @@ -51,7 +76,10 @@ async def get_gen_table_list_services( @classmethod async def get_gen_db_table_list_services( - cls, query_db: AsyncSession, query_object: GenTablePageQueryModel, is_page: bool = False + cls, + query_db: AsyncSession, + query_object: GenTablePageQueryModel, + is_page: bool = False, ) -> PageModel | list[dict[str, Any]]: """ 获取数据库列表信息service @@ -61,28 +89,47 @@ async def get_gen_db_table_list_services( :param is_page: 是否开启分页 :return: 数据库列表信息对象 """ - gen_db_table_list_result = await GenTableDao.get_gen_db_table_list(query_db, query_object, is_page) - - return gen_db_table_list_result + source_name = cls._source_name(query_object.data_source_name) + source_config = DataBaseConfig.get_source(source_name) + excluded = await GenTableDao.get_gen_table_names(query_db, source_name) + async with DataSourceRegistry.session(source_name) as target_db: + gen_db_table_list_result = await GenTableDao.get_gen_db_table_list( + target_db, + query_object, + is_page, + excluded_table_names=excluded, + source_config=source_config, + ) + return gen_db_table_list_result @classmethod async def get_gen_db_table_list_by_name_services( - cls, query_db: AsyncSession, table_names: list[str] + cls, query_db: AsyncSession, table_names: list[str], source_name: str | None = None ) -> list[GenTableModel]: """ 根据表名称组获取数据库列表信息service :param query_db: orm对象 :param table_names: 表名称组 + :param source_name: 数据源名称 :return: 数据库列表信息对象 """ - gen_db_table_list_result = await GenTableDao.get_gen_db_table_list_by_names(query_db, table_names) - - return [GenTableModel(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_db_table_list_result)] + source_name = cls._source_name(source_name) + source_config = DataBaseConfig.get_source(source_name) + async with DataSourceRegistry.session(source_name) as target_db: + rows = await GenTableDao.get_gen_db_table_list_by_names(target_db, table_names, source_config) + result = [ + GenTableModel(**gen_table, dataSourceName=source_name) + for gen_table in CamelCaseUtil.transform_result(rows) + ] + return result @classmethod async def import_gen_table_services( - cls, query_db: AsyncSession, gen_table_list: list[GenTableModel], current_user: CurrentUserModel + cls, + query_db: AsyncSession, + gen_table_list: list[GenTableModel], + current_user: CurrentUserModel, ) -> CrudResponseModel: """ 导入表结构service @@ -92,20 +139,26 @@ async def import_gen_table_services( :param current_user: 当前用户信息对象 :return: 导入结果 """ + source_name = cls._source_name(gen_table_list[0].data_source_name if gen_table_list else None) + source_config = DataBaseConfig.get_source(source_name) try: - for table in gen_table_list: - table_name = table.table_name - GenUtils.init_table(table, current_user.user.user_name) - add_gen_table = await GenTableDao.add_gen_table_dao(query_db, table) - if add_gen_table: - table.table_id = add_gen_table.table_id - gen_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name(query_db, table_name) - for column in [ - GenTableColumnModel(**gen_table_column) - for gen_table_column in CamelCaseUtil.transform_result(gen_table_columns) - ]: - GenUtils.init_column_field(column, table) - await GenTableColumnDao.add_gen_table_column_dao(query_db, column) + async with DataSourceRegistry.session(source_name) as target_db: + for table in gen_table_list: + table_name = table.table_name + table.data_source_name = source_name + GenUtils.init_table(table, current_user.user.user_name) + add_gen_table = await GenTableDao.add_gen_table_dao(query_db, table) + if add_gen_table: + table.table_id = add_gen_table.table_id + gen_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name( + target_db, table_name, source_config + ) + for column in [ + GenTableColumnModel(**gen_table_column) + for gen_table_column in CamelCaseUtil.transform_result(gen_table_columns) + ]: + GenUtils.init_column_field(column, table) + await GenTableColumnDao.add_gen_table_column_dao(query_db, column) await query_db.commit() return CrudResponseModel(is_success=True, message='导入成功') except Exception as e: @@ -125,6 +178,8 @@ async def edit_gen_table_services(cls, query_db: AsyncSession, page_object: Edit gen_table_info = await cls.get_gen_table_by_id_services(query_db, page_object.table_id) if gen_table_info.table_id: try: + # 数据源归属在导入时确定,编辑生成配置时不得隐式迁移到其他数据库。 + edit_gen_table['data_source_name'] = gen_table_info.data_source_name edit_gen_table['options'] = json.dumps(edit_gen_table.get('params')) await GenTableDao.edit_gen_table_dao(query_db, edit_gen_table) for gen_table_column in page_object.columns: @@ -183,21 +238,24 @@ async def get_gen_table_by_id_services(cls, query_db: AsyncSession, table_id: in return result @classmethod - async def get_gen_table_all_services(cls, query_db: AsyncSession) -> list[GenTableModel]: + async def get_gen_table_all_services( + cls, query_db: AsyncSession, source_name: str | None = None + ) -> list[GenTableModel]: """ 获取所有业务表信息service :param query_db: orm对象 + :param source_name: 数据源名称 :return: 所有业务表信息 """ - gen_table_all = await GenTableDao.get_gen_table_all(query_db) + gen_table_all = await GenTableDao.get_gen_table_all(query_db, source_name) result = [GenTableModel(**gen_table) for gen_table in CamelCaseUtil.transform_result(gen_table_all)] return result @classmethod async def create_table_services( - cls, query_db: AsyncSession, sql: str, current_user: CurrentUserModel + cls, query_db: AsyncSession, sql: str, current_user: CurrentUserModel, source_name: str | None = None ) -> CrudResponseModel: """ 创建表结构service @@ -205,14 +263,19 @@ async def create_table_services( :param query_db: orm对象 :param sql: 建表语句 :param current_user: 当前用户信息对象 + :param source_name: 数据源名称 :return: 创建表结构结果 """ - sql_statements = sqlglot_parse(sql, dialect=DataBaseConfig.sqlglot_parse_dialect) + source_name = cls._source_name(source_name) + target_config = DataBaseConfig.get_source(source_name) + sql_statements = sqlglot_parse(sql, dialect=target_config.sqlglot_parse_dialect) if cls.__is_valid_create_table(sql_statements): try: table_names = cls.__get_table_names(sql_statements) - await GenTableDao.create_table_by_sql_dao(query_db, sql_statements) - gen_table_list = await cls.get_gen_db_table_list_by_name_services(query_db, table_names) + async with DataSourceRegistry.session(source_name) as target_db: + await GenTableDao.create_table_by_sql_dao(target_db, sql_statements, source_config=target_config) + await target_db.commit() + gen_table_list = await cls.get_gen_db_table_list_by_name_services(query_db, table_names, source_name) await cls.import_gen_table_services(query_db, gen_table_list, current_user) return CrudResponseModel(is_success=True, message='创建表结构成功') @@ -277,16 +340,20 @@ async def preview_code_services(cls, query_db: AsyncSession, table_id: int) -> d return preview_code_result @classmethod - async def generate_code_services(cls, query_db: AsyncSession, table_name: str) -> CrudResponseModel: + async def generate_code_services( + cls, query_db: AsyncSession, table_name: str, source_name: str | None = None + ) -> CrudResponseModel: """ 生成代码至指定路径service :param query_db: orm对象 :param table_name: 业务表名称 + :param source_name: 数据源名称 :return: 生成代码结果 """ + source_name = cls._source_name(source_name) env = TemplateInitializer.init_jinja2() - render_info = await cls.__get_gen_render_info(query_db, table_name) + render_info = await cls.__get_gen_render_info(query_db, table_name, source_name) try: for template in render_info[0]: render_content = env.get_template(template).render(**render_info[2]) @@ -300,19 +367,23 @@ async def generate_code_services(cls, query_db: AsyncSession, table_name: str) - return CrudResponseModel(is_success=True, message='生成代码成功') @classmethod - async def batch_gen_code_services(cls, query_db: AsyncSession, table_names: list[str]) -> bytes: + async def batch_gen_code_services( + cls, query_db: AsyncSession, table_names: list[str], source_name: str | None = None + ) -> bytes: """ 批量生成代码service :param query_db: orm对象 :param table_names: 业务表名称组 + :param source_name: 数据源名称 :return: 下载代码结果 """ + source_name = cls._source_name(source_name) zip_buffer = io.BytesIO() with zipfile.ZipFile(zip_buffer, 'w', zipfile.ZIP_DEFLATED) as zip_file: for table_name in table_names: env = TemplateInitializer.init_jinja2() - render_info = await cls.__get_gen_render_info(query_db, table_name) + render_info = await cls.__get_gen_render_info(query_db, table_name, source_name) for template_file, output_file in zip(render_info[0], render_info[1], strict=False): render_content = env.get_template(template_file).render(**render_info[2]) zip_file.writestr(output_file, render_content) @@ -322,16 +393,19 @@ async def batch_gen_code_services(cls, query_db: AsyncSession, table_names: list return zip_data @classmethod - async def __get_gen_render_info(cls, query_db: AsyncSession, table_name: str) -> list: + async def __get_gen_render_info( + cls, query_db: AsyncSession, table_name: str, source_name: str | None = None + ) -> list: """ 获取生成代码渲染模板相关信息 :param query_db: orm对象 :param table_name: 业务表名称 + :param source_name: 数据源名称 :return: 生成代码渲染模板相关信息 """ gen_table = GenTableModel( - **CamelCaseUtil.transform_result(await GenTableDao.get_gen_table_by_name(query_db, table_name)) + **CamelCaseUtil.transform_result(await GenTableDao.get_gen_table_by_name(query_db, table_name, source_name)) ) await cls.set_sub_table(query_db, gen_table) await cls.set_pk_column(gen_table) @@ -357,22 +431,30 @@ def __get_gen_path(cls, gen_table: GenTableModel, template: str) -> str: return os.path.join(gen_path, TemplateUtils.get_file_name(template, gen_table)) @classmethod - async def sync_db_services(cls, query_db: AsyncSession, table_name: str) -> CrudResponseModel: + async def sync_db_services( + cls, query_db: AsyncSession, table_name: str, source_name: str | None = None + ) -> CrudResponseModel: """ 同步数据库service :param query_db: orm对象 :param table_name: 业务表名称 + :param source_name: 数据源名称 :return: 同步数据库结果 """ - gen_table = await GenTableDao.get_gen_table_by_name(query_db, table_name) + source_name = cls._source_name(source_name) + gen_table = await GenTableDao.get_gen_table_by_name(query_db, table_name, source_name) table = GenTableModel(**CamelCaseUtil.transform_result(gen_table)) table_columns = table.columns table_column_map = {column.column_name: column for column in table_columns} - query_db_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name(query_db, table_name) - db_table_columns = [ - GenTableColumnModel(**column) for column in CamelCaseUtil.transform_result(query_db_table_columns) - ] + source_config = DataBaseConfig.get_source(source_name) + async with DataSourceRegistry.session(source_name) as target_db: + query_db_table_columns = await GenTableColumnDao.get_gen_db_table_columns_by_name( + target_db, table_name, source_config + ) + db_table_columns = [ + GenTableColumnModel(**column) for column in CamelCaseUtil.transform_result(query_db_table_columns) + ] if not db_table_columns: raise ServiceException('同步数据失败,原表结构不存在') db_table_column_names = [column.column_name for column in db_table_columns] @@ -416,7 +498,11 @@ async def set_sub_table(cls, query_db: AsyncSession, gen_table: GenTableModel) - :return: """ if gen_table.sub_table_name: - sub_table = await GenTableDao.get_gen_table_by_name(query_db, gen_table.sub_table_name) + sub_table = await GenTableDao.get_gen_table_by_name( + query_db, + gen_table.sub_table_name, + cls._source_name(gen_table.data_source_name), + ) gen_table.sub_table = GenTableModel(**CamelCaseUtil.transform_result(sub_table)) @classmethod diff --git a/ruoyi-fastapi-backend/module_generator/templates/python/controller.py.jinja2 b/ruoyi-fastapi-backend/module_generator/templates/python/controller.py.jinja2 index b4340ea..a61fa47 100644 --- a/ruoyi-fastapi-backend/module_generator/templates/python/controller.py.jinja2 +++ b/ruoyi-fastapi-backend/module_generator/templates/python/controller.py.jinja2 @@ -48,7 +48,7 @@ from utils.response_util import ResponseUtil async def get_{{ moduleName }}_{{ businessName }}_list( request: Request, {% if table.crud or table.sub %}{{ businessName }}_page_query{% elif table.tree %}{{ businessName }}_query{% endif %}: Annotated[{{ BusinessName }}PageQueryModel, Query()], - query_db: Annotated[AsyncSession, DBSessionDependency()], + query_db: Annotated[AsyncSession, {{ dbSessionDependency }}], ) -> Response: {% if table.crud or table.sub %} # 获取分页数据 @@ -76,7 +76,7 @@ async def get_{{ moduleName }}_{{ businessName }}_list( async def add_{{ moduleName }}_{{ businessName }}( request: Request, add_{{ businessName }}: {{ BusinessName }}Model, - query_db: Annotated[AsyncSession, DBSessionDependency()], + query_db: Annotated[AsyncSession, {{ dbSessionDependency }}], current_user: Annotated[CurrentUserModel, CurrentUserDependency()], ) -> Response: {% for column in columns %} @@ -108,7 +108,7 @@ async def add_{{ moduleName }}_{{ businessName }}( async def edit_{{ moduleName }}_{{ businessName }}( request: Request, edit_{{ businessName }}: {{ BusinessName }}Model, - query_db: Annotated[AsyncSession, DBSessionDependency()], + query_db: Annotated[AsyncSession, {{ dbSessionDependency }}], current_user: Annotated[CurrentUserModel, CurrentUserDependency()], ) -> Response: {% for column in columns %} @@ -135,7 +135,7 @@ async def edit_{{ moduleName }}_{{ businessName }}( async def delete_{{ moduleName }}_{{ businessName }}( request: Request, {{ pk_field }}s: Annotated[str, Path(description='需要删除的{{ pk_field_comment }}')], - query_db: Annotated[AsyncSession, DBSessionDependency()], + query_db: Annotated[AsyncSession, {{ dbSessionDependency }}], ) -> Response: delete_{{ businessName }} = Delete{{ BusinessName }}Model({{ pkField }}s={{ pk_field }}s) delete_{{ businessName }}_result = await {{ BusinessName }}Service.delete_{{ businessName }}_services(query_db, delete_{{ businessName }}) @@ -154,7 +154,7 @@ async def delete_{{ moduleName }}_{{ businessName }}( async def query_detail_{{ moduleName }}_{{ businessName }}( request: Request, {{ pk_field }}: Annotated[{{ pkColumn.python_type }}, Path(description='{{ pk_field_comment }}')], - query_db: Annotated[AsyncSession, DBSessionDependency()], + query_db: Annotated[AsyncSession, {{ dbSessionDependency }}], ) -> Response: {{ businessName }}_detail_result = await {{ BusinessName }}Service.{{ businessName }}_detail_services(query_db, {{ pk_field }}) logger.info(f'获取{{ pk_field }}为{% raw %}{{% endraw %}{{ pk_field }}{% raw %}}{% endraw %}的信息成功') @@ -181,7 +181,7 @@ async def query_detail_{{ moduleName }}_{{ businessName }}( async def export_{{ moduleName }}_{{ businessName }}_list( request: Request, {{ businessName }}_page_query: Annotated[{{ BusinessName }}PageQueryModel, Form()], - query_db: Annotated[AsyncSession, DBSessionDependency()], + query_db: Annotated[AsyncSession, {{ dbSessionDependency }}], ) -> Response: # 获取全量数据 {{ businessName }}_query_result = await {{ BusinessName }}Service.get_{{ businessName }}_list_services(query_db, {{ businessName }}_page_query, is_page=False) diff --git a/ruoyi-fastapi-backend/module_generator/templates/python/do.py.jinja2 b/ruoyi-fastapi-backend/module_generator/templates/python/do.py.jinja2 index 76137e3..28ccf24 100644 --- a/ruoyi-fastapi-backend/module_generator/templates/python/do.py.jinja2 +++ b/ruoyi-fastapi-backend/module_generator/templates/python/do.py.jinja2 @@ -5,10 +5,17 @@ from sqlalchemy.orm import relationship {% endif %} +{% if dataSourceName == defaultDataSourceName %} from config.database import Base +{% else %} +from config.database import get_data_source_base -class {{ ClassName }}(Base): +DataSourceBase = get_data_source_base('{{ dataSourceName }}') +{% endif %} + + +class {{ ClassName }}({% if dataSourceName == defaultDataSourceName %}Base{% else %}DataSourceBase{% endif %}): """ {{ functionName }}表 """ @@ -17,7 +24,7 @@ class {{ ClassName }}(Base): __table_args__ = {'comment': '{{ tableComment }}'} {% for column in columns %} - {{ column.column_name }} = Column({{ column.column_type | get_sqlalchemy_type }}, {% if column.pk %}primary_key=True, {% endif %}{% if column.increment %}autoincrement=True, {% endif %}{% if column.required or column.pk %}nullable=False{% else %}nullable=True{% endif %}, comment='{{ column.column_comment }}') + {{ column.column_name }} = Column({{ column.column_type | get_sqlalchemy_type(dataSourceName) }}, {% if column.pk %}primary_key=True, {% endif %}{% if column.increment %}autoincrement=True, {% endif %}{% if column.required or column.pk %}nullable=False{% else %}nullable=True{% endif %}, comment='{{ column.column_comment }}') {% endfor %} {% if table.sub %} @@ -26,7 +33,7 @@ class {{ ClassName }}(Base): {% if table.sub %} -class {{ subClassName }}(Base): +class {{ subClassName }}({% if dataSourceName == defaultDataSourceName %}Base{% else %}DataSourceBase{% endif %}): """ {{ subTable.function_name }}表 """ @@ -34,7 +41,7 @@ class {{ subClassName }}(Base): __tablename__ = '{{ subTableName }}' {% for column in subTable.columns %} - {{ column.column_name }} = Column({{ column.column_type | get_sqlalchemy_type }}, {% if column.column_name == subTableFkName %}ForeignKey('{{ tableName }}.{{ subTableFkName }}'), {% endif %}{% if column.pk %}primary_key=True, {% endif %}{% if column.increment %}autoincrement=True, {% endif %}{% if column.required %}nullable=True{% else %}nullable=False{% endif %}, comment='{{ column.column_comment }}') + {{ column.column_name }} = Column({{ column.column_type | get_sqlalchemy_type(dataSourceName) }}, {% if column.column_name == subTableFkName %}ForeignKey('{{ tableName }}.{{ subTableFkName }}'), {% endif %}{% if column.pk %}primary_key=True, {% endif %}{% if column.increment %}autoincrement=True, {% endif %}{% if column.required %}nullable=True{% else %}nullable=False{% endif %}, comment='{{ column.column_comment }}') {% endfor %} {% if table.sub %} diff --git a/ruoyi-fastapi-backend/module_task/file_task.py b/ruoyi-fastapi-backend/module_task/file_task.py index 99a217a..a97bac1 100644 --- a/ruoyi-fastapi-backend/module_task/file_task.py +++ b/ruoyi-fastapi-backend/module_task/file_task.py @@ -1,4 +1,4 @@ -from config.database import AsyncSessionLocal +from config.database import DataSourceRegistry from module_admin.service.file_business_service import FileRetentionNoticeService from module_admin.service.file_service import FileLifecycleService, FileReconcileService from utils.log_util import logger @@ -16,7 +16,7 @@ async def scan_retention_reminders(remind_days: int = 7, batch_size: int = 500) """ expiring_count = 0 expired_count = 0 - async with AsyncSessionLocal() as query_db: + async with DataSourceRegistry.session() as query_db: for _ in range(MAX_TASK_BATCHES): scan_result = await FileRetentionNoticeService.scan_file_retention_notices_services( query_db, @@ -41,7 +41,7 @@ async def purge_recycle_bin(retention_days: int = 30, batch_size: int = 100) -> :return: None """ purge_count = 0 - async with AsyncSessionLocal() as query_db: + async with DataSourceRegistry.session() as query_db: for _ in range(MAX_TASK_BATCHES): current_count = await FileLifecycleService.purge_recycle_bin_services( query_db, diff --git a/ruoyi-fastapi-backend/plugins/ai/entity/do/ai_model_do.py b/ruoyi-fastapi-backend/plugins/ai/entity/do/ai_model_do.py index c13c279..6928d99 100644 --- a/ruoyi-fastapi-backend/plugins/ai/entity/do/ai_model_do.py +++ b/ruoyi-fastapi-backend/plugins/ai/entity/do/ai_model_do.py @@ -20,7 +20,7 @@ class AiModels(Base): model_name = Column( String(100), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='模型名称', ) provider = Column(String(50), nullable=False, comment='提供商') @@ -28,19 +28,19 @@ class AiModels(Base): api_key = Column( String(255), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='API Key', ) base_url = Column( String(255), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='Base URL', ) model_type = Column( String(50), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='模型类型', ) max_tokens = Column(Integer, nullable=True, comment='最大输出token') @@ -57,6 +57,6 @@ class AiModels(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/plugins/ai/utils/ai_util.py b/ruoyi-fastapi-backend/plugins/ai/utils/ai_util.py index 39dd781..6a07066 100644 --- a/ruoyi-fastapi-backend/plugins/ai/utils/ai_util.py +++ b/ruoyi-fastapi-backend/plugins/ai/utils/ai_util.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING from urllib.parse import urlparse -from config.database import async_engine +from config.database import DataSourceRegistry from config.env import DataBaseConfig if TYPE_CHECKING: @@ -111,11 +111,13 @@ def get_storage_engine(cls) -> 'AsyncBaseDb': :return: 存储引擎实例 """ - storage_engine_class = cls._resolve_storage_class(DataBaseConfig.db_type) + default_source = DataBaseConfig.db_default_source + source_config = DataBaseConfig.default_source + storage_engine_class = cls._resolve_storage_class(source_config.db_type) return storage_engine_class( - db_engine=async_engine, - db_schema=DataBaseConfig.db_database if DataBaseConfig.db_type == 'mysql' else 'public', + db_engine=DataSourceRegistry.get_async_engine(default_source), + db_schema=source_config.db_database if source_config.db_type == 'mysql' else 'public', session_table='ai_sessions', memory_table='ai_memories', metrics_table='ai_metrics', diff --git a/ruoyi-fastapi-backend/plugins/core/lifecycle/migration.py b/ruoyi-fastapi-backend/plugins/core/lifecycle/migration.py index 5ee70ee..32242ec 100644 --- a/ruoyi-fastapi-backend/plugins/core/lifecycle/migration.py +++ b/ruoyi-fastapi-backend/plugins/core/lifecycle/migration.py @@ -564,7 +564,7 @@ def _filter_current_database_migrations(cls, migration_paths: list[str]) -> list return PluginLifecycleScriptHelper.filter_current_database_paths( migration_paths, root_dir='migrations', - database_type=DataBaseConfig.db_type, + database_type=DataBaseConfig.default_source.db_type, ) def _load_migration_module(self, migration_file: Path) -> Any: diff --git a/ruoyi-fastapi-backend/plugins/core/lifecycle/seed.py b/ruoyi-fastapi-backend/plugins/core/lifecycle/seed.py index 5c93a28..9f09a26 100644 --- a/ruoyi-fastapi-backend/plugins/core/lifecycle/seed.py +++ b/ruoyi-fastapi-backend/plugins/core/lifecycle/seed.py @@ -138,7 +138,7 @@ def _filter_current_database_seeds(cls, seed_paths: list[str]) -> list[str]: return PluginLifecycleScriptHelper.filter_current_database_paths( seed_paths, root_dir='seeds', - database_type=DataBaseConfig.db_type, + database_type=DataBaseConfig.default_source.db_type, ) def _load_seed_module(self, seed_file: Path) -> Any: diff --git a/ruoyi-fastapi-backend/plugins/core/management/entity/do/models.py b/ruoyi-fastapi-backend/plugins/core/management/entity/do/models.py index 88b30c1..b0f2046 100644 --- a/ruoyi-fastapi-backend/plugins/core/management/entity/do/models.py +++ b/ruoyi-fastapi-backend/plugins/core/management/entity/do/models.py @@ -48,7 +48,7 @@ class SysPlugin(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) @@ -154,6 +154,6 @@ class SysPluginOperationLog(Base): remark = Column( String(500), nullable=True, - server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.db_type), + server_default=SqlalchemyUtil.get_server_default_null(DataBaseConfig.default_source.db_type), comment='备注', ) diff --git a/ruoyi-fastapi-backend/plugins/core/management/service/gateway.py b/ruoyi-fastapi-backend/plugins/core/management/service/gateway.py index a3d3028..c19f020 100644 --- a/ruoyi-fastapi-backend/plugins/core/management/service/gateway.py +++ b/ruoyi-fastapi-backend/plugins/core/management/service/gateway.py @@ -202,7 +202,7 @@ def get_async_session_local() -> AsyncSessionFactoryProtocol: :return: 异步数据库会话工厂 """ - return import_module('config.database').AsyncSessionLocal + return import_module('config.database').DataSourceRegistry.session @staticmethod def get_plugin_service() -> type[PluginManagementServiceProtocol]: diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/route_guard.py b/ruoyi-fastapi-backend/plugins/core/runtime/route_guard.py index ff9c707..74bcf2c 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/route_guard.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/route_guard.py @@ -3,7 +3,7 @@ from fastapi import Depends, params from sqlalchemy.ext.asyncio import AsyncSession -from config.get_db import get_db +from common.aspect.db_seesion import DBSessionDependency from exceptions.exception import PermissionException @@ -56,7 +56,7 @@ def __init__(self, plugin_id: str, state_gateway: PluginRouteStateGateway | None self.plugin_id = plugin_id self.state_gateway = state_gateway or UnavailablePluginRouteStateGateway() - async def __call__(self, db: AsyncSession = Depends(get_db)) -> bool: + async def __call__(self, db: AsyncSession = DBSessionDependency()) -> bool: """ 执行插件启用状态校验。 diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/startup.py b/ruoyi-fastapi-backend/plugins/core/runtime/startup.py index d84178a..48966ff 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/startup.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/startup.py @@ -8,9 +8,8 @@ from fastapi import FastAPI from common.router import auto_register_controller_files -from config.database import AsyncSessionLocal +from config.database import DataSourceRegistry from config.env import AppConfig, get_config -from config.get_db import get_db from plugins.core.discovery.registry import PluginRegistry, RegisteredPlugin from plugins.core.lifecycle.migration import ( PluginMigrationHistoryRecord, @@ -315,7 +314,7 @@ async def requires_startup_write(self) -> bool: if not discovered_plugin_ids: return False - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: plugin_list = await self.management_gateway.list_plugins(query_db) database_plugin_map = {plugin.plugin_id: plugin for plugin in plugin_list} return any( @@ -423,7 +422,7 @@ async def recover_plugin_dependency_errors( :return: None """ recovered = False - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: for plugin in plugins: result = await self.management_gateway.recover_plugin_dependency_error( query_db, @@ -589,7 +588,7 @@ async def load_registry_from_database(self, app: FastAPI) -> None: :param app: FastAPI对象 :return: None """ - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: plugin_list = await self.management_gateway.list_plugins(query_db) app.state.plugin_registry = self.builder.build_registry(plugin_list) @@ -662,7 +661,7 @@ async def install_plugin_resources_with_isolation( ): logger.info('🔄 开始同步单插件启动资源') try: - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: try: await self.management_gateway.install_plugin_resources( query_db, @@ -743,7 +742,7 @@ async def sync_plugin_install(self, discovered_plugin: Any, *, enabled: bool) -> ): logger.info('🔄 开始执行插件启动安装生命周期') self.validate_plugin_structure(discovered_plugin) - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: try: await self.management_gateway.upsert_discovered_plugin( query_db, @@ -820,10 +819,10 @@ async def run_plugin_install_scripts(self, query_db: Any, discovered_plugin: Any :param discovered_plugin: 已发现插件对象 :return: None """ - async with AsyncSessionLocal() as migration_session: + async with DataSourceRegistry.session() as migration_session: await PluginMigrationRunner( discovered_plugin, - PluginStartupMigrationHistoryStore(self.management_gateway, AsyncSessionLocal), + PluginStartupMigrationHistoryStore(self.management_gateway, DataSourceRegistry.session), manage_execution_transaction=True, ).run(migration_session) await self.run_plugin_seed_scripts(query_db, discovered_plugin) @@ -996,7 +995,7 @@ async def mark_plugin_runtime_error(self, app: FastAPI, plugin_id: str, error_me :param error_message: 错误信息 :return: None """ - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: result = await self.management_gateway.mark_plugin_error(query_db, plugin_id, error_message) if not result.is_success: await query_db.rollback() @@ -1037,7 +1036,7 @@ async def sync_default_enabled_builtin_plugin_install_states(self) -> set[str]: return set() failed_plugin_ids: set[str] = set() - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: plugin_list = await self.management_gateway.list_plugins(query_db) database_plugin_map = {plugin.plugin_id: plugin for plugin in plugin_list} plugins_to_sync = [ @@ -1083,7 +1082,7 @@ async def mark_discovered_plugin_startup_error( :return: None """ plugin_id = discovered_plugin.manifest.id - async for query_db in get_db(): + async with DataSourceRegistry.session() as query_db: try: await self.management_gateway.upsert_discovered_plugin( query_db, diff --git a/ruoyi-fastapi-backend/plugins/core/validation/manifest.py b/ruoyi-fastapi-backend/plugins/core/validation/manifest.py index 910e701..d1e660d 100644 --- a/ruoyi-fastapi-backend/plugins/core/validation/manifest.py +++ b/ruoyi-fastapi-backend/plugins/core/validation/manifest.py @@ -607,7 +607,7 @@ def _check_compatibility(self, manifest: PluginManifest) -> list[PluginValidatio ) ) - current_database = DataBaseConfig.db_type + current_database = DataBaseConfig.default_source.db_type if compatibility.databases and current_database not in compatibility.databases: issues.append( PluginValidationIssue( diff --git a/ruoyi-fastapi-backend/scripts/migrate_legacy_files.py b/ruoyi-fastapi-backend/scripts/migrate_legacy_files.py index 20dfa05..5c6600e 100644 --- a/ruoyi-fastapi-backend/scripts/migrate_legacy_files.py +++ b/ruoyi-fastapi-backend/scripts/migrate_legacy_files.py @@ -11,7 +11,7 @@ import aiofiles from pydantic import ValidationError -from config.database import AsyncSessionLocal +from config.database import DataSourceRegistry from config.env import UploadConfig from module_admin.dao.file_info_dao import FileInfoDao from module_admin.entity.vo.file_vo import FileInfoModel @@ -222,7 +222,7 @@ async def migrate_legacy_files( legacy_files, skipped_count = await asyncio.to_thread(collect_legacy_files) added_count = 0 pending_count = 0 - async with AsyncSessionLocal() as session: + async with DataSourceRegistry.session() as session: for legacy_file in legacy_files: if await FileInfoDao.get_file_info_by_storage_key(session, legacy_file.storage_key): skipped_count += 1 diff --git a/ruoyi-fastapi-backend/server.py b/ruoyi-fastapi-backend/server.py index 8829757..e17ac92 100644 --- a/ruoyi-fastapi-backend/server.py +++ b/ruoyi-fastapi-backend/server.py @@ -6,10 +6,11 @@ from common.constant import LockConstant from common.router import auto_register_routers +from config.database import DataSourceRegistry from config.env import AppConfig -from config.get_db import close_async_engine, init_create_table from config.get_redis import RedisUtil from config.get_scheduler import SchedulerUtil +from config.lifecycle import init_create_table from exceptions.handle import handle_exception from middlewares.handle import handle_middleware from module_admin.service.log_service import LogAggregatorService @@ -55,7 +56,7 @@ async def _stop_background_tasks(app: FastAPI) -> None: try: await RedisUtil.close_redis_pool(app) finally: - await close_async_engine() + await DataSourceRegistry.dispose_all() async def _shutdown_application_runtime(app: FastAPI) -> None: @@ -76,6 +77,40 @@ async def _shutdown_application_runtime(app: FastAPI) -> None: await logger.complete() +def _show_startup_addresses() -> None: + """ + 显示应用及接口文档访问地址 + + :return: None + """ + host = AppConfig.app_host + port = AppConfig.app_port + if host == '0.0.0.0': + local_ip = IPUtil.get_local_ip() + network_ips = IPUtil.get_network_ips() + else: + local_ip = host + network_ips = [host] + + app_links = [f'🏠 Local: http://{local_ip}:{port}'] + app_links.extend(f'📡 Network: http://{ip}:{port}' for ip in network_ips) + logger.opt(colors=True).info('💻 应用地址:\n' + '\n'.join(app_links)) + + if not AppConfig.app_disable_swagger: + swagger_links = [f'🏠 Local: http://{local_ip}:{port}{APIDocsUtil.docs_url()}'] + swagger_links.extend( + f'📡 Network: http://{ip}:{port}{APIDocsUtil.docs_url()}' for ip in network_ips + ) + logger.opt(colors=True).info('📄 Swagger文档:\n' + '\n'.join(swagger_links)) + + if not AppConfig.app_disable_redoc: + redoc_links = [f'🏠 Local: http://{local_ip}:{port}{APIDocsUtil.redoc_url()}'] + redoc_links.extend( + f'📡 Network: http://{ip}:{port}{APIDocsUtil.redoc_url()}' for ip in network_ips + ) + logger.opt(colors=True).info('📚 ReDoc文档:\n' + '\n'.join(redoc_links)) + + # 生命周期事件 @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: @@ -85,9 +120,10 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: :param app: FastAPI对象 :return: None """ - app.state.redis = await RedisUtil.create_redis_pool(log_enabled=False) + app.state.redis = None app.state.plugin_application_runtime_started = False try: + app.state.redis = await RedisUtil.create_redis_pool(log_enabled=False) application_lock_owner_token = SchedulerUtil.get_application_lock_owner_token() application_leader = await StartupUtil.acquire_application_leader( redis=app.state.redis, @@ -102,6 +138,8 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: if application_leader: SchedulerUtil.start_application_lock_renewal(app.state.redis) + await DataSourceRegistry.initialize() + startup_logger = logger.bind( startup_phase='application_startup', startup_role='application_leader', @@ -119,37 +157,16 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: # 短暂等待确保下面的启动日志在最后打印 await asyncio.sleep(1) startup_logger.info(f'🚀 {AppConfig.app_name}启动成功') - host = AppConfig.app_host - port = AppConfig.app_port - if host == '0.0.0.0': - local_ip = IPUtil.get_local_ip() - network_ips = IPUtil.get_network_ips() - else: - local_ip = host - network_ips = [host] - - app_links = [f'🏠 Local: http://{local_ip}:{port}'] - app_links.extend(f'📡 Network: http://{ip}:{port}' for ip in network_ips) - logger.opt(colors=True).info('💻 应用地址:\n' + '\n'.join(app_links)) - - if not AppConfig.app_disable_swagger: - swagger_links = [f'🏠 Local: http://{local_ip}:{port}{APIDocsUtil.docs_url()}'] - swagger_links.extend( - f'📡 Network: http://{ip}:{port}{APIDocsUtil.docs_url()}' for ip in network_ips - ) - logger.opt(colors=True).info('📄 Swagger文档:\n' + '\n'.join(swagger_links)) - - if not AppConfig.app_disable_redoc: - redoc_links = [f'🏠 Local: http://{local_ip}:{port}{APIDocsUtil.redoc_url()}'] - redoc_links.extend( - f'📡 Network: http://{ip}:{port}{APIDocsUtil.redoc_url()}' for ip in network_ips - ) - logger.opt(colors=True).info('📚 ReDoc文档:\n' + '\n'.join(redoc_links)) + _show_startup_addresses() # 确保启动阶段的插件摘要在ASGI lifespan启动完成前已写入stdout和日志文件。 await logger.complete() yield finally: - await _shutdown_application_runtime(app) + if app.state.redis is None: + await DataSourceRegistry.dispose_all() + await logger.complete() + else: + await _shutdown_application_runtime(app) async def _initialize_application_runtime(app: FastAPI, application_leader: bool) -> None: diff --git a/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql b/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql index 1e9e065..75db4e9 100644 --- a/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql +++ b/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql @@ -939,6 +939,7 @@ create table gen_table ( table_id bigserial not null, table_name varchar(200) default '', table_comment varchar(500) default '', + data_source_name varchar(64) not null default 'primary', sub_table_name varchar(64) default null, sub_table_fk_name varchar(64) default null, class_name varchar(100) default '', @@ -963,6 +964,7 @@ create table gen_table ( comment on column gen_table.table_id is '编号'; comment on column gen_table.table_name is '表名称'; comment on column gen_table.table_comment is '表描述'; +comment on column gen_table.data_source_name is '目标数据源名称'; comment on column gen_table.sub_table_name is '关联子表的表名'; comment on column gen_table.sub_table_fk_name is '子表关联的外键名'; comment on column gen_table.class_name is '实体类名称'; diff --git a/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql b/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql index 8b502a8..f4bdaf8 100644 --- a/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql +++ b/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql @@ -730,6 +730,7 @@ create table gen_table ( table_id bigint(20) not null auto_increment comment '编号', table_name varchar(200) default '' comment '表名称', table_comment varchar(500) default '' comment '表描述', + data_source_name varchar(64) not null default 'primary' comment '目标数据源名称', sub_table_name varchar(64) default null comment '关联子表的表名', sub_table_fk_name varchar(64) default null comment '子表关联的外键名', class_name varchar(100) default '' comment '实体类名称', diff --git a/ruoyi-fastapi-backend/tests/cli/runtime/test_app_runtime.py b/ruoyi-fastapi-backend/tests/cli/runtime/test_app_runtime.py index d470776..02d55bd 100644 --- a/ruoyi-fastapi-backend/tests/cli/runtime/test_app_runtime.py +++ b/ruoyi-fastapi-backend/tests/cli/runtime/test_app_runtime.py @@ -92,10 +92,12 @@ def test_app_snapshot_support_builds_config_snapshot() -> None: app_disable_redoc=False, ), DataBaseConfig=SimpleNamespace( - db_type='mysql', - db_host='127.0.0.1', - db_port=3306, - db_database='ruoyi', + default_source=SimpleNamespace( + db_type='mysql', + db_host='127.0.0.1', + db_port=3306, + db_database='ruoyi', + ) ), RedisConfig=SimpleNamespace(redis_host='127.0.0.1', redis_port=REDIS_PORT), LogConfig=SimpleNamespace(loguru_level='INFO'), diff --git a/ruoyi-fastapi-backend/tests/cli/tui/adapters/conftest.py b/ruoyi-fastapi-backend/tests/cli/tui/adapters/conftest.py index e86d268..813cbbf 100644 --- a/ruoyi-fastapi-backend/tests/cli/tui/adapters/conftest.py +++ b/ruoyi-fastapi-backend/tests/cli/tui/adapters/conftest.py @@ -209,14 +209,13 @@ def dispose_async_db_engine_after_test() -> None: 在每个 TUI adapter 测试结束后尝试释放全局异步数据库连接池。 这些适配器测试会按需导入运行时模块;若其中某些路径触发真实数据库访问, - 模块级 `async_engine` 可能在测试进程结束前仍持有连接,从而在 GC 阶段产生 - SQLAlchemy 未归还连接告警。这里统一在测试后主动 `dispose()`,将清理职责收口 - 到测试夹具而非业务代码。 + 注册表中的引擎可能在测试进程结束前仍持有连接,从而在 GC 阶段产生 + SQLAlchemy 未归还连接告警。这里统一在测试后主动释放注册表资源,将清理职责 + 收口到测试夹具而非业务代码。 :return: None """ yield database_module = sys.modules.get('config.database') - async_engine = getattr(database_module, 'async_engine', None) if database_module is not None else None - if async_engine is not None: - asyncio.run(async_engine.dispose()) + if database_module is not None: + asyncio.run(database_module.DataSourceRegistry.dispose_all()) diff --git a/ruoyi-fastapi-backend/tests/config/test_database_exception_handlers.py b/ruoyi-fastapi-backend/tests/config/test_database_exception_handlers.py new file mode 100644 index 0000000..5b0f396 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_database_exception_handlers.py @@ -0,0 +1,37 @@ +import json + +import pytest +from fastapi import FastAPI, status + +from common.constant import HttpStatusConstant +from exceptions.exception import ( + DataSourceInitializationException, + DataSourceNotFoundException, + DataSourceUnavailableException, + ServiceException, +) +from exceptions.handle import handle_exception + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'exception', + [ + DataSourceUnavailableException('reporting'), + DataSourceInitializationException('reporting'), + DataSourceNotFoundException('reporting'), + ], +) +async def test_data_source_exceptions_use_service_exception_handler(exception: Exception) -> None: + app = FastAPI() + handle_exception(app) + + handler = app.exception_handlers[ServiceException] + response = await handler(None, exception) + payload = json.loads(response.body) + + assert response.status_code == status.HTTP_200_OK + assert payload['code'] == HttpStatusConstant.ERROR + assert payload['success'] is False + assert 'reporting' in payload['msg'] + assert '://' not in payload['msg'] diff --git a/ruoyi-fastapi-backend/tests/config/test_database_registry.py b/ruoyi-fastapi-backend/tests/config/test_database_registry.py new file mode 100644 index 0000000..64ac107 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_database_registry.py @@ -0,0 +1,301 @@ +from __future__ import annotations + +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import TYPE_CHECKING +from unittest.mock import AsyncMock, MagicMock + +import pytest +from sqlalchemy import URL +from sqlalchemy.exc import OperationalError + +from common.aspect.db_seesion import DBSessionDependency, get_db_session_provider +from config import database +from exceptions.exception import DataSourceInitializationException, DataSourceUnavailableException + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + + +def _source(*, required: bool = True) -> SimpleNamespace: + return SimpleNamespace( + db_type='mysql', + db_host='db.example', + db_port=3306, + db_username='user', + db_password='p@ss/w0rd', + db_database='application', + db_echo=False, + db_connect_timeout=7, + db_max_overflow=1, + db_pool_size=1, + db_pool_recycle=60, + db_pool_timeout=2, + db_required=required, + ) + + +def test_database_urls_use_structured_url_and_hide_password() -> None: + config = _source() + async_url = database.build_async_sqlalchemy_database_url(config) + sync_url = database.build_sync_sqlalchemy_database_url(config) + + assert isinstance(async_url, URL) + assert async_url.drivername == 'mysql+asyncmy' + assert sync_url.drivername == 'mysql+pymysql' + assert async_url.password == 'p@ss/w0rd' + assert 'p@ss' not in repr(async_url) + + +@pytest.mark.parametrize( + ('db_type', 'db_port', 'async_timeout_key', 'sync_timeout_key'), + [ + ('mysql', 3306, 'connect_timeout', 'connect_timeout'), + ('postgresql', 5432, 'timeout', 'connect_timeout'), + ], +) +def test_engine_factories_use_driver_specific_connect_timeout( + monkeypatch: pytest.MonkeyPatch, + db_type: str, + db_port: int, + async_timeout_key: str, + sync_timeout_key: str, +) -> None: + config = _source() + config.db_type = db_type + config.db_port = db_port + async_engine = object() + sync_engine = object() + captured_options: dict[str, dict[str, object]] = {} + + def create_async_engine(_url: URL, **options: object) -> object: + captured_options['async'] = options + return async_engine + + def create_sync_engine(_url: URL, **options: object) -> object: + captured_options['sync'] = options + return sync_engine + + monkeypatch.setattr(database, 'create_async_engine', create_async_engine) + monkeypatch.setattr(database, 'create_engine', create_sync_engine) + + assert database.create_async_db_engine(config=config) is async_engine + assert database.create_sync_db_engine(config=config) is sync_engine + assert captured_options['async']['connect_args'] == {async_timeout_key: 7} + assert captured_options['sync']['connect_args'] == {sync_timeout_key: 7} + assert captured_options['async']['pool_use_lifo'] is True + assert captured_options['sync']['pool_use_lifo'] is True + + +def test_async_session_factory_disables_expiration_after_commit() -> None: + factory = database.create_async_session_factory(MagicMock()) + + assert factory.kw['expire_on_commit'] is False + + +class _Begin: + def __init__(self, should_fail: bool = False) -> None: + self.should_fail = should_fail + + async def __aenter__(self) -> _Begin: + if self.should_fail: + raise RuntimeError('password=secret') + return self + + async def __aexit__(self, *_args: object) -> bool: + return False + + async def execute(self, _statement: object) -> None: + return None + + +class _Engine: + def __init__(self) -> None: + self.should_fail = True + self.dispose = AsyncMock() + + def begin(self) -> _Begin: + return _Begin(self.should_fail) + + +@pytest.mark.asyncio +async def test_optional_source_recovers_after_cooldown(monkeypatch: pytest.MonkeyPatch) -> None: + primary_engine = _Engine() + primary_engine.should_fail = False + optional_engine = _Engine() + monkeypatch.setattr( + database, + 'create_async_db_engine', + lambda config: primary_engine if config.db_required else optional_engine, + ) + settings = SimpleNamespace( + db_default_source='primary', + db_sources={'primary': _source(), 'reporting': _source(required=False)}, + ) + registry = database._DataSourceRegistry(settings) + + await registry.initialize() + assert not registry._runtimes['reporting'].available + optional_engine.should_fail = False + runtime = registry._runtimes['reporting'] + runtime.next_retry_at = None + async with registry.connection('reporting'): + pass + assert runtime.available + await registry.dispose_all() + + +@pytest.mark.asyncio +async def test_optional_engine_creation_failure_degrades_and_recovers(monkeypatch: pytest.MonkeyPatch) -> None: + primary_config = _source() + optional_config = _source(required=False) + primary_engine = _Engine() + primary_engine.should_fail = False + optional_engine = _Engine() + optional_engine.should_fail = False + optional_creation_fails = True + + def create_engine(config: SimpleNamespace) -> _Engine: + nonlocal optional_creation_fails + if config is optional_config and optional_creation_fails: + raise ModuleNotFoundError('driver is not installed') + return optional_engine if config is optional_config else primary_engine + + monkeypatch.setattr(database, 'create_async_db_engine', create_engine) + registry = database._DataSourceRegistry( + SimpleNamespace( + db_default_source='primary', + db_sources={'primary': primary_config, 'reporting': optional_config}, + ) + ) + + await registry.initialize() + runtime = registry._runtimes['reporting'] + assert runtime.async_engine is None + assert not runtime.available + + optional_creation_fails = False + runtime.next_retry_at = None + async with registry.connection('reporting'): + pass + + assert runtime.async_engine is optional_engine + assert runtime.available + await registry.dispose_all() + + +@pytest.mark.asyncio +async def test_required_engine_creation_failure_disposes_other_sources(monkeypatch: pytest.MonkeyPatch) -> None: + primary_config = _source() + optional_config = _source(required=False) + optional_engine = _Engine() + optional_engine.should_fail = False + + def create_engine(config: SimpleNamespace) -> _Engine: + if config is primary_config: + raise ModuleNotFoundError('required driver is not installed') + return optional_engine + + monkeypatch.setattr(database, 'create_async_db_engine', create_engine) + registry = database._DataSourceRegistry( + SimpleNamespace( + db_default_source='primary', + db_sources={'primary': primary_config, 'reporting': optional_config}, + ) + ) + + with pytest.raises(DataSourceInitializationException): + await registry.initialize() + + optional_engine.dispose.assert_awaited_once_with() + assert registry._runtimes == {} + + +@pytest.mark.asyncio +async def test_runtime_operational_error_marks_source_down_then_recovers(monkeypatch: pytest.MonkeyPatch) -> None: + engine = _Engine() + engine.should_fail = False + monkeypatch.setattr(database, 'create_async_db_engine', lambda config: engine) + registry = database._DataSourceRegistry( + SimpleNamespace(db_default_source='primary', db_sources={'primary': _source()}) + ) + await registry.initialize() + runtime = registry._runtimes['primary'] + session_closed = False + + @asynccontextmanager + async def session_factory() -> AsyncGenerator[object, None]: + nonlocal session_closed + try: + yield object() + finally: + session_closed = True + + runtime.async_session_factory = session_factory + with pytest.raises(DataSourceUnavailableException): + async with registry.session(): + raise OperationalError('SELECT 1', {}, RuntimeError('connection lost')) + + assert session_closed + assert not runtime.available + with pytest.raises(DataSourceUnavailableException): + async with registry.session(): + pass + + runtime.next_retry_at = None + async with registry.session() as session: + assert session is not None + assert runtime.available + await registry.dispose_all() + + +@pytest.mark.asyncio +async def test_multiple_sources_have_independent_engines_and_factories(monkeypatch: pytest.MonkeyPatch) -> None: + primary_config = _source() + reporting_config = _source(required=False) + engines = {id(primary_config): _Engine(), id(reporting_config): _Engine()} + for engine in engines.values(): + engine.should_fail = False + monkeypatch.setattr(database, 'create_async_db_engine', lambda config: engines[id(config)]) + registry = database._DataSourceRegistry( + SimpleNamespace( + db_default_source='primary', + db_sources={'primary': primary_config, 'reporting': reporting_config}, + ) + ) + + await registry.initialize() + + assert registry.get_async_engine() is engines[id(primary_config)] + assert registry.get_async_engine('reporting') is engines[id(reporting_config)] + assert ( + registry._runtimes['primary'].async_session_factory is not registry._runtimes['reporting'].async_session_factory + ) + await registry.dispose_all() + + +@pytest.mark.asyncio +async def test_dispose_all_releases_sync_and_async_engines(monkeypatch: pytest.MonkeyPatch) -> None: + engine = _Engine() + engine.should_fail = False + sync_engine = MagicMock() + monkeypatch.setattr(database, 'create_async_db_engine', lambda config: engine) + monkeypatch.setattr(database, 'create_sync_db_engine', lambda config: sync_engine) + registry = database._DataSourceRegistry( + SimpleNamespace(db_default_source='primary', db_sources={'primary': _source()}) + ) + await registry.initialize() + assert registry.get_sync_engine() is sync_engine + + await registry.dispose_all() + + sync_engine.dispose.assert_called_once_with() + engine.dispose.assert_awaited_once_with() + assert registry._runtimes == {} + + +def test_dependency_provider_is_cached_per_source() -> None: + get_db_session_provider.cache_clear() + assert get_db_session_provider('reporting') is get_db_session_provider('reporting') + assert get_db_session_provider('reporting') is not get_db_session_provider('archive') + assert DBSessionDependency('reporting').dependency is get_db_session_provider('reporting') diff --git a/ruoyi-fastapi-backend/tests/config/test_database_settings.py b/ruoyi-fastapi-backend/tests/config/test_database_settings.py new file mode 100644 index 0000000..5fc7f48 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_database_settings.py @@ -0,0 +1,124 @@ +import json + +import pytest +from pydantic import SecretStr, ValidationError + +from config.env import DataBaseSettings, DataSourceNotFoundException, DataSourceSettings + +DEFAULT_CONNECT_TIMEOUT = 10 + + +def _source(**overrides: object) -> dict[str, object]: + values: dict[str, object] = { + 'db_type': 'mysql', + 'db_host': 'db.example.test', + 'db_port': 3306, + 'db_username': 'app', + 'db_password': 'super-secret', + 'db_database': 'appdb', + } + values.update(overrides) + return values + + +def test_db_sources_json_is_parsed_and_password_is_secret(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv('DB_DEFAULT_SOURCE', 'primary') + monkeypatch.setenv('DB_SOURCES', json.dumps({'primary': _source()})) + + settings = DataBaseSettings(_env_file=None) + source = settings.default_source + + assert settings.get_source() is source + assert isinstance(source.db_password, SecretStr) + assert source.db_password.get_secret_value() == 'super-secret' + assert source.db_connect_timeout == DEFAULT_CONNECT_TIMEOUT + assert 'super-secret' not in repr(source) + assert source.sqlglot_parse_dialect == 'mysql' + + +def test_postgresql_uses_sqlglot_postgres_dialect() -> None: + source = DataSourceSettings(**_source(db_type='postgresql', db_port=5432)) + + assert source.sqlglot_parse_dialect == 'postgres' + + +def test_data_source_allows_disabling_pool_recycle() -> None: + source = DataSourceSettings(**_source(db_pool_recycle=-1)) + + assert source.db_pool_recycle == -1 + + +@pytest.mark.parametrize('db_connect_timeout', [0, -1]) +def test_data_source_rejects_invalid_connect_timeout(db_connect_timeout: int) -> None: + with pytest.raises(ValidationError): + DataSourceSettings(**_source(db_connect_timeout=db_connect_timeout)) + + +@pytest.mark.parametrize( + ('field', 'value', 'message'), + [ + ('DB_SOURCES', '{}', 'DB_SOURCES 不能为空'), + ('DB_DEFAULT_SOURCE', 'missing', '默认数据源不存在'), + ], +) +def test_database_settings_rejects_empty_or_unknown_default( + monkeypatch: pytest.MonkeyPatch, + field: str, + value: str, + message: str, +) -> None: + monkeypatch.setenv('DB_SOURCES', json.dumps({'primary': _source()})) + monkeypatch.setenv('DB_DEFAULT_SOURCE', 'primary') + monkeypatch.setenv(field, value) + + with pytest.raises(ValidationError, match=message): + DataBaseSettings(_env_file=None) + + +def test_database_settings_rejects_invalid_source_name() -> None: + with pytest.raises(ValidationError, match='数据源名称不合法'): + DataBaseSettings( + _env_file=None, + db_default_source='Primary', + db_sources={'Primary': DataSourceSettings(**_source())}, + ) + + +def test_database_settings_does_not_expose_legacy_flat_fields() -> None: + settings = DataBaseSettings( + _env_file=None, + db_sources={'primary': DataSourceSettings(**_source())}, + ) + + assert not hasattr(settings, 'db_type') + assert not hasattr(settings, 'db_password') + + +def test_get_source_rejects_unconfigured_name() -> None: + settings = DataBaseSettings( + _env_file=None, + db_sources={'primary': DataSourceSettings(**_source())}, + ) + + with pytest.raises(DataSourceNotFoundException): + settings.get_source('reporting') + + +def test_malformed_sources_json_does_not_echo_password(monkeypatch: pytest.MonkeyPatch) -> None: + secret = 'must-not-leak' + monkeypatch.setenv('DB_SOURCES', f'{{"primary":{{"db_password":"{secret}"') + + with pytest.raises(ValueError, match='DB_SOURCES JSON 格式错误') as exc_info: + DataBaseSettings(_env_file=None) + + assert secret not in str(exc_info.value) + + +def test_invalid_source_fields_do_not_echo_password(monkeypatch: pytest.MonkeyPatch) -> None: + secret = 'valid-json-secret' + monkeypatch.setenv('DB_SOURCES', json.dumps({'primary': {'db_type': 'mysql', 'db_password': secret}})) + + with pytest.raises(ValidationError) as exc_info: + DataBaseSettings(_env_file=None) + + assert secret not in str(exc_info.value) diff --git a/ruoyi-fastapi-backend/tests/config/test_lifecycle.py b/ruoyi-fastapi-backend/tests/config/test_lifecycle.py new file mode 100644 index 0000000..8427ace --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_lifecycle.py @@ -0,0 +1,25 @@ +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from config.database import Base +from config.lifecycle import init_create_table + + +@pytest.mark.asyncio +async def test_init_create_table_uses_registry_connection(monkeypatch: pytest.MonkeyPatch) -> None: + connection = SimpleNamespace(run_sync=AsyncMock()) + + @asynccontextmanager + async def connection_context() -> AsyncGenerator[object, None]: + yield connection + + registry = SimpleNamespace(connection=connection_context) + monkeypatch.setattr('config.lifecycle.DataSourceRegistry', registry) + + await init_create_table(log_success_enabled=False) + + connection.run_sync.assert_awaited_once_with(Base.metadata.create_all) diff --git a/ruoyi-fastapi-backend/tests/config/test_multi_datasource_migration.py b/ruoyi-fastapi-backend/tests/config/test_multi_datasource_migration.py new file mode 100644 index 0000000..ecbd447 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_multi_datasource_migration.py @@ -0,0 +1,41 @@ +import importlib.util +from pathlib import Path +from types import ModuleType + +from alembic.migration import MigrationContext +from alembic.operations import Operations +from sqlalchemy import create_engine, inspect, text + +MIGRATION_PATH = ( + Path(__file__).resolve().parents[2] / 'alembic' / 'versions' / '2026_08_19_0001_add_generator_data_source.py' +) + + +def _load_migration() -> ModuleType: + spec = importlib.util.spec_from_file_location('multi_datasource_migration', MIGRATION_PATH) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_generator_data_source_migration_is_idempotent() -> None: + module = _load_migration() + engine = create_engine('sqlite:///:memory:') + with engine.begin() as connection: + connection.execute(text('create table gen_table (table_id integer primary key, table_name varchar(200))')) + module.op = Operations(MigrationContext.configure(connection)) + + module.upgrade() + module.upgrade() + + inspector = inspect(connection) + assert 'data_source_name' in {column['name'] for column in inspector.get_columns('gen_table')} + + +def test_generator_data_source_migration_allows_empty_database() -> None: + module = _load_migration() + engine = create_engine('sqlite:///:memory:') + with engine.begin() as connection: + module.op = Operations(MigrationContext.configure(connection)) + module.upgrade() diff --git a/ruoyi-fastapi-backend/tests/config/test_scheduler_database.py b/ruoyi-fastapi-backend/tests/config/test_scheduler_database.py new file mode 100644 index 0000000..971307e --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_scheduler_database.py @@ -0,0 +1,40 @@ +from unittest.mock import MagicMock + +import pytest + +from config.env import DataBaseConfig +from config.get_scheduler import SchedulerUtil + + +def test_scheduler_database_engines_keep_jobstore_logging_isolated(monkeypatch: pytest.MonkeyPatch) -> None: + jobstore_engine = MagicMock() + listener_engine = MagicMock() + create_engine = MagicMock(return_value=jobstore_engine) + + monkeypatch.setattr(SchedulerUtil, '_jobstore_engine', None) + monkeypatch.setattr(SchedulerUtil, '_listener_engine', None) + monkeypatch.setattr(SchedulerUtil, '_session_local', None) + monkeypatch.setattr(SchedulerUtil, '_disposed_sync_engines', False) + monkeypatch.setattr('config.get_scheduler.create_sync_db_engine', create_engine) + monkeypatch.setattr('config.get_scheduler.DataSourceRegistry.get_sync_engine', lambda _name: listener_engine) + + assert SchedulerUtil._get_jobstore_engine() is jobstore_engine + assert SchedulerUtil._get_listener_engine() is listener_engine + assert jobstore_engine is not listener_engine + create_engine.assert_called_once_with(echo=False, config=DataBaseConfig.get_source()) + + +def test_scheduler_cleanup_does_not_dispose_registry_listener_engine(monkeypatch: pytest.MonkeyPatch) -> None: + jobstore_engine = MagicMock() + listener_engine = MagicMock() + monkeypatch.setattr(SchedulerUtil, '_jobstore_engine', jobstore_engine) + monkeypatch.setattr(SchedulerUtil, '_listener_engine', listener_engine) + monkeypatch.setattr(SchedulerUtil, '_session_local', MagicMock()) + monkeypatch.setattr(SchedulerUtil, '_disposed_sync_engines', False) + + SchedulerUtil._dispose_sync_engines() + + jobstore_engine.dispose.assert_called_once_with() + listener_engine.dispose.assert_not_called() + assert SchedulerUtil._jobstore_engine is None + assert SchedulerUtil._listener_engine is None diff --git a/ruoyi-fastapi-backend/tests/config/test_scheduler_leader_lease.py b/ruoyi-fastapi-backend/tests/config/test_scheduler_leader_lease.py index e7d4314..d8291fc 100644 --- a/ruoyi-fastapi-backend/tests/config/test_scheduler_leader_lease.py +++ b/ruoyi-fastapi-backend/tests/config/test_scheduler_leader_lease.py @@ -178,11 +178,6 @@ async def test_scheduler_close_releases_owner_lease_before_forgetting_redis() -> 'stop_application_lock_renewal', new_callable=AsyncMock, ) as stop_renewal, - patch.object( - SchedulerUtil, - '_dispose_sync_async_engine', - new_callable=AsyncMock, - ), patch.object(SchedulerUtil, '_dispose_sync_engines'), patch('config.get_scheduler.scheduler', running=False), patch( diff --git a/ruoyi-fastapi-backend/tests/module_admin/service/test_file_access_log.py b/ruoyi-fastapi-backend/tests/module_admin/service/test_file_access_log.py index e7ea796..e30f064 100644 --- a/ruoyi-fastapi-backend/tests/module_admin/service/test_file_access_log.py +++ b/ruoyi-fastapi-backend/tests/module_admin/service/test_file_access_log.py @@ -31,6 +31,11 @@ async def __aexit__(self, exc_type: type | None, exc_value: BaseException | None return None +def make_database_registry(session: SimpleNamespace) -> SimpleNamespace: + """构造使用测试会话的数据库注册表。""" + return SimpleNamespace(session=lambda: AsyncSessionContext(session)) + + def make_file_access_log() -> FileAccessLogModel: return FileAccessLogModel( fileId='file-id', @@ -94,8 +99,8 @@ def test_file_access_log_event_is_persisted_and_acknowledged() -> None: with ( patch( - 'module_admin.service.log_service.AsyncSessionLocal', - return_value=AsyncSessionContext(session), + 'module_admin.service.log_service.DataSourceRegistry', + make_database_registry(session), ), patch.object(FileAccessLogDao, 'add_file_access_log_dao', new_callable=AsyncMock) as add_file_access_log, ): diff --git a/ruoyi-fastapi-backend/tests/module_generator/dao/test_gen_dao.py b/ruoyi-fastapi-backend/tests/module_generator/dao/test_gen_dao.py index 9a00938..f21c29b 100644 --- a/ruoyi-fastapi-backend/tests/module_generator/dao/test_gen_dao.py +++ b/ruoyi-fastapi-backend/tests/module_generator/dao/test_gen_dao.py @@ -1,4 +1,5 @@ -from unittest.mock import AsyncMock, patch +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -7,6 +8,32 @@ from utils.page_util import PageUtil +@pytest.mark.asyncio +async def test_get_gen_table_by_name_is_scoped_to_data_source() -> None: + result = MagicMock() + result.scalars.return_value.first.return_value = None + db = MagicMock() + db.execute = AsyncMock(return_value=result) + + await GenTableDao.get_gen_table_by_name(db, 'orders', 'reporting') + + statement = db.execute.await_args.args[0] + params = statement.compile().params + assert 'orders' in params.values() + assert 'reporting' in params.values() + + +@pytest.mark.asyncio +async def test_get_gen_table_list_filters_source_when_requested() -> None: + query_object = GenTablePageQueryModel(dataSourceName='reporting') + + with patch.object(PageUtil, 'paginate', new=AsyncMock(return_value=[])) as paginate: + await GenTableDao.get_gen_table_list(object(), query_object, is_page=True) + + query = paginate.await_args.args[1] + assert 'reporting' in query.compile().params.values() + + @pytest.mark.asyncio async def test_get_gen_db_table_list_does_not_bind_unreferenced_model_defaults() -> None: query_object = GenTablePageQueryModel(formColNum=3) @@ -41,3 +68,20 @@ async def test_get_gen_db_table_list_only_binds_active_sql_filters() -> None: 'end_time': '2026-07-30', } assert ')and ' not in compiled_sql + + +@pytest.mark.asyncio +async def test_get_gen_db_table_list_uses_target_source_config() -> None: + query_object = GenTablePageQueryModel() + + with patch.object(PageUtil, 'paginate', new=AsyncMock(return_value=[])) as paginate: + await GenTableDao.get_gen_db_table_list( + object(), + query_object, + excluded_table_names={'sys_user'}, + source_config=SimpleNamespace(db_type='postgresql'), + ) + + query = paginate.await_args.args[1] + assert 'list_table' in str(query.compile()) + assert query.compile().params['excluded_table_names'] == ('sys_user',) diff --git a/ruoyi-fastapi-backend/tests/module_generator/test_gen_service.py b/ruoyi-fastapi-backend/tests/module_generator/test_gen_service.py new file mode 100644 index 0000000..391c84d --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_generator/test_gen_service.py @@ -0,0 +1,57 @@ +from collections.abc import AsyncGenerator +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest +from sqlalchemy import create_engine, text + +from module_generator.service.gen_service import GenTableService + + +def test_get_data_source_list_services_returns_camel_case_response(monkeypatch: pytest.MonkeyPatch) -> None: + settings = SimpleNamespace( + db_default_source='primary', + db_sources={ + 'primary': SimpleNamespace(db_type='mysql'), + 'reporting': SimpleNamespace(db_type='postgresql'), + }, + ) + monkeypatch.setattr('module_generator.service.gen_service.DataBaseConfig', settings) + + result = GenTableService.get_data_source_list_services() + + assert [item.model_dump(by_alias=True) for item in result] == [ + {'name': 'primary', 'dbType': 'mysql', 'isDefault': True}, + {'name': 'reporting', 'dbType': 'postgresql', 'isDefault': False}, + ] + + +@pytest.mark.asyncio +async def test_get_gen_db_table_list_by_name_services_transforms_rows(monkeypatch: pytest.MonkeyPatch) -> None: + engine = create_engine('sqlite:///:memory:') + with engine.connect() as connection: + rows = connection.execute( + text("select 'sys_user' as table_name, '用户表' as table_comment, null as create_time, null as update_time") + ).fetchall() + + @asynccontextmanager + async def session(_source_name: str) -> AsyncGenerator[object, None]: + yield SimpleNamespace() + + settings = SimpleNamespace( + db_default_source='primary', + get_source=lambda _source_name: SimpleNamespace(db_type='mysql'), + ) + registry = SimpleNamespace(session=session) + monkeypatch.setattr('module_generator.service.gen_service.DataBaseConfig', settings) + monkeypatch.setattr('module_generator.service.gen_service.DataSourceRegistry', registry) + + with patch( + 'module_generator.service.gen_service.GenTableDao.get_gen_db_table_list_by_names', + new=AsyncMock(return_value=rows), + ): + result = await GenTableService.get_gen_db_table_list_by_name_services(object(), ['sys_user'], 'reporting') + + assert result[0].table_name == 'sys_user' + assert result[0].data_source_name == 'reporting' diff --git a/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_migration.py b/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_migration.py index 47da0b4..12c2419 100644 --- a/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_migration.py +++ b/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_migration.py @@ -215,7 +215,7 @@ async def test_plugin_migration_runner_filters_database_dialect_migrations( discovered_plugin = PluginScanner(tmp_path / 'plugins').load_manifest(plugin_root / 'plugin.yaml') engine = create_async_engine('sqlite+aiosqlite:///:memory:') session_maker = async_sessionmaker(bind=engine, expire_on_commit=False) - monkeypatch.setattr('plugins.core.lifecycle.migration.DataBaseConfig.db_type', 'mysql') + monkeypatch.setattr('plugins.core.lifecycle.migration.DataBaseConfig.default_source.db_type', 'mysql') try: async with session_maker() as session: diff --git a/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_seed.py b/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_seed.py index 8788326..961c088 100644 --- a/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_seed.py +++ b/ruoyi-fastapi-backend/tests/plugins/core/lifecycle/test_seed.py @@ -112,7 +112,7 @@ async def test_plugin_seed_runner_filters_seed_by_database_dialect( """, encoding='utf-8', ) - monkeypatch.setattr('plugins.core.lifecycle.seed.DataBaseConfig.db_type', 'mysql') + monkeypatch.setattr('plugins.core.lifecycle.seed.DataBaseConfig.default_source.db_type', 'mysql') discovered_plugin = PluginScanner(tmp_path / 'plugins').load_manifest(plugin_root / 'plugin.yaml') engine = create_async_engine('sqlite+aiosqlite:///:memory:') session_maker = async_sessionmaker(bind=engine, expire_on_commit=False) diff --git a/ruoyi-fastapi-backend/tests/plugins/core/runtime/test_startup.py b/ruoyi-fastapi-backend/tests/plugins/core/runtime/test_startup.py index df27c18..a4bb614 100644 --- a/ruoyi-fastapi-backend/tests/plugins/core/runtime/test_startup.py +++ b/ruoyi-fastapi-backend/tests/plugins/core/runtime/test_startup.py @@ -1,5 +1,8 @@ +from contextlib import asynccontextmanager from pathlib import Path from subprocess import CompletedProcess +from types import SimpleNamespace +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -12,9 +15,13 @@ BACKEND_ROOT = Path(__file__).resolve().parents[4] -def patch_startup_get_db(fake_get_db: object) -> object: - """patch 启动管理器方法实际使用的 get_db 全局引用。""" - return patch.dict(PluginRuntimeStartupManager.load_registry_from_database.__globals__, {'get_db': fake_get_db}) +def patch_startup_get_db(fake_get_db: Any) -> object: + """patch 启动管理器方法实际使用的默认数据源Session。""" + registry = SimpleNamespace(session=asynccontextmanager(fake_get_db)) + return patch.dict( + PluginRuntimeStartupManager.load_registry_from_database.__globals__, + {'DataSourceRegistry': registry}, + ) def patch_startup_global(name: str, value: object) -> object: @@ -445,10 +452,11 @@ async def test_run_plugin_install_scripts_runs_migrations_and_seeds() -> None: migration_session_context = MagicMock() migration_session_context.__aenter__ = AsyncMock(return_value=fake_migration_session) migration_session_context.__aexit__ = AsyncMock(return_value=None) - async_session_local = MagicMock(return_value=migration_session_context) + database_registry = MagicMock() + database_registry.session.return_value = migration_session_context with ( - patch_startup_global('AsyncSessionLocal', async_session_local), + patch_startup_global('DataSourceRegistry', database_registry), patch_startup_global('PluginMigrationRunner', runner_class), patch_startup_global('PluginSeedRunner', seed_runner_class), ): @@ -457,7 +465,7 @@ async def test_run_plugin_install_scripts_runs_migrations_and_seeds() -> None: runner_class.assert_called_once() assert runner_class.call_args.args[0] is discovered_plugin assert isinstance(runner_class.call_args.args[1], PluginStartupMigrationHistoryStore) - assert runner_class.call_args.args[1].async_session_local is async_session_local + assert runner_class.call_args.args[1].async_session_local is database_registry.session assert runner_class.call_args.kwargs['manage_execution_transaction'] is True migration_runner.run.assert_awaited_once_with(fake_migration_session) seed_runner_class.assert_called_once_with(discovered_plugin) diff --git a/ruoyi-fastapi-backend/tests/plugins/core/validation/test_manifest.py b/ruoyi-fastapi-backend/tests/plugins/core/validation/test_manifest.py index 2fa63c6..3052813 100644 --- a/ruoyi-fastapi-backend/tests/plugins/core/validation/test_manifest.py +++ b/ruoyi-fastapi-backend/tests/plugins/core/validation/test_manifest.py @@ -217,7 +217,7 @@ def test_manifest_checker_accepts_satisfied_compatibility(tmp_path: Path) -> Non def test_manifest_checker_accepts_supported_database(monkeypatch: object) -> None: """校验当前数据库在插件支持列表内时不产生问题。""" - monkeypatch.setattr(manifest_module.DataBaseConfig, 'db_type', 'postgresql') + monkeypatch.setattr(manifest_module.DataBaseConfig.default_source, 'db_type', 'postgresql') manifest = PluginManifest.model_validate( { 'id': 'demo', @@ -302,7 +302,7 @@ def test_manifest_checker_reports_unsatisfied_compatibility(tmp_path: Path) -> N def test_manifest_checker_reports_unsupported_database(monkeypatch: object) -> None: """校验当前数据库不在插件支持列表内时产生 error。""" - monkeypatch.setattr(manifest_module.DataBaseConfig, 'db_type', 'postgresql') + monkeypatch.setattr(manifest_module.DataBaseConfig.default_source, 'db_type', 'postgresql') manifest = PluginManifest.model_validate( { 'id': 'demo', diff --git a/ruoyi-fastapi-backend/tests/scripts/test_migrate_legacy_files.py b/ruoyi-fastapi-backend/tests/scripts/test_migrate_legacy_files.py index 6ef0b07..c26345d 100644 --- a/ruoyi-fastapi-backend/tests/scripts/test_migrate_legacy_files.py +++ b/ruoyi-fastapi-backend/tests/scripts/test_migrate_legacy_files.py @@ -31,6 +31,11 @@ async def __aexit__(self, exc_type: type | None, exc_value: BaseException | None return None +def make_database_registry(session: SimpleNamespace) -> SimpleNamespace: + """构造使用测试会话的数据库注册表。""" + return SimpleNamespace(session=lambda: AsyncSessionContext(session)) + + def test_collect_legacy_files_filters_disallowed_extensions(tmp_path: Path) -> None: allowed_file = tmp_path / 'upload' / 'report.txt' allowed_file.parent.mkdir() @@ -82,7 +87,7 @@ def test_migrate_legacy_files_dry_run_performs_full_validation_without_commit(tm with ( patch.object(UploadConfig, 'UPLOAD_PATH', str(tmp_path)), - patch('scripts.migrate_legacy_files.AsyncSessionLocal', return_value=AsyncSessionContext(session)), + patch('scripts.migrate_legacy_files.DataSourceRegistry', make_database_registry(session)), patch.object(FileInfoDao, 'get_file_info_by_storage_key', new=AsyncMock(return_value=None)) as get_file_info, patch.object(FileInfoDao, 'add_file_info_dao', new_callable=AsyncMock) as add_file_info, ): @@ -100,7 +105,7 @@ def test_migrate_legacy_files_writes_file_info_model_and_commits(tmp_path: Path) with ( patch.object(UploadConfig, 'UPLOAD_PATH', str(tmp_path)), - patch('scripts.migrate_legacy_files.AsyncSessionLocal', return_value=AsyncSessionContext(session)), + patch('scripts.migrate_legacy_files.DataSourceRegistry', make_database_registry(session)), patch.object(FileInfoDao, 'get_file_info_by_storage_key', new=AsyncMock(return_value=None)), patch.object(FileInfoDao, 'add_file_info_dao', new_callable=AsyncMock) as add_file_info, ): @@ -125,7 +130,7 @@ def test_migrate_legacy_files_skips_existing_storage_location(tmp_path: Path) -> with ( patch.object(UploadConfig, 'UPLOAD_PATH', str(tmp_path)), - patch('scripts.migrate_legacy_files.AsyncSessionLocal', return_value=AsyncSessionContext(session)), + patch('scripts.migrate_legacy_files.DataSourceRegistry', make_database_registry(session)), patch.object( FileInfoDao, 'get_file_info_by_storage_key', @@ -147,7 +152,7 @@ def test_migrate_legacy_files_commits_by_batch(tmp_path: Path) -> None: with ( patch.object(UploadConfig, 'UPLOAD_PATH', str(tmp_path)), - patch('scripts.migrate_legacy_files.AsyncSessionLocal', return_value=AsyncSessionContext(session)), + patch('scripts.migrate_legacy_files.DataSourceRegistry', make_database_registry(session)), patch.object(FileInfoDao, 'get_file_info_by_storage_key', new=AsyncMock(return_value=None)), patch.object(FileInfoDao, 'add_file_info_dao', new_callable=AsyncMock) as add_file_info, ): diff --git a/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py b/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py index 8d0519f..41cf7f3 100644 --- a/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py +++ b/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py @@ -126,11 +126,14 @@ async def test_lifespan_only_application_leader_outputs_banner_and_addresses() - """校验仅Application leader输出横幅、成功摘要和访问地址。""" app = SimpleNamespace(state=SimpleNamespace()) redis = MagicMock() + database_registry = MagicMock() + database_registry.initialize = AsyncMock() fake_logger = MagicMock() fake_logger.complete = AsyncMock() with ( patch('server.RedisUtil.create_redis_pool', new=AsyncMock(return_value=redis)), + patch('server.DataSourceRegistry', database_registry), patch('server.SchedulerUtil.get_application_lock_owner_token', return_value='owner-1'), patch('server.StartupUtil.acquire_application_leader', new=AsyncMock(return_value=True)), patch('server.SchedulerUtil.start_application_lock_renewal') as start_renewal, @@ -148,6 +151,7 @@ async def test_lifespan_only_application_leader_outputs_banner_and_addresses() - pass start_renewal.assert_called_once_with(redis) + database_registry.initialize.assert_awaited_once_with() initialize_runtime.assert_awaited_once_with(app, application_leader=True) worship.assert_called_once_with() fake_logger.bind.return_value.info.assert_has_calls( @@ -169,11 +173,14 @@ async def test_lifespan_non_leader_runs_local_initialization_without_display_log """校验非leader执行本地初始化,但不输出leader专属展示日志。""" app = SimpleNamespace(state=SimpleNamespace()) redis = MagicMock() + database_registry = MagicMock() + database_registry.initialize = AsyncMock() fake_logger = MagicMock() fake_logger.complete = AsyncMock() with ( patch('server.RedisUtil.create_redis_pool', new=AsyncMock(return_value=redis)), + patch('server.DataSourceRegistry', database_registry), patch('server.SchedulerUtil.get_application_lock_owner_token', return_value='owner-2'), patch('server.StartupUtil.acquire_application_leader', new=AsyncMock(return_value=False)), patch('server.SchedulerUtil.start_application_lock_renewal') as start_renewal, @@ -188,6 +195,7 @@ async def test_lifespan_non_leader_runs_local_initialization_without_display_log pass start_renewal.assert_not_called() + database_registry.initialize.assert_awaited_once_with() is_application_leader.assert_not_called() initialize_runtime.assert_awaited_once_with(app, application_leader=False) worship.assert_not_called() @@ -201,11 +209,14 @@ async def test_lifespan_non_leader_runs_local_initialization_without_display_log async def test_lifespan_non_leader_initialization_error_propagates_and_still_cleans_up() -> None: """校验非leader初始化异常不被门禁吞掉,并继续执行finally清理。""" app = SimpleNamespace(state=SimpleNamespace()) + database_registry = MagicMock() + database_registry.initialize = AsyncMock() fake_logger = MagicMock() fake_logger.complete = AsyncMock() with ( patch('server.RedisUtil.create_redis_pool', new=AsyncMock(return_value=MagicMock())), + patch('server.DataSourceRegistry', database_registry), patch('server.SchedulerUtil.get_application_lock_owner_token', return_value='owner-3'), patch('server.StartupUtil.acquire_application_leader', new=AsyncMock(return_value=False)), patch('server.TransportKeyProvider.validate_runtime_configuration'), @@ -221,11 +232,62 @@ async def test_lifespan_non_leader_initialization_error_propagates_and_still_cle pass initialize_runtime.assert_awaited_once_with(app, application_leader=False) + database_registry.initialize.assert_awaited_once_with() fake_logger.bind.return_value.info.assert_not_called() fake_logger.complete.assert_not_awaited() shutdown_runtime.assert_awaited_once_with(app) +@pytest.mark.asyncio +async def test_lifespan_skips_database_initialization_when_redis_creation_fails() -> None: + """校验Redis前置初始化失败时不创建数据库资源,并执行幂等清理。""" + app = SimpleNamespace(state=SimpleNamespace()) + database_registry = MagicMock() + database_registry.initialize = AsyncMock() + database_registry.dispose_all = AsyncMock() + fake_logger = MagicMock() + fake_logger.complete = AsyncMock() + + with ( + patch('server.DataSourceRegistry', database_registry), + patch( + 'server.RedisUtil.create_redis_pool', + new=AsyncMock(side_effect=RuntimeError('redis unavailable')), + ), + patch('server.logger', fake_logger), + pytest.raises(RuntimeError, match='redis unavailable'), + ): + async with lifespan(app): + pass + + database_registry.initialize.assert_not_awaited() + database_registry.dispose_all.assert_awaited_once_with() + fake_logger.complete.assert_awaited_once_with() + + +@pytest.mark.asyncio +async def test_lifespan_database_initialization_failure_releases_redis_and_database_resources() -> None: + """校验Redis和租约创建后数据库初始化失败仍走统一关闭流程。""" + app = SimpleNamespace(state=SimpleNamespace()) + redis = MagicMock() + database_registry = MagicMock() + database_registry.initialize = AsyncMock(side_effect=RuntimeError('database unavailable')) + + with ( + patch('server.RedisUtil.create_redis_pool', new=AsyncMock(return_value=redis)), + patch('server.DataSourceRegistry', database_registry), + patch('server.SchedulerUtil.get_application_lock_owner_token', return_value='owner-db-failure'), + patch('server.StartupUtil.acquire_application_leader', new=AsyncMock(return_value=False)), + patch('server._shutdown_application_runtime', new_callable=AsyncMock) as shutdown_runtime, + pytest.raises(RuntimeError, match='database unavailable'), + ): + async with lifespan(app): + pass + + database_registry.initialize.assert_awaited_once_with() + shutdown_runtime.assert_awaited_once_with(app) + + @pytest.mark.asyncio async def test_shutdown_application_runtime_preserves_cleanup_order() -> None: """校验插件关闭先执行,随后按Scheduler、Redis、数据库顺序释放资源。""" @@ -264,7 +326,7 @@ async def record_log_complete() -> None: 'server.RedisUtil.close_redis_pool', new=AsyncMock(side_effect=record_redis_shutdown), ), - patch('server.close_async_engine', new=AsyncMock(side_effect=record_database_shutdown)), + patch('server.DataSourceRegistry.dispose_all', new=AsyncMock(side_effect=record_database_shutdown)), patch('server.logger.complete', new=AsyncMock(side_effect=record_log_complete)), ): await _shutdown_application_runtime(app) @@ -302,10 +364,10 @@ async def test_stop_background_tasks_closes_redis_and_database_when_scheduler_cl new=AsyncMock(side_effect=RuntimeError('scheduler close failed')), ), patch('server.RedisUtil.close_redis_pool', new_callable=AsyncMock) as close_redis_pool, - patch('server.close_async_engine', new_callable=AsyncMock) as close_async_engine, + patch('server.DataSourceRegistry.dispose_all', new_callable=AsyncMock) as dispose_all, pytest.raises(RuntimeError, match='scheduler close failed'), ): await _stop_background_tasks(app) close_redis_pool.assert_awaited_once_with(app) - close_async_engine.assert_awaited_once_with() + dispose_all.assert_awaited_once_with() diff --git a/ruoyi-fastapi-backend/tests/utils/test_template_util.py b/ruoyi-fastapi-backend/tests/utils/test_template_util.py index cd8cb91..05caf31 100644 --- a/ruoyi-fastapi-backend/tests/utils/test_template_util.py +++ b/ruoyi-fastapi-backend/tests/utils/test_template_util.py @@ -1,7 +1,9 @@ import json +from types import SimpleNamespace import pytest +from config import database from module_generator.entity.vo.gen_vo import GenTableColumnModel, GenTableModel, GenTableParamsModel from module_generator.service.gen_service import GenTableService from utils.template_util import TemplateInitializer, TemplateUtils @@ -151,3 +153,61 @@ async def test_set_table_from_options_restores_view_flag() -> None: result = await GenTableService.set_table_from_options(gen_table) assert result.view is True + + +def test_secondary_source_templates_use_named_dependency_and_isolated_metadata( + monkeypatch: pytest.MonkeyPatch, +) -> None: + gen_table = _gen_table(gen_view=False) + gen_table.data_source_name = 'reporting' + monkeypatch.setattr( + type(database.DataBaseConfig), + 'get_source', + lambda _self, _name=None: SimpleNamespace(db_type='postgresql'), + ) + + context = TemplateUtils.prepare_context(gen_table) + env = TemplateInitializer.init_jinja2() + controller = env.get_template('python/controller.py.jinja2').render(**context) + entity = env.get_template('python/do.py.jinja2').render(**context) + + assert "DBSessionDependency('reporting')" in controller + assert "DataSourceBase = get_data_source_base('reporting')" in entity + assert 'class GenItem(DataSourceBase):' in entity + assert 'from config.database import Base' not in entity + compile(entity, '', 'exec') + + +def test_named_data_source_bases_are_cached_and_isolated(monkeypatch: pytest.MonkeyPatch) -> None: + database.get_data_source_base.cache_clear() + monkeypatch.setattr(type(database.DataBaseConfig), 'get_source', lambda _self, _name: object()) + + reporting_base = database.get_data_source_base('reporting') + archive_base = database.get_data_source_base('archive') + + assert reporting_base is database.get_data_source_base('reporting') + assert reporting_base is not archive_base + assert reporting_base.metadata is not archive_base.metadata + + +@pytest.mark.parametrize( + ('db_type', 'column_type', 'expected'), + [ + ('postgresql', 'JSONB', 'JSONB'), + ('postgresql', 'INET', 'INET'), + ('postgresql', 'ARRAY', 'ARRAY'), + ('postgresql', 'VARCHAR(64)', 'String(64)'), + ('mysql', 'DECIMAL(10, 2)', 'DECIMAL(10, 2)'), + ('mysql', 'LONGBLOB', 'LargeBinary'), + ], +) +def test_sqlalchemy_type_mapping_uses_target_source( + monkeypatch: pytest.MonkeyPatch, db_type: str, column_type: str, expected: str +) -> None: + monkeypatch.setattr( + type(database.DataBaseConfig), + 'get_source', + lambda _self, _name=None: SimpleNamespace(db_type=db_type), + ) + + assert TemplateUtils.get_sqlalchemy_type(column_type) == expected diff --git a/ruoyi-fastapi-backend/utils/gen_util.py b/ruoyi-fastapi-backend/utils/gen_util.py index dc8e881..33a4824 100644 --- a/ruoyi-fastapi-backend/utils/gen_util.py +++ b/ruoyi-fastapi-backend/utils/gen_util.py @@ -2,7 +2,7 @@ from datetime import datetime from common.constant import GenConstant -from config.env import GenConfig +from config.env import DataBaseConfig, GenConfig from module_generator.entity.vo.gen_vo import GenTableColumnModel, GenTableModel from utils.string_util import StringUtil @@ -41,33 +41,35 @@ def init_column_field(cls, column: GenTableColumnModel, table: GenTableModel) -> param table: 业务表对象 :return: """ - data_type = cls.get_db_type(column.column_type) + data_type = cls.get_db_type(column.column_type).lower() + source_config = DataBaseConfig.get_source(table.data_source_name) + db_type = source_config.db_type + python_mapping = GenConstant.DB_TO_PYTHON_TYPE_MAPPING[db_type] + string_types = GenConstant.COLUMNTYPE_STR[db_type] + text_types = GenConstant.COLUMNTYPE_TEXT[db_type] + time_types = GenConstant.COLUMNTYPE_TIME[db_type] + number_types = GenConstant.COLUMNTYPE_NUMBER[db_type] column_name = column.column_name column.table_id = table.table_id column.create_by = table.create_by # 设置Python字段名 column.python_field = cls.to_camel_case(column_name) # 设置默认类型 - column.python_type = StringUtil.get_mapping_value_by_key_ignore_case( - GenConstant.DB_TO_PYTHON_TYPE_MAPPING, data_type - ) + column.python_type = StringUtil.get_mapping_value_by_key_ignore_case(python_mapping, data_type) or 'str' column.query_type = GenConstant.QUERY_EQ - if cls.arrays_contains(GenConstant.COLUMNTYPE_STR, data_type) or cls.arrays_contains( - GenConstant.COLUMNTYPE_TEXT, data_type - ): + if data_type in string_types or data_type in text_types: # 字符串长度超过500设置为文本域 column_length = cls.get_column_length(column.column_type) html_type = ( GenConstant.HTML_TEXTAREA - if column_length >= cls.TEXTAREA_COLUMN_LENGTH - or cls.arrays_contains(GenConstant.COLUMNTYPE_TEXT, data_type) + if column_length >= cls.TEXTAREA_COLUMN_LENGTH or data_type in text_types else GenConstant.HTML_INPUT ) column.html_type = html_type - elif cls.arrays_contains(GenConstant.COLUMNTYPE_TIME, data_type): + elif data_type in time_types: column.html_type = GenConstant.HTML_DATETIME - elif cls.arrays_contains(GenConstant.COLUMNTYPE_NUMBER, data_type): + elif data_type in number_types: column.html_type = GenConstant.HTML_INPUT # 插入字段(默认所有字段都需要插入) diff --git a/ruoyi-fastapi-backend/utils/template_util.py b/ruoyi-fastapi-backend/utils/template_util.py index 12ad67d..21f5653 100644 --- a/ruoyi-fastapi-backend/utils/template_util.py +++ b/ruoyi-fastapi-backend/utils/template_util.py @@ -73,6 +73,9 @@ def prepare_context(cls, gen_table: GenTableModel) -> dict[str, Any]: tpl_category = gen_table.tpl_category function_name = gen_table.function_name + source_name = gen_table.data_source_name or DataBaseConfig.db_default_source + source_config = DataBaseConfig.get_source(source_name) + default_source_name = DataBaseConfig.db_default_source context = { 'tplCategory': tpl_category, 'tableName': gen_table.table_name, @@ -95,7 +98,14 @@ def prepare_context(cls, gen_table: GenTableModel) -> dict[str, Any]: 'columns': gen_table.columns, 'table': gen_table, 'dicts': cls.get_dicts(gen_table), - 'dbType': DataBaseConfig.db_type, + 'dbType': source_config.db_type, + 'dataSourceName': source_name, + 'defaultDataSourceName': default_source_name, + 'dbSessionDependency': ( + f"DBSessionDependency('{source_name}')" + if source_name != default_source_name + else 'DBSessionDependency()' + ), 'column_not_add_show': GenConstant.COLUMNNAME_NOT_ADD_SHOW, 'column_not_edit_show': GenConstant.COLUMNNAME_NOT_EDIT_SHOW, } @@ -288,20 +298,20 @@ def get_do_import_list(cls, gen_table: GenTableModel) -> list[str]: import_list = set() import_list.add('from sqlalchemy import Column') for column in columns: - data_type = cls.get_db_type(column.column_type) - if data_type in GenConstant.COLUMNTYPE_GEOMETRY: + sqlalchemy_type = cls.get_sqlalchemy_type(column.column_type, gen_table.data_source_name) + if sqlalchemy_type == 'Geometry': import_list.add('from geoalchemy2 import Geometry') - import_list.add( - f'from sqlalchemy import {StringUtil.get_mapping_value_by_key_ignore_case(GenConstant.DB_TO_SQLALCHEMY_TYPE_MAPPING, data_type)}' - ) + else: + import_list.add(f'from sqlalchemy import {sqlalchemy_type.split("(", 1)[0]}') if gen_table.sub: import_list.add('from sqlalchemy import ForeignKey') sub_columns = gen_table.sub_table.columns or [] for sub_column in sub_columns: - data_type = cls.get_db_type(sub_column.column_type) - import_list.add( - f'from sqlalchemy import {StringUtil.get_mapping_value_by_key_ignore_case(GenConstant.DB_TO_SQLALCHEMY_TYPE_MAPPING, data_type)}' - ) + sqlalchemy_type = cls.get_sqlalchemy_type(sub_column.column_type, gen_table.data_source_name) + if sqlalchemy_type == 'Geometry': + import_list.add('from geoalchemy2 import Geometry') + else: + import_list.add(f'from sqlalchemy import {sqlalchemy_type.split("(", 1)[0]}') return cls.merge_same_imports(list(import_list), 'from sqlalchemy import') @classmethod @@ -491,30 +501,19 @@ def to_camel_case(cls, text: str) -> str: return parts[0] + ''.join(word.capitalize() for word in parts[1:]) @classmethod - def get_sqlalchemy_type(cls, column_type: str) -> str: + def get_sqlalchemy_type(cls, column_type: str, source_name: str | None = None) -> str: """ 获取SQLAlchemy类型 :param column_type: 列类型 + :param source_name: 数据源名称 :return: SQLAlchemy类型 """ - if '(' in column_type: - column_type_list = column_type.split('(') - if column_type_list[0] in GenConstant.COLUMNTYPE_STR: - sqlalchemy_type = ( - StringUtil.get_mapping_value_by_key_ignore_case( - GenConstant.DB_TO_SQLALCHEMY_TYPE_MAPPING, column_type_list[0] - ) - + '(' - + column_type_list[1] - ) - else: - sqlalchemy_type = StringUtil.get_mapping_value_by_key_ignore_case( - GenConstant.DB_TO_SQLALCHEMY_TYPE_MAPPING, column_type_list[0] - ) - else: - sqlalchemy_type = StringUtil.get_mapping_value_by_key_ignore_case( - GenConstant.DB_TO_SQLALCHEMY_TYPE_MAPPING, column_type - ) - + source_config = DataBaseConfig.get_source(source_name) + normalized = column_type.lower().strip() + base = normalized.split('(', 1)[0].removesuffix(' unsigned').strip() + type_mapping = GenConstant.DB_TO_SQLALCHEMY_TYPE_MAPPING[source_config.db_type] + sqlalchemy_type = StringUtil.get_mapping_value_by_key_ignore_case(type_mapping, base) or 'String' + if '(' in normalized and sqlalchemy_type in {'String', 'CHAR', 'Numeric', 'DECIMAL'}: + return f'{sqlalchemy_type}({normalized.split("(", 1)[1]}' return sqlalchemy_type diff --git a/ruoyi-fastapi-frontend/src/api/tool/gen.js b/ruoyi-fastapi-frontend/src/api/tool/gen.js index 2075677..db5a5f7 100644 --- a/ruoyi-fastapi-frontend/src/api/tool/gen.js +++ b/ruoyi-fastapi-frontend/src/api/tool/gen.js @@ -1,5 +1,13 @@ import request from '@/utils/request' +// 查询代码生成数据源选项 +export function listDataSources() { + return request({ + url: '/tool/gen/dataSources', + method: 'get' + }) +} + // 查询生成表数据 export function listTable(query) { return request({ @@ -69,17 +77,19 @@ export function delTable(tableId) { } // 生成代码(自定义路径) -export function genCode(tableName) { +export function genCode(tableName, dataSourceName) { return request({ url: '/tool/gen/genCode/' + tableName, - method: 'get' + method: 'get', + params: { dataSourceName } }) } // 同步数据库 -export function synchDb(tableName) { +export function synchDb(tableName, dataSourceName) { return request({ url: '/tool/gen/synchDb/' + tableName, - method: 'get' + method: 'get', + params: { dataSourceName } }) } diff --git a/ruoyi-fastapi-frontend/src/views/tool/gen/createTable.vue b/ruoyi-fastapi-frontend/src/views/tool/gen/createTable.vue index f914b5d..1f10601 100644 --- a/ruoyi-fastapi-frontend/src/views/tool/gen/createTable.vue +++ b/ruoyi-fastapi-frontend/src/views/tool/gen/createTable.vue @@ -1,6 +1,23 @@