test_rule_sql.py 7.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242
  1. from __future__ import annotations
  2. from contextlib import contextmanager
  3. import pytest
  4. from app.core.common.identifiers import new_governance_uid
  5. from app.core.data_source.errors import DataSourceWriteOutcomeUnknown
  6. from app.runner.nodes import NodeExecutionError
  7. class Definition:
  8. def __init__(self, dialect):
  9. self.database_type = dialect
  10. self.extra_properties = {
  11. "sql_rule_capabilities": {
  12. "dialect": dialect,
  13. "timezone": "Asia/Shanghai",
  14. "collation": "C" if dialect == "postgresql" else "utf8mb4_0900_bin",
  15. "rounding_mode": "half_away_from_zero",
  16. "regex_engine": "posix" if dialect == "postgresql" else "icu",
  17. }
  18. }
  19. class Definitions:
  20. def __init__(self, definition):
  21. self.definition = definition
  22. def get(self, _uid):
  23. return self.definition
  24. class Result:
  25. def __init__(self, *, scalar=None, rowcount=0):
  26. self._scalar = scalar
  27. self.rowcount = rowcount
  28. def scalar_one(self):
  29. return self._scalar
  30. class Connection:
  31. def __init__(self):
  32. self.calls = []
  33. def execute(self, statement, parameters):
  34. self.calls.append((str(statement), parameters))
  35. if len(self.calls) == 1:
  36. return Result(scalar=3)
  37. if len(self.calls) == 2:
  38. return Result(scalar=2)
  39. return Result(rowcount=2)
  40. class Manager:
  41. def __init__(self, dialect="postgresql", *, unknown_commit=False):
  42. self.definitions = Definitions(Definition(dialect))
  43. self.connection = Connection()
  44. self.unknown_commit = unknown_commit
  45. self.calls = []
  46. @contextmanager
  47. def connect(self, uid, purpose):
  48. self.calls.append((uid, purpose))
  49. yield self.connection
  50. if self.unknown_commit:
  51. raise DataSourceWriteOutcomeUnknown()
  52. def sql_plan(dialect="postgresql"):
  53. from app.core.data_rules.compilers.sql import bound_sql_plan_hash
  54. uid = new_governance_uid()
  55. capabilities = Definition(dialect).extra_properties["sql_rule_capabilities"]
  56. quote = '"' if dialect == "postgresql" else "`"
  57. plan = {
  58. "schema_version": "1.0",
  59. "dialect": dialect,
  60. "capabilities": capabilities,
  61. "data_source_uid": uid,
  62. "rule_version_id": new_governance_uid(),
  63. "input_binding_id": new_governance_uid(),
  64. "output_binding_id": new_governance_uid(),
  65. "statements": [
  66. {
  67. "purpose": "write",
  68. "sql": (
  69. f"INSERT INTO {quote}clean{quote}.{quote}customer{quote} "
  70. f"({quote}id{quote}) SELECT {quote}id{quote} "
  71. f"FROM {quote}raw{quote}.{quote}customer{quote}"
  72. ),
  73. "parameters": {},
  74. }
  75. ],
  76. "result_contract": {
  77. "rows_in": "counted",
  78. "rows_out": "counted",
  79. "rows_rejected": "counted",
  80. },
  81. }
  82. return plan, bound_sql_plan_hash(plan)
  83. def node_for(plan, plan_hash):
  84. return {
  85. "id": "task4_rule",
  86. "type": "rule.apply",
  87. "purpose": "write",
  88. "idempotency": {"strategy": "upsert", "key": "id"},
  89. "config": {
  90. "component_binding_id": new_governance_uid(),
  91. "rule_version_id": plan["rule_version_id"],
  92. "execution_plan_hash": plan_hash,
  93. },
  94. }
  95. def test_sqlglot_rule_adapter_verifies_and_executes_one_transaction():
  96. from app.runner.rule_sql import SqlGlotRulePlanAdapter
  97. plan, plan_hash = sql_plan()
  98. node = node_for(plan, plan_hash)
  99. manager = Manager()
  100. result = SqlGlotRulePlanAdapter(manager).execute(
  101. plan=plan,
  102. node=node,
  103. parameters={},
  104. write_authorized=True,
  105. )
  106. assert result == {
  107. "rows_in": 3,
  108. "rows_out": 2,
  109. "rows_rejected": 1,
  110. "commit_outcome": "committed",
  111. }
  112. assert manager.calls == [(plan["data_source_uid"], "dataflow_write")]
  113. assert len(manager.connection.calls) == 3
  114. def test_sqlglot_rule_adapter_fails_closed_for_dialect_hash_and_authorization():
  115. from app.runner.rule_sql import SqlGlotRulePlanAdapter
  116. plan, plan_hash = sql_plan()
  117. node = node_for(plan, plan_hash)
  118. with pytest.raises(NodeExecutionError, match="dialect"):
  119. SqlGlotRulePlanAdapter(Manager("mysql")).execute(
  120. plan=plan,
  121. node=node,
  122. parameters={},
  123. write_authorized=True,
  124. )
  125. node["config"]["execution_plan_hash"] = "0" * 64
  126. with pytest.raises(NodeExecutionError, match="hash"):
  127. SqlGlotRulePlanAdapter(Manager()).execute(
  128. plan=plan,
  129. node=node,
  130. parameters={},
  131. write_authorized=True,
  132. )
  133. node["config"]["execution_plan_hash"] = plan_hash
  134. with pytest.raises(NodeExecutionError, match="authorization"):
  135. SqlGlotRulePlanAdapter(Manager()).execute(
  136. plan=plan,
  137. node=node,
  138. parameters={},
  139. write_authorized=False,
  140. )
  141. def test_sqlglot_rule_adapter_rejects_unimplemented_idempotency_strategy():
  142. from app.runner.rule_sql import SqlGlotRulePlanAdapter
  143. plan, plan_hash = sql_plan()
  144. node = node_for(plan, plan_hash)
  145. node["idempotency"] = {
  146. "strategy": "partition_replace",
  147. "key": "task4",
  148. }
  149. with pytest.raises(NodeExecutionError, match="idempotency"):
  150. SqlGlotRulePlanAdapter(Manager()).execute(
  151. plan=plan,
  152. node=node,
  153. parameters={},
  154. write_authorized=True,
  155. )
  156. def test_rule_executor_matches_node_idempotency_to_persisted_component():
  157. from app.runner.rules import RulePlanExecutor
  158. plan, plan_hash = sql_plan()
  159. node = node_for(plan, plan_hash)
  160. class Repository:
  161. def load(self, **_kwargs):
  162. return {
  163. "component_binding_id": node["config"]["component_binding_id"],
  164. "rule_version_id": plan["rule_version_id"],
  165. "backend": "sql_pushdown",
  166. "plan": plan,
  167. "plan_hash": plan_hash,
  168. "plan_status": "published",
  169. "rule_status": "published",
  170. "component_kind": "rule.apply",
  171. "binding_idempotency": {
  172. "strategy": "upsert",
  173. "key": "different_id",
  174. },
  175. }
  176. class Adapter:
  177. def execute(self, **_kwargs):
  178. return {"rows_in": 0, "rows_out": 0, "rows_rejected": 0}
  179. with pytest.raises(NodeExecutionError, match="idempotency"):
  180. RulePlanExecutor(
  181. Repository(),
  182. adapters={"sql_pushdown": Adapter()},
  183. ).execute(node, {}, write_authorized=True)
  184. def test_sqlglot_rule_adapter_reports_unknown_commit_outcome():
  185. from app.runner.rule_sql import SqlGlotRulePlanAdapter
  186. plan, plan_hash = sql_plan()
  187. with pytest.raises(NodeExecutionError) as error:
  188. SqlGlotRulePlanAdapter(Manager(unknown_commit=True)).execute(
  189. plan=plan,
  190. node=node_for(plan, plan_hash),
  191. parameters={},
  192. write_authorized=True,
  193. )
  194. assert error.value.commit_outcome == "unknown"