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