Quellcode durchsuchen

fix: enforce expression backend capabilities

马小龙 vor 4 Wochen
Ursprung
Commit
509f66db97

+ 35 - 2
app/core/data_rules/compiler.py

@@ -8,6 +8,7 @@ from typing import Any
 
 from app.core.common.identifiers import ensure_governance_uid
 from app.core.data_rules.contracts import read_rule_spec, rule_spec_hash
+from app.core.data_rules.expressions import SUPPORTED_BACKENDS, backend_support
 
 
 COMPILER_VERSION = "dataops-rulespec-1.0"
@@ -23,7 +24,18 @@ def _canonical_hash(value: Any) -> str:
     return hashlib.sha256(canonical.encode("utf-8")).hexdigest()
 
 
-def compile_rule_plan(rule_version: dict[str, Any]) -> dict[str, Any]:
+def _rule_supported_backends(spec: dict[str, Any]) -> frozenset[str]:
+    supported = SUPPORTED_BACKENDS
+    for step in spec["steps"]:
+        ast = step.get("expression_ast")
+        if ast is not None:
+            supported = supported & backend_support(ast)
+    return frozenset(supported)
+
+
+def compile_rule_plan(
+    rule_version: dict[str, Any], *, supported_backends: frozenset[str] | None = None
+) -> dict[str, Any]:
     """Compile a published rule into a target-neutral immutable plan.
 
     This compiler deliberately emits no SQL or Python. M3 runtime adapters
@@ -45,8 +57,28 @@ def compile_rule_plan(rule_version: dict[str, Any]) -> dict[str, Any]:
     if rule_version.get("spec_hash") != digest:
         raise ValueError("published rule version spec hash does not match")
 
+    derived_backends = _rule_supported_backends(spec)
+    if supported_backends is None:
+        supported = derived_backends
+    elif (
+        not isinstance(supported_backends, frozenset)
+        or not supported_backends <= SUPPORTED_BACKENDS
+    ):
+        raise ValueError("supported_backends must be a concrete backend set")
+    else:
+        supported = supported_backends & derived_backends
+    if not supported:
+        raise ValueError("rule has no supported execution backend")
+
     quality_only = all(step["op"] == "assert" for step in spec["steps"])
-    backend = "quality_check" if quality_only else "polars_batch"
+    if quality_only:
+        backend = "quality_check"
+    elif "polars" in supported:
+        backend = "polars_batch"
+    elif supported & {"postgresql", "mysql"}:
+        backend = "sql_pushdown"
+    else:
+        raise ValueError("rule has no supported compiler target")
     plan = {
         "schema_version": "1.0",
         "compiler_version": COMPILER_VERSION,
@@ -56,6 +88,7 @@ def compile_rule_plan(rule_version: dict[str, Any]) -> dict[str, Any]:
         "output_schema_ref": spec["output_schema_ref"],
         "null_policy": spec["null_policy"],
         "timezone": spec["timezone"],
+        "supported_execution_backends": sorted(supported),
         "steps": spec["steps"],
     }
     return {

+ 9 - 12
app/core/data_rules/expressions.py

@@ -37,8 +37,10 @@ ALLOWED_FUNCTIONS = frozenset(
 SUPPORTED_BACKENDS = frozenset({"postgresql", "mysql", "polars"})
 _FUNCTION_BACKENDS = {
     "matches": SUPPORTED_BACKENDS,
-    "lower": SUPPORTED_BACKENDS,
-    "upper": SUPPORTED_BACKENDS,
+    # The reference oracle verifies Python/Polars Unicode case folding. SQL
+    # engines require a deployment-pinned collation before they may opt in.
+    "lower": frozenset({"polars"}),
+    "upper": frozenset({"polars"}),
     "trim": SUPPORTED_BACKENDS,
     "length": SUPPORTED_BACKENDS,
     "coalesce": SUPPORTED_BACKENDS,
@@ -565,7 +567,7 @@ def type_check_expression(ast: dict, fields: dict[str, str]) -> str:
 
 def validate_rule_expressions(
     rule_spec: dict[str, Any], fields: dict[str, str]
-) -> dict[str, dict[str, Any]]:
+) -> frozenset[str]:
     """Type-check every canonical expression in a schema-bound RuleSpec.
 
     Callers must supply fields from a server-owned schema snapshot, never an
@@ -577,8 +579,8 @@ def validate_rule_expressions(
         rule_spec.get("steps"), list
     ):
         raise ValueError("rule spec steps are required for expression validation")
-    result: dict[str, dict[str, Any]] = {}
-    for index, step in enumerate(rule_spec["steps"]):
+    supported_for_rule = SUPPORTED_BACKENDS
+    for step in rule_spec["steps"]:
         if not isinstance(step, dict):
             raise ValueError("rule step must be an object")
         operation = step.get("op")
