178 lines
6.3 KiB
Python
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()
|