301 lines
9.4 KiB
Python
301 lines
9.4 KiB
Python
#!/usr/bin/env python
|
||
# -*- coding: utf-8 -*-
|
||
"""表单数据数据库执行适配层:平台库 vs 第三方连接"""
|
||
from __future__ import annotations
|
||
|
||
import logging
|
||
from contextlib import asynccontextmanager
|
||
from typing import Any, Dict, List, Optional, Tuple
|
||
|
||
from sqlalchemy import text
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from online_dev.form_data_manager.dynamic_sql_builder import DynamicSQLBuilder
|
||
from online_dev.form_data_manager.exceptions import (
|
||
FormDataConnectionError,
|
||
FormDataException,
|
||
QueryError,
|
||
)
|
||
from utils.sql_param_compile import compile_sql_with_named_params
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
PLATFORM_DB_TYPE = "postgresql"
|
||
PLATFORM_SCHEMA = "public"
|
||
|
||
|
||
class FormDataWriteForbidden(FormDataException):
|
||
"""系统连接不允许表单业务数据写入"""
|
||
|
||
error_code = "FORM_WRITE_FORBIDDEN"
|
||
http_status = 403
|
||
|
||
def __init__(self, db_config: str = ""):
|
||
super().__init__(
|
||
"系统数据库连接不允许写入表单业务数据,请使用第三方数据库连接",
|
||
context={"db_config": db_config},
|
||
)
|
||
|
||
|
||
class FormDataDbAdapter:
|
||
"""表单 CRUD SQL 执行抽象"""
|
||
|
||
db_type: str
|
||
db_config: str
|
||
is_system: bool
|
||
is_external: bool
|
||
|
||
async def execute_query(
|
||
self,
|
||
sql: str,
|
||
params: Optional[Dict[str, Any]] = None,
|
||
*,
|
||
database: Optional[str] = None,
|
||
) -> List[Dict[str, Any]]:
|
||
raise NotImplementedError
|
||
|
||
async def execute_command(
|
||
self,
|
||
sql: str,
|
||
params: Optional[Dict[str, Any]] = None,
|
||
*,
|
||
database: Optional[str] = None,
|
||
) -> int:
|
||
raise NotImplementedError
|
||
|
||
def ensure_write_allowed(self) -> None:
|
||
"""写操作前校验(系统连接禁止写业务表)"""
|
||
|
||
@asynccontextmanager
|
||
async def transaction(self, database: Optional[str] = None):
|
||
"""请求级事务;平台库由 API Session 提交,第三方由 handler 事务域。"""
|
||
yield
|
||
|
||
|
||
class PlatformSessionAdapter(FormDataDbAdapter):
|
||
"""平台 AsyncSession(db_config=default)"""
|
||
|
||
def __init__(
|
||
self,
|
||
platform_db: AsyncSession,
|
||
db_type: str,
|
||
db_config: str = "default",
|
||
default_database: str = "",
|
||
):
|
||
self._db = platform_db
|
||
self.db_type = db_type
|
||
self.db_config = db_config
|
||
self.default_database = (default_database or "").strip()
|
||
self.is_system = db_config == "default"
|
||
self.is_external = False
|
||
|
||
def ensure_write_allowed(self) -> None:
|
||
return
|
||
|
||
async def execute_query(
|
||
self,
|
||
sql: str,
|
||
params: Optional[Dict[str, Any]] = None,
|
||
*,
|
||
database: Optional[str] = None,
|
||
) -> List[Dict[str, Any]]:
|
||
result = await self._db.execute(text(sql), params or {})
|
||
rows = result.fetchall()
|
||
columns = result.keys()
|
||
return [dict(zip(columns, row)) for row in rows]
|
||
|
||
async def execute_command(
|
||
self,
|
||
sql: str,
|
||
params: Optional[Dict[str, Any]] = None,
|
||
*,
|
||
database: Optional[str] = None,
|
||
) -> int:
|
||
result = await self._db.execute(text(sql), params or {})
|
||
return result.rowcount
|
||
|
||
@asynccontextmanager
|
||
async def transaction(self, database: Optional[str] = None):
|
||
yield
|
||
|
||
|
||
class ExternalConnectionAdapter(FormDataDbAdapter):
|
||
"""第三方连接:经 ConnectionResolver + database_manager 执行"""
|
||
|
||
def __init__(
|
||
self,
|
||
platform_db: AsyncSession,
|
||
db_config: str,
|
||
db_type: str,
|
||
is_system: bool,
|
||
default_database: str = "",
|
||
):
|
||
self._platform_db = platform_db
|
||
self.db_config = db_config
|
||
self.db_type = db_type
|
||
self.default_database = (default_database or "").strip()
|
||
self.is_system = is_system
|
||
self.is_external = True
|
||
self._manager_service = None
|
||
|
||
def ensure_write_allowed(self) -> None:
|
||
if self.is_system:
|
||
raise FormDataWriteForbidden(self.db_config)
|
||
|
||
async def _get_manager_service(self):
|
||
if self._manager_service is None:
|
||
from core.database_manager.service import AsyncDatabaseManagerService
|
||
|
||
self._manager_service = await AsyncDatabaseManagerService.create(
|
||
self.db_config, self._platform_db
|
||
)
|
||
return self._manager_service
|
||
|
||
async def execute_query(
|
||
self,
|
||
sql: str,
|
||
params: Optional[Dict[str, Any]] = None,
|
||
*,
|
||
database: Optional[str] = None,
|
||
) -> List[Dict[str, Any]]:
|
||
compiled = compile_sql_with_named_params(sql, params or {}, self.db_type)
|
||
service = await self._get_manager_service()
|
||
result_data = await service.execute_sql(
|
||
compiled, is_query=True, database=database
|
||
)
|
||
if not result_data.get("success"):
|
||
msg = result_data.get("message", "查询失败")
|
||
logger.error(
|
||
"External form query failed [%s] db_type=%s: %s | SQL: %s",
|
||
self.db_config,
|
||
self.db_type,
|
||
msg,
|
||
compiled[:500] if len(compiled) > 500 else compiled,
|
||
)
|
||
raise QueryError(detail=msg)
|
||
return result_data.get("rows") or []
|
||
|
||
async def execute_command(
|
||
self,
|
||
sql: str,
|
||
params: Optional[Dict[str, Any]] = None,
|
||
*,
|
||
database: Optional[str] = None,
|
||
) -> int:
|
||
self.ensure_write_allowed()
|
||
compiled = compile_sql_with_named_params(sql, params or {}, self.db_type)
|
||
service = await self._get_manager_service()
|
||
result_data = await service.execute_sql(
|
||
compiled, is_query=False, database=database
|
||
)
|
||
if not result_data.get("success"):
|
||
msg = result_data.get("message", "执行失败")
|
||
logger.error(
|
||
"External form command failed [%s] db_type=%s: %s | SQL: %s",
|
||
self.db_config,
|
||
self.db_type,
|
||
msg,
|
||
compiled[:500] if len(compiled) > 500 else compiled,
|
||
)
|
||
raise QueryError(detail=msg)
|
||
affected = result_data.get("affected_rows")
|
||
return int(affected) if affected is not None else 0
|
||
|
||
@asynccontextmanager
|
||
async def transaction(self, database: Optional[str] = None):
|
||
service = await self._get_manager_service()
|
||
if not await service.begin_transaction(database=database):
|
||
raise QueryError(detail="无法开启数据库事务")
|
||
try:
|
||
yield
|
||
await service.commit_transaction()
|
||
except Exception:
|
||
await service.rollback_transaction()
|
||
raise
|
||
|
||
async def create_savepoint(self, name: str) -> None:
|
||
service = await self._get_manager_service()
|
||
await service.create_savepoint(name)
|
||
|
||
async def release_savepoint(self, name: str) -> None:
|
||
service = await self._get_manager_service()
|
||
await service.release_savepoint(name)
|
||
|
||
async def rollback_to_savepoint(self, name: str) -> None:
|
||
service = await self._get_manager_service()
|
||
await service.rollback_to_savepoint(name)
|
||
|
||
|
||
async def create_platform_adapter(platform_db: AsyncSession) -> PlatformSessionAdapter:
|
||
"""平台库适配器(固定 default,用于 core_* 元数据查询)。"""
|
||
from core.database_connection.resolver import ConnectionResolver
|
||
|
||
info = await ConnectionResolver.resolve("default", platform_db)
|
||
return PlatformSessionAdapter(
|
||
platform_db,
|
||
PLATFORM_DB_TYPE,
|
||
"default",
|
||
default_database=info.database or "",
|
||
)
|
||
|
||
|
||
async def create_form_data_adapter(
|
||
db_config: str,
|
||
platform_db: AsyncSession,
|
||
) -> FormDataDbAdapter:
|
||
"""按表单 db_config 创建业务库执行适配器"""
|
||
from core.database_connection.resolver import ConnectionResolver
|
||
|
||
code = (db_config or "default").strip() or "default"
|
||
try:
|
||
info = await ConnectionResolver.resolve(code, platform_db)
|
||
except ValueError as e:
|
||
raise FormDataConnectionError(detail=str(e)) from e
|
||
|
||
if code == "default":
|
||
return PlatformSessionAdapter(
|
||
platform_db,
|
||
info.db_type,
|
||
code,
|
||
default_database=info.database or "",
|
||
)
|
||
|
||
return ExternalConnectionAdapter(
|
||
platform_db,
|
||
code,
|
||
info.db_type,
|
||
info.is_system,
|
||
default_database=info.database or "",
|
||
)
|
||
|
||
|
||
def create_platform_sql_builder() -> DynamicSQLBuilder:
|
||
return DynamicSQLBuilder(PLATFORM_DB_TYPE)
|
||
|
||
|
||
def create_sql_builder_for_adapter(adapter: FormDataDbAdapter) -> DynamicSQLBuilder:
|
||
return DynamicSQLBuilder(
|
||
adapter.db_type,
|
||
default_database=getattr(adapter, "default_database", "") or "",
|
||
)
|
||
|
||
|
||
async def resolve_form_sql_context(
|
||
platform_db: AsyncSession,
|
||
form_meta,
|
||
*,
|
||
adapter_cache: Optional[Dict[str, FormDataDbAdapter]] = None,
|
||
builder_cache: Optional[Dict[str, DynamicSQLBuilder]] = None,
|
||
) -> Tuple[FormDataDbAdapter, DynamicSQLBuilder]:
|
||
"""按 FormMeta 解析业务 adapter 与 sql_builder(支持请求内缓存)。"""
|
||
code = (form_meta.db_config or "default").strip() or "default"
|
||
adapters = adapter_cache if adapter_cache is not None else {}
|
||
builders = builder_cache if builder_cache is not None else {}
|
||
|
||
if code not in adapters:
|
||
adapters[code] = await create_form_data_adapter(code, platform_db)
|
||
if code not in builders:
|
||
builders[code] = create_sql_builder_for_adapter(adapters[code])
|
||
|
||
return adapters[code], builders[code]
|