Build lightweight AI agent admin
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user