Files
RuoYi-Vue3-FastAPI/ruoyi-fastapi-backend/config/database.py
T
insistence ae852b8501 feat: 新增多数据源功能 (#123)
* feat: 新增多数据源功能

* perf: 优化多worker下数据源打印重复日志的问题

* perf: 优化代码

* perf: 优化代码

* chore: 固定ai插件相关依赖
2026-08-21 15:26:24 +08:00

589 lines
20 KiB
Python

import asyncio
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from functools import cache
from typing import Any
from pydantic import SecretStr
from sqlalchemy import URL, Engine, create_engine, text
from sqlalchemy.exc import DBAPIError
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, DataBaseSettings, DataSourceSettings
from exceptions.exception import (
DataSourceException,
DataSourceInitializationException,
DataSourceNotFoundException,
DataSourceUnavailableException,
)
from utils.log_util import logger
_HEALTH_RETRY_COOLDOWN = timedelta(seconds=5)
def _error_details(exc: BaseException) -> tuple[str, int | None]:
"""
提取不含连接凭据和SQL参数的安全错误摘要
:param exc: 原始异常或数据源异常
:return: 异常类型和数字错误码
"""
if isinstance(exc, DataSourceException):
return exc.error_type or type(exc).__name__, exc.error_code
original = exc.orig if isinstance(exc, DBAPIError) else exc
error_code = original.args[0] if original.args and isinstance(original.args[0], int) else None
return type(original).__name__, error_code
def _error_log_suffix(error_type: str, error_code: int | None) -> str:
"""
构建数据源错误日志摘要
:param error_type: 异常类型
:param error_code: 数字错误码
:return: 不含敏感信息的日志后缀
"""
code_text = f',错误码:{error_code}' if error_code is not None else ''
return f',错误类型:{error_type}{code_text}'
@dataclass(frozen=True, slots=True)
class DatabaseDriverAdapter:
"""
数据库驱动适配器
"""
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,
)
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',
),
)
}
def _secret_value(value: SecretStr | str) -> str:
"""
获取密码配置的原始值
:param value: 密码配置
:return: 密码原始值
"""
return value.get_secret_value() if isinstance(value, SecretStr) else value
def _database_source(config: DataBaseSettings | DataSourceSettings) -> DataSourceSettings:
"""
获取指定配置对应的数据源配置
:param config: 数据库集合配置或单个数据源配置
:return: 单个数据源配置
"""
return config.get_source() if isinstance(config, DataBaseSettings) else config
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 _build_url(config: DataSourceSettings, *, sync: bool) -> URL:
"""
使用对应的数据库驱动适配器构建连接URL
: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(
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, config: DataBaseSettings | DataSourceSettings | None = None
) -> Engine:
"""
创建同步SQLAlchemy Engine
:param echo: 是否输出SQLAlchemy SQL日志
:param config: 数据库集合配置或单个数据源配置
:return: 同步SQLAlchemy Engine
"""
source = _database_source(config or DataBaseConfig)
adapter = _driver_adapter(source)
return create_engine(
adapter.build_url(source, sync=True),
connect_args=adapter.build_connect_args(source, sync=True),
**_engine_options(source, echo),
)
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工厂
:param engine: 同步SQLAlchemy Engine
:return: 同步Session工厂
"""
return sessionmaker(bind=engine, autocommit=False, autoflush=False)
@dataclass(slots=True)
class DataSourceRuntime:
"""
数据源运行时状态
"""
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 = field(default_factory=asyncio.Lock)
class _DataSourceRegistry:
"""
数据源注册中心
"""
def __init__(self, settings: DataBaseSettings | None = None) -> None:
self.settings = settings or DataBaseConfig
self._runtimes: dict[str, DataSourceRuntime] = {}
self._initialized = False
self._log_enabled = True
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])
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
@staticmethod
def _data_source_error(
exception_type: type[DataSourceException],
runtime: DataSourceRuntime,
exc: BaseException,
) -> DataSourceException:
"""将底层异常转换为不泄露连接信息的数据源异常。"""
error_type, error_code = _error_details(exc)
return exception_type(runtime.name, error_type=error_type, error_code=error_code)
@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 as exc:
cls._mark_unavailable(runtime)
raise cls._data_source_error(DataSourceInitializationException, runtime, exc) from None
runtime.async_engine = engine
runtime.async_session_factory = session_factory
async def initialize(self, log_enabled: bool = True) -> None:
"""
初始化并检查所有数据源的连接状态
:param log_enabled: 是否输出数据源启动日志
:return: None
"""
self._log_enabled = log_enabled
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()
required_failure: tuple[str, BaseException] | None = None
for name, result in zip(names, results, strict=True):
config = self._configs[name]
required = config.db_required or name == default_name
healthy = not isinstance(result, BaseException)
error_type = None
error_code = None
if not healthy:
error_type, error_code = _error_details(result)
if log_enabled:
log_context: dict[str, Any] = {
'data_source': name,
'database_type': config.db_type,
'required': required,
}
if error_type is not None:
log_context['error_type'] = error_type
log_context['error_code'] = error_code
source_logger = logger.bind(**log_context)
if healthy:
source_logger.info(f'✅ 数据源 {name} 初始化成功')
elif required:
source_logger.error(f'❌ 必需数据源 {name} 连接检查失败{_error_log_suffix(error_type, error_code)}')
else:
source_logger.warning(
f'⚠️ 非必需数据源 {name} 连接检查失败,应用将降级启动{_error_log_suffix(error_type, error_code)}'
)
if not healthy and required and required_failure is None:
required_failure = (name, result)
if required_failure is not None:
name, result = required_failure
error_type, error_code = _error_details(result)
await self.dispose_all()
raise DataSourceInitializationException(
name,
error_type=error_type,
error_code=error_code,
) from None
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 as exc:
raise self._data_source_error(DataSourceInitializationException, runtime, exc) from None
return runtime.sync_engine
async def _check_health(self, name: str) -> None:
"""
检查指定数据源的连接状态
:param name: 数据源名称
:return: None
"""
runtime = self._runtime(name)
async with runtime.health_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 as exc:
self._mark_unavailable(runtime)
raise self._data_source_error(DataSourceUnavailableException, runtime, exc) 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
"""
async with runtime.health_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)
if self._log_enabled:
logger.bind(data_source=runtime.name).info(f'✅ 数据源 {runtime.name} 连接已恢复')
@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 DBAPIError as exc:
if not exc.connection_invalidated:
raise
self._mark_unavailable(runtime)
raise self._data_source_error(DataSourceUnavailableException, runtime, exc) 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 DBAPIError as exc:
if not exc.connection_invalidated:
raise
self._mark_unavailable(runtime)
raise self._data_source_error(DataSourceUnavailableException, runtime, exc) 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:
try:
runtime.sync_engine.dispose()
except Exception as exc:
error_type, error_code = _error_details(exc)
logger.bind(
data_source=runtime.name,
engine_type='sync',
error_type=error_type,
error_code=error_code,
).warning(f'⚠️ 数据源 {runtime.name} 同步Engine释放失败{_error_log_suffix(error_type, error_code)}')
async_resources = [(runtime, engine) for runtime in runtimes if (engine := runtime.async_engine) is not None]
results = await asyncio.gather(
*(engine.dispose() for _, engine in async_resources),
return_exceptions=True,
)
for (runtime, _), result in zip(async_resources, results, strict=True):
if isinstance(result, Exception):
error_type, error_code = _error_details(result)
logger.bind(
data_source=runtime.name,
engine_type='async',
error_type=error_type,
error_code=error_code,
).warning(f'⚠️ 数据源 {runtime.name} 异步Engine释放失败{_error_log_suffix(error_type, error_code)}')
elif isinstance(result, BaseException):
raise result
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