""" 数据库操作节点 """ 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'], }