""" 工作流数据库节点:连接解析与方言 SQL 构建 """ from __future__ import annotations import json from dataclasses import dataclass from datetime import date, datetime from decimal import Decimal from typing import Any, Dict, List, Optional, Tuple from uuid import UUID from core.database_manager.sql_utils import quote_identifier, quote_table DEFAULT_CONNECTION_CODE = "default" DEFAULT_CONNECTION_WRITE_WARNING = "default_connection_write" OPERATOR_MAP = { "=": "=", "!=": "!=", ">": ">", ">=": ">=", "<": "<", "<=": "<=", "like": "LIKE", "in": "IN", "is_null": "IS NULL", "is_not_null": "IS NOT NULL", } @dataclass class DbTarget: """工作流节点数据库目标""" db_name: str = DEFAULT_CONNECTION_CODE db_type: str = "postgresql" database: str = "" schema: str = "" @property def is_external(self) -> bool: return is_external_connection(self.db_name) def is_external_connection(db_name: Optional[str]) -> bool: code = (db_name or DEFAULT_CONNECTION_CODE).strip() or DEFAULT_CONNECTION_CODE return code != DEFAULT_CONNECTION_CODE def resolve_db_target(config: Optional[Dict[str, Any]]) -> DbTarget: """从节点 db_config 解析连接目标""" db_config = config or {} db_name = (db_config.get("dbName") or DEFAULT_CONNECTION_CODE).strip() or DEFAULT_CONNECTION_CODE db_type = (db_config.get("dbType") or "postgresql").lower() database = (db_config.get("database") or "").strip() schema = (db_config.get("schema") or "").strip() return DbTarget( db_name=db_name, db_type=db_type, database=database, schema=schema, ) def default_connection_write_warnings(operation: str, target: DbTarget) -> List[str]: """default 连接写操作返回运行时警告码""" if target.is_external: return [] write_ops = {"insert", "update", "delete", "upsert", "execute"} if operation.lower() in write_ops: return [DEFAULT_CONNECTION_WRITE_WARNING] return [] def resolve_schema_for_handler(db_type: str, schema: str, default_schema: str = "") -> str: """解析 handler 使用的 schema/database 参数""" db = (db_type or "postgresql").lower() if db == "mysql": return schema or default_schema or "" if not schema or (schema == "public" and db != "postgresql"): return default_schema or schema or "" return schema async def resolve_handler_schema_name(db_service, target: DbTarget) -> str: """结合 AsyncDatabaseManagerService 解析 schema""" db_type = (db_service.db_type or "postgresql").lower() if db_type == "mysql": default_schema = db_service._default_schema() if hasattr(db_service, "_default_schema") else "" return target.database or target.schema or default_schema or "" default_schema = "" if hasattr(db_service, "_default_schema"): default_schema = db_service._default_schema() or "" return resolve_schema_for_handler(db_type, target.schema, default_schema) def format_sql_literal(value: Any, db_type: str) -> str: """将 Python 值格式化为 SQL 字面量(用于 handler raw WHERE)""" if value is None: return "NULL" if isinstance(value, bool): if (db_type or "").lower() == "postgresql": return "TRUE" if value else "FALSE" return "1" if value else "0" if isinstance(value, (int, float, Decimal)): return str(value) if isinstance(value, (datetime, date)): return f"'{value.isoformat()}'" if isinstance(value, UUID): return f"'{value}'" if isinstance(value, (list, tuple)): inner = ", ".join(format_sql_literal(v, db_type) for v in value) return f"({inner})" if isinstance(value, (dict, list)): encoded = json.dumps(value, ensure_ascii=False) escaped = encoded.replace("'", "''") return f"'{escaped}'" escaped = str(value).replace("'", "''") return f"'{escaped}'" def build_where_clause_raw( conditions: List[Dict[str, Any]], db_type: str, ) -> str: """ 构建 raw WHERE 子句(不含 WHERE 关键字),供 database_manager handler 使用。 """ if not conditions: return "" clauses: List[str] = [] for condition in conditions: field = condition.get("field", "") op = (condition.get("operator") or "=").lower() value = condition.get("value") sql_op = OPERATOR_MAP.get(op, "=") quoted_field = quote_identifier(field, db_type) if op in ("is_null", "is_not_null"): clauses.append(f"{quoted_field} {sql_op}") elif op == "in": if isinstance(value, list): literals = ", ".join(format_sql_literal(v, db_type) for v in value) clauses.append(f"{quoted_field} IN ({literals})") else: clauses.append(f"{quoted_field} IN ({format_sql_literal(value, db_type)})") elif op == "like": clauses.append(f"{quoted_field} LIKE {format_sql_literal(value, db_type)}") else: clauses.append(f"{quoted_field} {sql_op} {format_sql_literal(value, db_type)}") return " AND ".join(clauses) def build_where_clause_platform( conditions: List[Dict[str, Any]], ) -> Tuple[str, Dict[str, Any]]: """构建平台 PostgreSQL WHERE 子句(SQLAlchemy 命名参数)""" if not conditions: return "", {} clauses: List[str] = [] params: Dict[str, Any] = {} param_index = 0 for condition in conditions: field = condition["field"] op = condition["operator"].lower() value = condition["value"] sql_op = OPERATOR_MAP.get(op, "=") if op in ("is_null", "is_not_null"): clauses.append(f'"{field}" {sql_op}') elif op == "in": if isinstance(value, list): param_placeholders = [] for v in value: param_name = f"p{param_index}" param_placeholders.append(f":{param_name}") params[param_name] = v param_index += 1 clauses.append(f'"{field}" IN ({", ".join(param_placeholders)})') else: param_name = f"p{param_index}" clauses.append(f'"{field}" {sql_op} :{param_name}') params[param_name] = value param_index += 1 else: param_name = f"p{param_index}" clauses.append(f'"{field}" {sql_op} :{param_name}') params[param_name] = value param_index += 1 where_clause = " AND ".join(clauses) return f"WHERE {where_clause}", params def quote_table_for_target(table: str, target: DbTarget) -> str: """按目标连接类型构建完整表名""" schema = target.schema or None if target.db_type == "mysql" and not schema and target.database: schema = target.database return quote_table(schema, table, target.db_type) def format_select_sql( table: str, target: DbTarget, return_fields: Any = None, conditions: Optional[List[Dict[str, Any]]] = None, order_by: str = "", limit: int = 100, ) -> str: """按 dialect 生成 SELECT SQL""" normalized_fields = normalize_return_fields(return_fields) if normalized_fields == "*": field_list = "*" else: field_list = ", ".join( quote_identifier(f, target.db_type) for f in normalized_fields ) full_table = quote_table_for_target(table, target) where_raw = build_where_clause_raw(conditions or [], target.db_type) where_part = f" WHERE {where_raw}" if where_raw else "" sql = f"SELECT {field_list} FROM {full_table}{where_part}" db = (target.db_type or "postgresql").lower() if order_by: sql += f" ORDER BY {order_by}" elif db == "sqlserver": sql += " ORDER BY (SELECT NULL)" if db == "sqlserver": sql += f" OFFSET 0 ROWS FETCH NEXT {int(limit)} ROWS ONLY" elif db == "oracle": sql += f" FETCH FIRST {int(limit)} ROWS ONLY" else: sql += f" LIMIT {int(limit)}" return sql def format_limit_clause(db_type: str, limit: int = 1) -> str: db = (db_type or "postgresql").lower() if db == "sqlserver": return f" ORDER BY (SELECT NULL) OFFSET 0 ROWS FETCH NEXT {int(limit)} ROWS ONLY" if db == "oracle": return f" FETCH FIRST {int(limit)} ROWS ONLY" return f" LIMIT {int(limit)}" def normalize_return_fields(return_fields: Any) -> Any: """将 return_fields 规范化为 * 或字段名列表""" if return_fields in (None, "", "*", ["*"]): return "*" if isinstance(return_fields, str): stripped = return_fields.strip() if not stripped or stripped == "*": return "*" return [field.strip() for field in stripped.split(",") if field.strip()] if isinstance(return_fields, list): if not return_fields or "*" in return_fields: return "*" return return_fields return "*" def merge_result_metadata( base_metadata: Optional[Dict[str, Any]], warnings: Optional[List[str]] = None, ) -> Dict[str, Any]: metadata = dict(base_metadata or {}) if warnings: existing = list(metadata.get("warnings") or []) for code in warnings: if code not in existing: existing.append(code) metadata["warnings"] = existing return metadata def convert_sql_param_type(value: Any, param_type: str) -> Any: """将参数值转换为指定 SQL 参数类型(与数据源 _convert_param_type 对齐)""" if value is None: return None try: if param_type == "integer": return int(value) if param_type == "float": return float(value) if param_type == "boolean": if isinstance(value, bool): return value return str(value).lower() in ("true", "1", "yes") if param_type == "date": if isinstance(value, date): return value if isinstance(value, datetime): return value.date() return date.fromisoformat(str(value).strip()[:10]) if param_type == "datetime": if isinstance(value, datetime): return value return datetime.fromisoformat(str(value).strip()) return str(value) except (ValueError, TypeError): return value def resolve_sql_param_value(raw: Any, context: Any) -> Any: """解析单个 SQL 参数值:模板变量 + JSON 字面量""" if raw is None: return None if isinstance(raw, str): resolved = context.resolve_template(raw) if resolved == "": return None try: return json.loads(resolved) except (json.JSONDecodeError, TypeError): return resolved return raw def build_sql_param_dict( param_defs: Optional[List[Dict[str, Any]]], context: Any, ) -> Dict[str, Any]: """ 从节点 params 配置构建 SQLAlchemy 命名参数字典。 参数定义字段:name, type, value(或 default), required SQL 中使用 :name 占位符,与数据源一致。 """ result: Dict[str, Any] = {} for param in param_defs or []: name = (param.get("name") or "").strip() if not name: continue param_type = param.get("type") or "string" required = bool(param.get("required", False)) raw = param.get("value") if raw is None or raw == "": raw = param.get("default") if raw is None or raw == "": if required: raise ValueError(f"缺少必填 SQL 参数: {name}") result[name] = None continue value = resolve_sql_param_value(raw, context) result[name] = convert_sql_param_type(value, param_type) return result