#!/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]