Files
ai-agent-admin/backend-fastapi/ai_platform/nodes/utils/db_execution.py
T
2026-06-08 18:14:59 +08:00

362 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
工作流数据库节点:连接解析与方言 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