Restore AI workflow design nodes
This commit is contained in:
@@ -0,0 +1,300 @@
|
||||
#!/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]
|
||||
Reference in New Issue
Block a user