Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,794 @@
|
||||
"""
|
||||
数据库操作节点
|
||||
"""
|
||||
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'],
|
||||
}
|
||||
Reference in New Issue
Block a user