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

178 lines
6.3 KiB
Python

"""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()