Build lightweight AI agent admin

This commit is contained in:
Codex
2026-06-08 18:14:59 +08:00
commit e164840f43
2530 changed files with 435693 additions and 0 deletions
@@ -0,0 +1 @@
"""AI workflow node utilities."""
@@ -0,0 +1,66 @@
"""
节点配置解析工具
从节点配置或上下文变量中解析 object / list,避免 dict 被 stringify 后无法反序列化。
"""
import ast
import json
import re
from typing import Any, List, Optional
from ai_platform.nodes.base import NodeContext
_VAR_REF_PATTERN = re.compile(r'^\{\{\s*([^}]+)\s*\}\}$')
def _parse_structured_string(value: str) -> Any:
text = value.strip()
if not text:
return None
try:
return json.loads(text)
except (json.JSONDecodeError, TypeError):
pass
try:
return ast.literal_eval(text)
except (ValueError, SyntaxError):
pass
return value
def resolve_config_value(context: NodeContext, raw: Any) -> Any:
"""解析配置值:支持 dict/list 直传、{{var}} 引用、JSON/Python 字面量字符串。"""
if raw is None or raw == '':
return None
if isinstance(raw, (dict, list)):
return raw
if not isinstance(raw, str):
return raw
stripped = raw.strip()
var_match = _VAR_REF_PATTERN.match(stripped)
if var_match:
var_name = var_match.group(1).strip()
if var_name in context.variables:
return context.variables[var_name]
resolved = context.resolve_template(raw)
if isinstance(resolved, (dict, list)):
return resolved
if isinstance(resolved, str):
return _parse_structured_string(resolved)
return resolved
def resolve_object_config(context: NodeContext, raw: Any) -> Optional[dict]:
value = resolve_config_value(context, raw)
return value if isinstance(value, dict) else None
def resolve_list_config(context: NodeContext, raw: Any) -> Optional[List[Any]]:
value = resolve_config_value(context, raw)
return value if isinstance(value, list) else None
@@ -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
@@ -0,0 +1,177 @@
"""Tests for workflow database execution utilities."""
import unittest
from ai_platform.nodes.utils.db_execution import (
DbTarget,
build_sql_param_dict,
build_where_clause_platform,
build_where_clause_raw,
convert_sql_param_type,
default_connection_write_warnings,
format_select_sql,
format_limit_clause,
merge_result_metadata,
normalize_return_fields,
resolve_db_target,
resolve_handler_schema_name,
resolve_schema_for_handler,
resolve_sql_param_value,
)
class DbExecutionTestCase(unittest.TestCase):
def test_resolve_db_target_default(self):
target = resolve_db_target({})
self.assertEqual(target.db_name, "default")
self.assertFalse(target.is_external)
def test_resolve_db_target_external(self):
target = resolve_db_target(
{
"dbName": "erp_mysql",
"dbType": "mysql",
"database": "sales",
"schema": "",
}
)
self.assertEqual(target.db_name, "erp_mysql")
self.assertTrue(target.is_external)
def test_default_connection_write_warnings(self):
target = resolve_db_target({"dbName": "default"})
self.assertEqual(
default_connection_write_warnings("insert", target),
["default_connection_write"],
)
self.assertEqual(default_connection_write_warnings("select", target), [])
def test_build_where_clause_raw_postgresql(self):
where = build_where_clause_raw(
[
{"field": "status", "operator": "=", "value": "active"},
{"field": "age", "operator": ">", "value": 18},
],
"postgresql",
)
self.assertIn('"status" = \'active\'', where)
self.assertIn('"age" > 18', where)
def test_build_where_clause_raw_mysql_like(self):
where = build_where_clause_raw(
[{"field": "name", "operator": "like", "value": "%"}],
"mysql",
)
self.assertIn("`name` LIKE '%'", where)
def test_build_where_clause_platform_named_params(self):
clause, params = build_where_clause_platform(
[{"field": "id", "operator": "=", "value": "1"}]
)
self.assertTrue(clause.startswith("WHERE"))
self.assertEqual(params["p0"], "1")
def test_format_select_sql_postgresql(self):
target = DbTarget(db_name="erp", db_type="postgresql", schema="public")
sql = format_select_sql(
"users",
target,
return_fields=["id", "name"],
conditions=[{"field": "status", "operator": "=", "value": 1}],
limit=10,
)
self.assertTrue(sql.startswith("SELECT"))
self.assertIn('"public"."users"', sql)
self.assertIn("LIMIT 10", sql)
def test_resolve_schema_for_handler_mysql(self):
self.assertEqual(resolve_schema_for_handler("mysql", "", "sales_db"), "sales_db")
def test_merge_result_metadata(self):
metadata = merge_result_metadata({}, ["default_connection_write"])
self.assertEqual(metadata["warnings"], ["default_connection_write"])
def test_normalize_return_fields_comma_string(self):
self.assertEqual(normalize_return_fields("id, name"), ["id", "name"])
self.assertEqual(normalize_return_fields("*"), "*")
def test_format_select_sql_sqlserver_order_by(self):
target = DbTarget(db_name="erp", db_type="sqlserver", schema="dbo")
sql = format_select_sql("users", target, limit=10)
self.assertIn("ORDER BY (SELECT NULL)", sql)
self.assertIn("OFFSET 0 ROWS FETCH NEXT 10 ROWS ONLY", sql)
def test_format_limit_clause_sqlserver(self):
clause = format_limit_clause("sqlserver", 1)
self.assertIn("ORDER BY (SELECT NULL)", clause)
self.assertIn("OFFSET 0 ROWS FETCH NEXT 1 ROWS ONLY", clause)
def test_resolve_handler_schema_mysql_prefers_database(self):
class FakeService:
db_type = "mysql"
def _default_schema(self):
return "fallback"
target = DbTarget(db_name="erp", db_type="mysql", database="sales_db", schema="")
import asyncio
schema = asyncio.run(resolve_handler_schema_name(FakeService(), target))
self.assertEqual(schema, "sales_db")
def test_convert_sql_param_type(self):
self.assertEqual(convert_sql_param_type("42", "integer"), 42)
self.assertEqual(convert_sql_param_type("3.14", "float"), 3.14)
self.assertTrue(convert_sql_param_type("true", "boolean"))
self.assertEqual(convert_sql_param_type("2024-01-15", "date").isoformat(), "2024-01-15")
def test_build_sql_param_dict_basic(self):
class FakeContext:
def resolve_template(self, raw: str) -> str:
return raw.replace("{{user_id}}", "99")
params = build_sql_param_dict(
[
{"name": "status", "type": "string", "value": "active"},
{"name": "limit", "type": "integer", "value": "10"},
{"name": "user_id", "type": "integer", "value": "{{user_id}}"},
],
FakeContext(),
)
self.assertEqual(params["status"], "active")
self.assertEqual(params["limit"], 10)
self.assertEqual(params["user_id"], 99)
def test_build_sql_param_dict_required_missing(self):
class FakeContext:
def resolve_template(self, raw: str) -> str:
return raw
with self.assertRaises(ValueError) as ctx:
build_sql_param_dict(
[{"name": "id", "type": "integer", "required": True}],
FakeContext(),
)
self.assertIn("id", str(ctx.exception))
def test_build_sql_param_dict_uses_default(self):
class FakeContext:
def resolve_template(self, raw: str) -> str:
return raw
params = build_sql_param_dict(
[{"name": "offset", "type": "integer", "default": "0"}],
FakeContext(),
)
self.assertEqual(params["offset"], 0)
def test_resolve_sql_param_value_json(self):
class FakeContext:
def resolve_template(self, raw: str) -> str:
return '["a", "b"]'
value = resolve_sql_param_value("{{ids}}", FakeContext())
self.assertEqual(value, ["a", "b"])
if __name__ == "__main__":
unittest.main()