Files
ai-agent-admin/backend-fastapi/online_dev/form_data_manager/db_adapter.py
T
2026-06-09 21:18:33 +08:00

301 lines
9.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/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):
"""平台 AsyncSessiondb_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]