795 lines
29 KiB
Python
795 lines
29 KiB
Python
"""
|
|
数据库操作节点
|
|
"""
|
|
import json
|
|
import logging
|
|
import time
|
|
import uuid
|
|
from datetime import date, datetime
|
|
from decimal import Decimal
|
|
from typing import Any, Dict, List, Optional
|
|
from uuid import UUID
|
|
|
|
from sqlalchemy import text
|
|
|
|
from ..base import BaseNode, NodeContext, NodeResult
|
|
from ..registry import NodeRegistry
|
|
from ..utils.db_execution import (
|
|
DbTarget,
|
|
build_sql_param_dict,
|
|
build_where_clause_platform,
|
|
build_where_clause_raw,
|
|
default_connection_write_warnings,
|
|
format_limit_clause,
|
|
format_select_sql,
|
|
merge_result_metadata,
|
|
normalize_return_fields,
|
|
quote_table_for_target,
|
|
resolve_db_target,
|
|
resolve_handler_schema_name,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def serialize_value(value: Any) -> Any:
|
|
"""将数据库值转换为可 JSON 序列化的格式"""
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, datetime):
|
|
return value.isoformat()
|
|
if isinstance(value, date):
|
|
return value.isoformat()
|
|
if isinstance(value, Decimal):
|
|
return float(value)
|
|
if isinstance(value, UUID):
|
|
return str(value)
|
|
if isinstance(value, bytes):
|
|
return value.decode('utf-8', errors='replace')
|
|
if isinstance(value, (list, tuple)):
|
|
return [serialize_value(v) for v in value]
|
|
if isinstance(value, dict):
|
|
return {k: serialize_value(v) for k, v in value.items()}
|
|
return value
|
|
|
|
|
|
def serialize_row(row: Dict[str, Any]) -> Dict[str, Any]:
|
|
"""序列化数据库行"""
|
|
return {k: serialize_value(v) for k, v in row.items()}
|
|
|
|
|
|
def prepare_value_for_db(value: Any) -> Any:
|
|
"""将值转换为数据库可接受的格式(dict/list 转为 JSON 字符串)"""
|
|
if value is None:
|
|
return None
|
|
if isinstance(value, (dict, list)):
|
|
return json.dumps(value, ensure_ascii=False)
|
|
if isinstance(value, datetime):
|
|
return value.isoformat()
|
|
if isinstance(value, date):
|
|
return value.isoformat()
|
|
if isinstance(value, UUID):
|
|
return str(value)
|
|
return value
|
|
|
|
|
|
ALLOWED_TABLES = []
|
|
|
|
PROTECTED_FIELDS = ['password', 'token', 'secret', 'api_key', 'private_key']
|
|
|
|
|
|
class BaseDatabaseNode(BaseNode):
|
|
"""
|
|
数据库操作节点基类
|
|
|
|
default 连接走平台 AsyncSession;第三方连接走 AsyncDatabaseManagerService。
|
|
"""
|
|
|
|
node_type = 'database'
|
|
node_name = '数据库操作'
|
|
node_category = 'data'
|
|
node_icon = 'database'
|
|
node_description = '对数据库进行增删改查操作'
|
|
|
|
inputs = [
|
|
{
|
|
'name': 'data',
|
|
'type': 'object',
|
|
'description': '要操作的数据',
|
|
},
|
|
]
|
|
|
|
outputs = [
|
|
{
|
|
'name': 'result',
|
|
'type': 'object',
|
|
'description': '操作结果',
|
|
},
|
|
{
|
|
'name': 'affected_rows',
|
|
'type': 'number',
|
|
'description': '影响的行数',
|
|
},
|
|
]
|
|
|
|
def _get_db_target(self) -> DbTarget:
|
|
return resolve_db_target(self.config.get('db_config'))
|
|
|
|
def _build_full_table_name(self, table: str, target: Optional[DbTarget] = None) -> str:
|
|
"""构建完整的表名(平台 PG 路径)"""
|
|
db_config = self.config.get('db_config', {})
|
|
schema = db_config.get('schema', '')
|
|
if schema:
|
|
return f'"{schema}"."{table}"'
|
|
return f'"{table}"'
|
|
|
|
def execute(self, context: NodeContext) -> NodeResult:
|
|
import asyncio
|
|
loop = asyncio.get_event_loop()
|
|
if loop.is_running():
|
|
import concurrent.futures
|
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
|
future = executor.submit(asyncio.run, self.execute_async(context))
|
|
return future.result()
|
|
return loop.run_until_complete(self.execute_async(context))
|
|
|
|
async def execute_async(self, context: NodeContext) -> NodeResult:
|
|
start_time = time.time()
|
|
|
|
try:
|
|
operation = self.config.get('operation', 'select').lower()
|
|
table = self.config.get('table', '')
|
|
output_variable = self.config.get('output_variable', 'db_result')
|
|
frontend_max_rows = int(self.config.get('frontend_max_rows', 100))
|
|
target = self._get_db_target()
|
|
|
|
if not table:
|
|
raise ValueError('未指定目标表')
|
|
|
|
if ALLOWED_TABLES and table not in ALLOWED_TABLES:
|
|
is_allowed = any(allowed == '*' or table == allowed for allowed in ALLOWED_TABLES)
|
|
if not is_allowed:
|
|
raise ValueError(f'表 {table} 不在允许操作的白名单中')
|
|
|
|
if operation == 'insert':
|
|
result = await self._execute_insert(table, context, target)
|
|
elif operation == 'update':
|
|
result = await self._execute_update(table, context, target)
|
|
elif operation == 'upsert':
|
|
result = await self._execute_upsert(table, context, target)
|
|
elif operation == 'select':
|
|
result = await self._execute_select(table, context, target)
|
|
elif operation == 'delete':
|
|
result = await self._execute_delete(table, context, target)
|
|
else:
|
|
raise ValueError(f'不支持的操作类型: {operation}')
|
|
|
|
elapsed_time = int((time.time() - start_time) * 1000)
|
|
full_data = result.get('data')
|
|
affected_rows = result.get('affected_rows', 0)
|
|
|
|
output_variables = {
|
|
output_variable: full_data,
|
|
f'{output_variable}_count': affected_rows,
|
|
}
|
|
|
|
frontend_output_variables = output_variables
|
|
frontend_output = full_data
|
|
|
|
if operation == 'select' and isinstance(full_data, list) and len(full_data) > frontend_max_rows:
|
|
truncated = full_data[:frontend_max_rows]
|
|
frontend_output = truncated
|
|
frontend_output_variables = {
|
|
output_variable: truncated,
|
|
f'{output_variable}_count': affected_rows,
|
|
f'{output_variable}_total': len(full_data),
|
|
}
|
|
|
|
warnings = default_connection_write_warnings(operation, target)
|
|
metadata = merge_result_metadata(
|
|
{'frontend_output_variables': frontend_output_variables},
|
|
warnings,
|
|
)
|
|
|
|
return NodeResult(
|
|
success=True,
|
|
output=frontend_output,
|
|
output_variables=output_variables,
|
|
metadata=metadata,
|
|
elapsed_time=elapsed_time,
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.exception(f'数据库节点执行失败: {e}')
|
|
return NodeResult(
|
|
success=False,
|
|
error=str(e),
|
|
elapsed_time=int((time.time() - start_time) * 1000),
|
|
)
|
|
|
|
async def _create_db_service(self, context: NodeContext, target: DbTarget):
|
|
from core.database_manager.service import AsyncDatabaseManagerService
|
|
|
|
try:
|
|
return await AsyncDatabaseManagerService.create(target.db_name, context.db_session)
|
|
except Exception as e:
|
|
raise ValueError(f'无法连接数据库 {target.db_name}: {e}') from e
|
|
|
|
def _resolve_field_mapping(self, context: NodeContext) -> Dict[str, Any]:
|
|
field_mapping = self.config.get('field_mapping', {})
|
|
resolved = {}
|
|
|
|
for field, value in field_mapping.items():
|
|
if field.lower() in PROTECTED_FIELDS:
|
|
logger.warning(f'跳过保护字段: {field}')
|
|
continue
|
|
|
|
if isinstance(value, str):
|
|
resolved_value = context.resolve_template(value)
|
|
try:
|
|
resolved[field] = json.loads(resolved_value)
|
|
except (json.JSONDecodeError, TypeError):
|
|
resolved[field] = resolved_value
|
|
else:
|
|
resolved[field] = value
|
|
|
|
return resolved
|
|
|
|
def _resolve_conditions(self, context: NodeContext) -> List[Dict[str, Any]]:
|
|
conditions = self.config.get('where_conditions', [])
|
|
resolved = []
|
|
|
|
for condition in conditions:
|
|
field = condition.get('field', '')
|
|
operator = condition.get('operator', '=')
|
|
value = condition.get('value', '')
|
|
|
|
if isinstance(value, str):
|
|
resolved_value = context.resolve_template(value)
|
|
try:
|
|
resolved_value = json.loads(resolved_value)
|
|
except (json.JSONDecodeError, TypeError):
|
|
pass
|
|
else:
|
|
resolved_value = value
|
|
|
|
resolved.append({
|
|
'field': field,
|
|
'operator': operator,
|
|
'value': resolved_value,
|
|
})
|
|
|
|
return resolved
|
|
|
|
async def _execute_insert(
|
|
self,
|
|
table: str,
|
|
context: NodeContext,
|
|
target: DbTarget,
|
|
) -> Dict[str, Any]:
|
|
data = self._resolve_field_mapping(context)
|
|
if not data:
|
|
raise ValueError('没有要插入的数据')
|
|
if 'id' not in data:
|
|
data['id'] = str(uuid.uuid4())
|
|
|
|
if target.is_external:
|
|
db_service = await self._create_db_service(context, target)
|
|
schema_name = await resolve_handler_schema_name(db_service, target)
|
|
payload = {k: prepare_value_for_db(v) for k, v in data.items()}
|
|
result = await db_service.insert_data(table, payload, schema_name)
|
|
if not result.get('success'):
|
|
raise ValueError(result.get('message', '插入失败'))
|
|
return {
|
|
'data': {'id': data['id'], **data},
|
|
'affected_rows': result.get('affected_rows', 1),
|
|
}
|
|
|
|
db = context.db_session
|
|
if not db:
|
|
raise ValueError('数据库会话不可用,请确保工作流已配置数据库连接')
|
|
|
|
full_table_name = self._build_full_table_name(table)
|
|
fields = list(data.keys())
|
|
params = {f: prepare_value_for_db(v) for f, v in data.items()}
|
|
placeholders = ', '.join([f':{f}' for f in fields])
|
|
field_names = ', '.join([f'"{f}"' for f in fields])
|
|
sql = f'INSERT INTO {full_table_name} ({field_names}) VALUES ({placeholders})'
|
|
await db.execute(text(sql), params)
|
|
return {
|
|
'data': {'id': data['id'], **data},
|
|
'affected_rows': 1,
|
|
}
|
|
|
|
async def _execute_update(
|
|
self,
|
|
table: str,
|
|
context: NodeContext,
|
|
target: DbTarget,
|
|
) -> Dict[str, Any]:
|
|
data = self._resolve_field_mapping(context)
|
|
conditions = self._resolve_conditions(context)
|
|
|
|
if not data:
|
|
raise ValueError('没有要更新的数据')
|
|
if not conditions:
|
|
raise ValueError('UPDATE 操作必须指定条件,防止误更新全表')
|
|
|
|
if target.is_external:
|
|
db_service = await self._create_db_service(context, target)
|
|
schema_name = await resolve_handler_schema_name(db_service, target)
|
|
where_raw = build_where_clause_raw(conditions, db_service.db_type)
|
|
payload = {k: prepare_value_for_db(v) for k, v in data.items()}
|
|
result = await db_service.update_data(table, payload, where_raw, schema_name)
|
|
if not result.get('success'):
|
|
raise ValueError(result.get('message', '更新失败'))
|
|
affected_rows = result.get('affected_rows', 0)
|
|
return {
|
|
'data': {'updated': True, 'affected_rows': affected_rows, **data},
|
|
'affected_rows': affected_rows,
|
|
}
|
|
|
|
db = context.db_session
|
|
if not db:
|
|
raise ValueError('数据库会话不可用')
|
|
|
|
set_clauses = []
|
|
params = {}
|
|
for field, value in data.items():
|
|
param_name = f's_{field}'
|
|
set_clauses.append(f'"{field}" = :{param_name}')
|
|
params[param_name] = prepare_value_for_db(value)
|
|
|
|
where_clause, where_params = build_where_clause_platform(conditions)
|
|
params.update(where_params)
|
|
full_table_name = self._build_full_table_name(table)
|
|
sql = f'UPDATE {full_table_name} SET {", ".join(set_clauses)} {where_clause}'
|
|
result = await db.execute(text(sql), params)
|
|
affected_rows = result.rowcount
|
|
return {
|
|
'data': {'updated': True, 'affected_rows': affected_rows, **data},
|
|
'affected_rows': affected_rows,
|
|
}
|
|
|
|
async def _execute_upsert(
|
|
self,
|
|
table: str,
|
|
context: NodeContext,
|
|
target: DbTarget,
|
|
) -> Dict[str, Any]:
|
|
data = self._resolve_field_mapping(context)
|
|
conditions = self._resolve_conditions(context)
|
|
if not data:
|
|
raise ValueError('没有要操作的数据')
|
|
|
|
if conditions:
|
|
if target.is_external:
|
|
db_service = await self._create_db_service(context, target)
|
|
schema_name = await resolve_handler_schema_name(db_service, target)
|
|
where_raw = build_where_clause_raw(conditions, db_service.db_type)
|
|
full_table = quote_table_for_target(table, target)
|
|
check_sql = f'SELECT 1 FROM {full_table}'
|
|
if where_raw:
|
|
check_sql += f' WHERE {where_raw}'
|
|
check_sql += format_limit_clause(db_service.db_type, 1)
|
|
check_result = await db_service.execute_sql(check_sql, is_query=True)
|
|
rows = check_result.get('rows') or check_result.get('data') or []
|
|
if rows:
|
|
return await self._execute_update(table, context, target)
|
|
else:
|
|
db = context.db_session
|
|
if not db:
|
|
raise ValueError('数据库会话不可用')
|
|
full_table_name = self._build_full_table_name(table)
|
|
where_clause, where_params = build_where_clause_platform(conditions)
|
|
check_sql = f'SELECT id FROM {full_table_name} {where_clause} LIMIT 1'
|
|
result = await db.execute(text(check_sql), where_params)
|
|
if result.fetchone():
|
|
return await self._execute_update(table, context, target)
|
|
|
|
return await self._execute_insert(table, context, target)
|
|
|
|
async def _execute_select(
|
|
self,
|
|
table: str,
|
|
context: NodeContext,
|
|
target: DbTarget,
|
|
) -> Dict[str, Any]:
|
|
conditions = self._resolve_conditions(context)
|
|
return_fields = self.config.get('return_fields', ['*'])
|
|
limit = self.config.get('limit', 100)
|
|
order_by = self.config.get('order_by', '')
|
|
|
|
if target.is_external:
|
|
db_service = await self._create_db_service(context, target)
|
|
sql = format_select_sql(
|
|
table,
|
|
target,
|
|
return_fields=return_fields,
|
|
conditions=conditions,
|
|
order_by=order_by,
|
|
limit=int(limit),
|
|
)
|
|
result_data = await db_service.execute_sql(sql, is_query=True)
|
|
if result_data.get('success') is False:
|
|
raise ValueError(result_data.get('message') or '查询失败')
|
|
rows = result_data.get('rows') or result_data.get('data') or []
|
|
result_data_list = []
|
|
for row in rows:
|
|
if isinstance(row, dict):
|
|
result_data_list.append(serialize_row(row))
|
|
else:
|
|
result_data_list.append(serialize_row(dict(row)))
|
|
return {
|
|
'data': result_data_list,
|
|
'affected_rows': len(result_data_list),
|
|
}
|
|
|
|
db = context.db_session
|
|
if not db:
|
|
raise ValueError('数据库会话不可用')
|
|
|
|
normalized_fields = normalize_return_fields(return_fields)
|
|
if normalized_fields == '*':
|
|
field_list = '*'
|
|
else:
|
|
field_list = ', '.join([f'"{f}"' for f in normalized_fields])
|
|
|
|
where_clause, params = build_where_clause_platform(conditions)
|
|
full_table_name = self._build_full_table_name(table)
|
|
sql = f'SELECT {field_list} FROM {full_table_name} {where_clause}'
|
|
if order_by:
|
|
sql += f' ORDER BY {order_by}'
|
|
sql += f' LIMIT {int(limit)}'
|
|
|
|
result = await db.execute(text(sql), params)
|
|
columns = result.keys()
|
|
rows = result.fetchall()
|
|
result_data = [serialize_row(dict(zip(columns, row))) for row in rows]
|
|
return {
|
|
'data': result_data,
|
|
'affected_rows': len(result_data),
|
|
}
|
|
|
|
async def _execute_delete(
|
|
self,
|
|
table: str,
|
|
context: NodeContext,
|
|
target: DbTarget,
|
|
) -> Dict[str, Any]:
|
|
conditions = self._resolve_conditions(context)
|
|
if not conditions:
|
|
raise ValueError('DELETE 操作必须指定条件,防止误删全表')
|
|
|
|
if target.is_external:
|
|
db_service = await self._create_db_service(context, target)
|
|
schema_name = await resolve_handler_schema_name(db_service, target)
|
|
where_raw = build_where_clause_raw(conditions, db_service.db_type)
|
|
result = await db_service.delete_data(table, where_raw, schema_name)
|
|
if not result.get('success'):
|
|
raise ValueError(result.get('message', '删除失败'))
|
|
affected_rows = result.get('affected_rows', 0)
|
|
return {
|
|
'data': {'deleted': True, 'affected_rows': affected_rows},
|
|
'affected_rows': affected_rows,
|
|
}
|
|
|
|
db = context.db_session
|
|
if not db:
|
|
raise ValueError('数据库会话不可用')
|
|
|
|
where_clause, params = build_where_clause_platform(conditions)
|
|
full_table_name = self._build_full_table_name(table)
|
|
sql = f'DELETE FROM {full_table_name} {where_clause}'
|
|
result = await db.execute(text(sql), params)
|
|
affected_rows = result.rowcount
|
|
return {
|
|
'data': {'deleted': True, 'affected_rows': affected_rows},
|
|
'affected_rows': affected_rows,
|
|
}
|
|
|
|
@classmethod
|
|
def get_config_schema(cls) -> Dict[str, Any]:
|
|
return {
|
|
'type': 'object',
|
|
'properties': {
|
|
'operation': {
|
|
'type': 'string',
|
|
'title': '操作类型',
|
|
'enum': ['insert', 'update', 'upsert', 'select', 'delete'],
|
|
'enumNames': ['插入', '更新', '插入或更新', '查询', '删除'],
|
|
'default': 'insert',
|
|
},
|
|
'table': {
|
|
'type': 'string',
|
|
'title': '目标表',
|
|
'description': '数据库表名',
|
|
},
|
|
'field_mapping': {
|
|
'type': 'object',
|
|
'title': '字段映射',
|
|
'description': '数据库字段与变量的映射关系',
|
|
'additionalProperties': {'type': 'string'},
|
|
},
|
|
'where_conditions': {
|
|
'type': 'array',
|
|
'title': '条件',
|
|
'description': '查询/更新/删除的条件',
|
|
'items': {
|
|
'type': 'object',
|
|
'properties': {
|
|
'field': {'type': 'string', 'title': '字段'},
|
|
'operator': {
|
|
'type': 'string',
|
|
'title': '操作符',
|
|
'enum': ['=', '!=', '>', '>=', '<', '<=', 'like', 'in', 'is_null', 'is_not_null'],
|
|
'default': '=',
|
|
},
|
|
'value': {'type': 'string', 'title': '值'},
|
|
},
|
|
},
|
|
},
|
|
'return_fields': {
|
|
'type': 'array',
|
|
'title': '返回字段',
|
|
'description': '查询时返回的字段列表',
|
|
'items': {'type': 'string'},
|
|
'default': ['*'],
|
|
},
|
|
'limit': {
|
|
'type': 'integer',
|
|
'title': '限制条数',
|
|
'description': 'SQL 查询时的最大返回条数',
|
|
'default': 100,
|
|
},
|
|
'frontend_max_rows': {
|
|
'type': 'integer',
|
|
'title': '前端返回最大条数',
|
|
'description': '前端 SSE 事件中返回的最大数据条数(默认100),超过此值仅截断前端传输,后续节点仍可获取全量数据',
|
|
'default': 100,
|
|
'minimum': 1,
|
|
'maximum': 10000,
|
|
},
|
|
'order_by': {
|
|
'type': 'string',
|
|
'title': '排序',
|
|
'description': '排序字段,如 created_at DESC',
|
|
},
|
|
'output_variable': {
|
|
'type': 'string',
|
|
'title': '输出变量名',
|
|
'default': 'db_result',
|
|
},
|
|
},
|
|
'required': ['operation', 'table'],
|
|
}
|
|
|
|
|
|
@NodeRegistry.register
|
|
class DbInsertNode(BaseDatabaseNode):
|
|
node_type = 'db_insert'
|
|
node_name = 'DB 插入'
|
|
node_icon = 'database-zap'
|
|
node_description = '向数据库插入数据'
|
|
|
|
def __init__(self, config: Dict[str, Any] = None):
|
|
super().__init__(config)
|
|
if config:
|
|
self.config['operation'] = 'upsert' if config.get('upsert') else 'insert'
|
|
|
|
|
|
@NodeRegistry.register
|
|
class DbUpdateNode(BaseDatabaseNode):
|
|
node_type = 'db_update'
|
|
node_name = 'DB 更新'
|
|
node_icon = 'database-backup'
|
|
node_description = '更新数据库记录'
|
|
|
|
def __init__(self, config: Dict[str, Any] = None):
|
|
super().__init__(config)
|
|
if config:
|
|
self.config['operation'] = 'update'
|
|
|
|
|
|
@NodeRegistry.register
|
|
class DbQueryNode(BaseDatabaseNode):
|
|
node_type = 'db_query'
|
|
node_name = 'DB 查询'
|
|
node_icon = 'search'
|
|
node_description = '从数据库查询数据'
|
|
|
|
def __init__(self, config: Dict[str, Any] = None):
|
|
super().__init__(config)
|
|
if config:
|
|
self.config['operation'] = 'select'
|
|
|
|
|
|
@NodeRegistry.register
|
|
class DbDeleteNode(BaseDatabaseNode):
|
|
node_type = 'db_delete'
|
|
node_name = 'DB 删除'
|
|
node_icon = 'trash-2'
|
|
node_description = '从数据库删除数据'
|
|
|
|
def __init__(self, config: Dict[str, Any] = None):
|
|
super().__init__(config)
|
|
if config:
|
|
self.config['operation'] = 'delete'
|
|
|
|
|
|
@NodeRegistry.register
|
|
class DbSqlNode(BaseNode):
|
|
"""自定义 SQL 执行节点"""
|
|
|
|
node_type = 'db_sql'
|
|
node_name = 'SQL 执行'
|
|
node_category = 'data'
|
|
node_icon = 'database'
|
|
node_description = '执行自定义 SQL 语句'
|
|
|
|
inputs = [
|
|
{
|
|
'name': 'data',
|
|
'type': 'object',
|
|
'description': '输入数据',
|
|
},
|
|
]
|
|
|
|
outputs = [
|
|
{
|
|
'name': 'result',
|
|
'type': 'any',
|
|
'description': 'SQL 执行结果',
|
|
},
|
|
]
|
|
|
|
def execute(self, context: NodeContext) -> NodeResult:
|
|
import asyncio
|
|
loop = asyncio.get_event_loop()
|
|
if loop.is_running():
|
|
import concurrent.futures
|
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
|
future = executor.submit(asyncio.run, self.execute_async(context))
|
|
return future.result()
|
|
return loop.run_until_complete(self.execute_async(context))
|
|
|
|
async def execute_async(self, context: NodeContext) -> NodeResult:
|
|
start_time = time.time()
|
|
|
|
sql_type = self.config.get('sql_type', 'query')
|
|
sql = self.config.get('sql', '')
|
|
target = resolve_db_target(self.config.get('db_config'))
|
|
output_variable = self.config.get('output_variable', 'sql_result')
|
|
is_query = sql_type == 'query'
|
|
|
|
if not sql:
|
|
return NodeResult(
|
|
success=False,
|
|
error='SQL 语句不能为空',
|
|
elapsed_time=int((time.time() - start_time) * 1000),
|
|
)
|
|
|
|
try:
|
|
resolved_sql = context.resolve_template(sql)
|
|
param_dict = build_sql_param_dict(self.config.get('params'), context)
|
|
logger.info(
|
|
'执行 SQL [%s]: %s, params=%s',
|
|
target.db_name,
|
|
resolved_sql,
|
|
list(param_dict.keys()),
|
|
)
|
|
|
|
operation = 'query' if is_query else 'execute'
|
|
warnings = default_connection_write_warnings(operation, target)
|
|
|
|
if not target.is_external:
|
|
db = context.db_session
|
|
if not db:
|
|
raise ValueError('数据库会话不可用')
|
|
result = await db.execute(text(resolved_sql), param_dict)
|
|
if is_query:
|
|
rows = result.mappings().all()
|
|
output_result = [serialize_row(dict(row)) for row in rows]
|
|
row_count = len(output_result)
|
|
else:
|
|
row_count = max(result.rowcount or 0, 0)
|
|
output_result = row_count
|
|
else:
|
|
from core.database_manager.service import AsyncDatabaseManagerService
|
|
from utils.sql_param_compile import compile_sql_with_named_params
|
|
|
|
try:
|
|
db_service = await AsyncDatabaseManagerService.create(
|
|
target.db_name,
|
|
context.db_session,
|
|
)
|
|
except Exception as e:
|
|
raise ValueError(f'无法连接数据库 {target.db_name}: {e}') from e
|
|
|
|
executable_sql = compile_sql_with_named_params(
|
|
resolved_sql,
|
|
param_dict,
|
|
target.db_type,
|
|
)
|
|
result_data = await db_service.execute_sql(executable_sql, is_query=is_query)
|
|
if result_data.get('success') is False:
|
|
raise Exception(result_data.get('message') or 'SQL 执行失败')
|
|
|
|
if is_query:
|
|
rows = result_data.get('rows') or result_data.get('data') or []
|
|
output_result = []
|
|
for row in rows:
|
|
if isinstance(row, dict):
|
|
output_result.append(serialize_row(row))
|
|
else:
|
|
output_result.append(serialize_row(dict(row)))
|
|
row_count = len(output_result)
|
|
else:
|
|
output_result = result_data.get('affected_rows', 0)
|
|
row_count = output_result
|
|
|
|
elapsed_time = int((time.time() - start_time) * 1000)
|
|
metadata = merge_result_metadata({}, warnings)
|
|
|
|
return NodeResult(
|
|
success=True,
|
|
output=output_result,
|
|
output_variables={
|
|
output_variable: output_result,
|
|
f'{output_variable}_count': row_count,
|
|
},
|
|
metadata=metadata,
|
|
elapsed_time=elapsed_time,
|
|
)
|
|
|
|
except Exception as e:
|
|
logger.error('SQL 执行失败: %s', e)
|
|
return NodeResult(
|
|
success=False,
|
|
error=f'SQL 执行失败: {str(e)}',
|
|
elapsed_time=int((time.time() - start_time) * 1000),
|
|
)
|
|
|
|
@classmethod
|
|
def get_config_schema(cls) -> Dict[str, Any]:
|
|
return {
|
|
'type': 'object',
|
|
'properties': {
|
|
'sql_type': {
|
|
'type': 'string',
|
|
'title': '执行类型',
|
|
'enum': ['query', 'execute'],
|
|
'default': 'query',
|
|
'description': 'query: 查询返回结果, execute: 执行不返回结果',
|
|
},
|
|
'sql': {
|
|
'type': 'string',
|
|
'title': 'SQL 语句',
|
|
'description': '要执行的 SQL 语句,使用 :param_name 作为命名参数占位符',
|
|
},
|
|
'params': {
|
|
'type': 'array',
|
|
'title': '参数列表',
|
|
'description': 'SQL 命名参数,与 SQL 中 :param_name 对应',
|
|
'items': {
|
|
'type': 'object',
|
|
'properties': {
|
|
'name': {'type': 'string', 'title': '参数名'},
|
|
'type': {
|
|
'type': 'string',
|
|
'enum': ['string', 'integer', 'float', 'boolean', 'date', 'datetime'],
|
|
'default': 'string',
|
|
},
|
|
'value': {'type': 'string', 'title': '参数值'},
|
|
},
|
|
},
|
|
},
|
|
'output_variable': {
|
|
'type': 'string',
|
|
'title': '输出变量名',
|
|
'default': 'sql_result',
|
|
},
|
|
},
|
|
'required': ['sql'],
|
|
}
|