@@ -593,13 +595,8 @@ def validate_rule_expressions(
         supported = backend_support(ast)
         if not supported:
             raise ValueError("expression has no supported execution backend")
-        step_id = step.get("id")
-        key = step_id if isinstance(step_id, str) else str(index)
-        result[key] = {
-            "result_type": expression_type,
-            "backends": supported,
-        }
-    return result
+        supported_for_rule = supported_for_rule & supported
+    return frozenset(supported_for_rule)
 
 
 def _portable_regex(pattern: str) -> bool:

+ 52 - 10
app/core/data_rules/release.py

@@ -59,6 +59,52 @@ class ProductionLineReleaseService:
         output = output_snapshot["schema_hash"]
 
         standards, rules = self.repository.load_published_assets(flow)
+        rule_backend_support: dict[str, frozenset[str]] = {}
+
+        def validate_published_rule(rule_version_id: str) -> None:
+            if rule_version_id in rule_backend_support:
+                return
+            rule = rules.get(rule_version_id)
+            if not isinstance(rule, dict):
+                raise ValueError(
+                    f"published rule version {rule_version_id} was not found"
+                )
+            rule_spec = read_rule_spec(rule.get("rule_spec"))
+            rule_snapshot = self.schema_resolver.resolve(
+                rule_spec["input_schema_ref"]
+            )
+            snapshot_fields = {
+                field["name"]: field["type"]
+                for field in rule_snapshot["fields"]
+            }
+            supported = validate_rule_expressions(
+                rule_spec, snapshot_fields
+            )
+            if not supported:
+                raise ValueError("rule has no supported execution backend")
+            rule_backend_support[rule_version_id] = supported
+
+        # Validate every loaded rule before allocating any release version.
+        # The repository is expected to scope this mapping to published assets
+        # available to the release, so no loaded expression is left unchecked.
+        for loaded_rule_version_id in rules:
+            validate_published_rule(loaded_rule_version_id)
+
+        # Ensure referenced standards and their clauses resolve to loaded,
+        # already-validated rule versions before beginning the release.
+        for component in flow["components"]:
+            if component["type"] == "standard.enforce":
+                standard = standards.get(component["standard_version_id"])
+                if not isinstance(standard, dict):
+                    raise ValueError(
+                        "published standard version "
+                        f"{component['standard_version_id']} was not found"
+                    )
+                for clause in standard["clauses"]:
+                    validate_published_rule(str(clause["rule_version_id"]))
+            else:
+                validate_published_rule(component["rule_version_id"])
+
         version = self.repository.begin_dataflow_release(
             dataflow_spec=flow,
             source_text=source,
@@ -86,16 +132,12 @@ class ProductionLineReleaseService:
                 raise ValueError(
                     f"published rule version {rule_version_id} was not found"
                 )
-            rule_spec = read_rule_spec(rule.get("rule_spec"))
-            rule_snapshot = self.schema_resolver.resolve(
-                rule_spec["input_schema_ref"]
-            )
-            snapshot_fields = {
-                field["name"]: field["type"]
-                for field in rule_snapshot["fields"]
-            }
-            validate_rule_expressions(rule_spec, snapshot_fields)
-            plan = compiled.setdefault(rule_version_id, compile_rule_plan(rule))
+            if rule_version_id not in compiled:
+                compiled[rule_version_id] = compile_rule_plan(
+                    rule,
+                    supported_backends=rule_backend_support[rule_version_id],
+                )
+            plan = compiled[rule_version_id]
             binding_id = new_governance_uid()
             binding_ids[binding_key] = binding_id
             self.repository.persist_component_plan(

+ 27 - 1
tests/core/data_rules/test_expressions.py

@@ -128,7 +128,7 @@ def test_decimal_literals_preserve_their_original_lexeme():
             {"display_name": "string"},
             "boolean",
             [("STRAßE", True)],
-            frozenset({"postgresql", "mysql", "polars"}),
+            frozenset({"polars"}),
         ),
     ],
 )
@@ -177,6 +177,32 @@ def test_backend_support_is_conservative_for_rounding_and_nonportable_regex():
     assert backend_support(parse_expression("matches(mobile, '(?<=0)[0-9]+')")) == frozenset()
 
 
