Files
2026-06-08 18:14:59 +08:00

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'],
}