Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,361 @@
|
||||
"""
|
||||
工作流数据库节点:连接解析与方言 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
|
||||
Reference in New Issue
Block a user