+def test_rule_expression_capability_is_the_intersection_of_every_expression():
+    from app.core.data_rules.expressions import (
+        parse_expression,
+        validate_rule_expressions,
+    )
+
+    rule_spec = {
+        "steps": [
+            {
+                "id": "unicode_name",
+                "op": "assert",
+                "expression_ast": parse_expression("lower(name) == 'straße'"),
+            },
+            {
+                "id": "rounded_amount",
+                "op": "assert",
+                "expression_ast": parse_expression("round(amount, 2) >= 1.00"),
+            },
+        ]
+    }
+
+    assert validate_rule_expressions(
+        rule_spec, {"name": "string", "amount": "decimal"}
+    ) == frozenset()
+
+
 @pytest.mark.parametrize(
     ("source", "fields", "message"),
     [

+ 107 - 0
tests/core/data_rules/test_release.py

@@ -96,6 +96,39 @@ def test_rulespec_compiler_is_deterministic_and_never_emits_source_code():
         _published_rule(new_governance_uid(), assertion_only_rule())
     )
     assert quality["backend"] == "quality_check"
+    assert quality["plan"]["supported_execution_backends"] == [
+        "mysql",
+        "polars",
+        "postgresql",
+    ]
+
+
+def test_rulespec_compiler_never_selects_polars_without_polars_capability():
+    import pytest
+
+    from app.core.data_rules.compiler import compile_rule_plan
+
+    spec = valid_rule_spec()
+    spec["steps"] = [
+        {
+            "id": "rounded_balance",
+            "op": "derive",
+            "expression": "round(balance, 2)",
+        }
+    ]
+
+    compiled = compile_rule_plan(_published_rule(new_governance_uid(), spec))
+
+    assert compiled["backend"] == "sql_pushdown"
+    assert compiled["plan"]["supported_execution_backends"] == [
+        "mysql",
+        "postgresql",
+    ]
+    with pytest.raises(ValueError, match="no supported"):
+        compile_rule_plan(
+            _published_rule(new_governance_uid(), spec),
+            supported_backends=frozenset(),
+        )
 
 
 def test_release_expands_standard_and_persists_fixed_bindings_and_plans():
@@ -218,6 +251,80 @@ def test_release_type_checks_persisted_rule_expressions_against_server_schema(
             created_by=new_governance_uid(),
         )
 
+    assert "begin_dataflow_release" not in [
+        method for method, _kwargs in repository.calls
+    ]
+
+
+def test_release_selects_an_sql_target_when_rounding_cannot_run_on_polars():
+    from app.core.data_rules.release import ProductionLineReleaseService
+
+    rule_id = new_governance_uid()
+    spec = valid_rule_spec()
+    spec["steps"] = [
+        {
+            "id": "rounded_balance",
+            "op": "derive",
+            "expression": "round(balance, 2)",
+        }
+    ]
+    repository = ReleaseRepository(
+        standards={}, rules={rule_id: _published_rule(rule_id, spec)}
+    )
+    flow = valid_dataflow_spec(rule_version_id=rule_id)
+    flow["components"] = [flow["components"][0]]
+
+    ProductionLineReleaseService(
+        repository, schema_resolver=FakeSchemaResolver()
+    ).release(
+        dataflow_uid=flow["dataflow_uid"],
+        dataflow_spec=flow,
+        source_text="四舍五入需要安全的 SQL 目标",
+        created_by=new_governance_uid(),
+    )
+
+    plan = next(
+        kwargs["plan"]
+        for method, kwargs in repository.calls
+        if method == "persist_component_plan"
+    )
+    assert plan["backend"] == "sql_pushdown"
+    assert plan["plan"]["supported_execution_backends"] == ["mysql", "postgresql"]
+
+
+def test_release_validates_every_loaded_rule_before_beginning_a_release():
+    from app.core.data_rules.release import ProductionLineReleaseService
+
+    direct_rule_id = new_governance_uid()
+    unrelated_invalid_rule_id = new_governance_uid()
+    invalid_spec = valid_rule_spec()
+    invalid_spec["steps"][1]["expression"] = "missing_schema_field == 'x'"
+    repository = ReleaseRepository(
+        standards={},
+        rules={
+            direct_rule_id: _published_rule(direct_rule_id, valid_rule_spec()),
+            unrelated_invalid_rule_id: _published_rule(
+                unrelated_invalid_rule_id, invalid_spec
+            ),
+        },
+    )
+    flow = valid_dataflow_spec(rule_version_id=direct_rule_id)
+    flow["components"] = [flow["components"][0]]
+
+    with pytest.raises(ValueError, match="unknown expression field"):
+        ProductionLineReleaseService(
+            repository, schema_resolver=FakeSchemaResolver()
+        ).release(
+            dataflow_uid=flow["dataflow_uid"],
+            dataflow_spec=flow,
+            source_text="所有已加载规则必须先通过验证",
+            created_by=new_governance_uid(),
+        )
+
+    assert "begin_dataflow_release" not in [
+        method for method, _kwargs in repository.calls
+    ]
+
 
 def test_release_rejects_a_persisted_unknown_expression_function_before_compile():
     from app.core.data_rules.release import ProductionLineReleaseService