362 lines
12 KiB
Python
362 lines
12 KiB
Python
"""
|
||
工作流数据库节点:连接解析与方言 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
|