feat: add extensible operators and golden flows

This commit is contained in:
2026-07-29 15:50:15 +02:00
parent 09c98087c5
commit 946202ef01
27 changed files with 3584 additions and 218 deletions
+539 -79
View File
@@ -8,7 +8,17 @@ from dataclasses import dataclass
from decimal import Decimal
from typing import Any, Callable
from govoplan_dataflow.backend.expressions import (
convert_value,
evaluate_expression,
parse_expression,
)
from govoplan_dataflow.backend.graph import graph_inputs_by_port, topological_order, validate_graph
from govoplan_dataflow.backend.operator_registry import (
OPERATOR_REGISTRY,
OperatorExecutionContext,
OperatorExecutionResult,
)
from govoplan_dataflow.backend.schemas import (
DataflowDiagnostic,
GraphNode,
@@ -17,6 +27,7 @@ from govoplan_dataflow.backend.schemas import (
PipelineGraph,
PreviewColumn,
)
from govoplan_dataflow.backend.subflows import substitute_parameters
EXECUTOR_VERSION = "dataflow-preview-v2"
@@ -79,7 +90,10 @@ def execute_preview(
row_limit: int,
source_resolver: SourceResolver | None = None,
preview_node_id: str | None = None,
_execution_depth: int = 0,
) -> PipelineExecutionResult:
if _execution_depth > 5:
raise PipelineExecutionError("Subflows are limited to five nested levels.")
validation = validate_graph(graph)
errors = [item for item in validation if item.severity == "error"]
if errors:
@@ -133,87 +147,37 @@ def execute_preview(
node_preview=_node_preview(outputs, preview_node_id, row_limit),
)
try:
if node.type == "source.inline":
output_rows = [dict(row) for row in node.config.get("rows", [])]
input_row_count += len(output_rows)
source_fingerprints.append(
{
"node_id": node.id,
"source_name": node.config.get("source_name"),
"kind": "inline",
"fingerprint": _rows_fingerprint(output_rows),
"row_count": len(output_rows),
}
)
elif node.type == "source.reference":
if source_resolver is None:
raise PipelineExecutionError(
"Datasource-backed preview requires the Datasources catalogue capability.",
node_id=node.id,
)
resolved = source_resolver(node, MAX_SOURCE_ROWS)
output_rows = [dict(row) for row in resolved.rows]
input_row_count += len(output_rows)
source_fingerprints.append(
{
"node_id": node.id,
"source_ref": resolved.source_ref,
"source_name": node.config.get("source_name"),
"kind": "datasource",
"provider": resolved.provider,
"fingerprint": resolved.fingerprint,
"row_count": resolved.total_rows,
"preview_rows": len(output_rows),
"truncated": resolved.truncated,
}
)
if resolved.truncated:
message = (
f"Source preview used {len(output_rows):,} of "
f"{resolved.total_rows:,} rows."
)
node_messages.append(message)
validation.append(
DataflowDiagnostic(
severity="warning",
code="source.preview_truncated",
message=message,
node_id=node.id,
)
)
elif node.type == "combine.union":
ordered_inputs = [
outputs[source_id]
for source_id in node_inputs.get("input", [])
]
output_rows = _union_rows(ordered_inputs, node.config)
elif node.type == "combine.join":
left_rows = outputs[node_inputs["left"][0]]
right_rows = outputs[node_inputs["right"][0]]
output_rows = _join_rows(
left_rows,
right_rows,
node.config,
executor = OPERATOR_REGISTRY.executor(node.type)
if executor is None:
raise PipelineExecutionError(
f"Node type {node.type!r} has no registered executor.",
node_id=node.id,
)
elif node.type == "filter":
output_rows = _filter_rows(input_rows, node.config, node_id=node.id)
elif node.type == "distinct":
output_rows = _distinct_rows(input_rows, node.config)
elif node.type == "select":
output_rows = _select_rows(input_rows, node.config)
elif node.type == "derive":
output_rows = _derive_rows(input_rows, node.config, node_id=node.id)
elif node.type == "aggregate":
output_rows = _aggregate_rows(input_rows, node.config, node_id=node.id)
elif node.type == "sort":
output_rows = _sort_rows(input_rows, node.config)
elif node.type == "limit":
output_rows = input_rows[: int(node.config["count"])]
elif node.type == "output":
output_rows = [dict(row) for row in input_rows]
else:
raise PipelineExecutionError(f"Unsupported node type: {node.type}", node_id=node.id)
execution = executor(
OperatorExecutionContext(
node=node,
inputs_by_port=node_inputs,
outputs=outputs,
input_sets=input_sets,
input_rows=input_rows,
source_resolver=source_resolver,
execution_depth=_execution_depth,
)
)
output_rows = execution.rows
input_row_count += execution.input_row_count
source_fingerprints.extend(execution.source_fingerprints)
node_messages.extend(execution.messages)
if node.type == "source.reference" and execution.messages:
validation.extend(
DataflowDiagnostic(
severity="warning",
code="source.preview_truncated",
message=message,
node_id=node.id,
)
for message in execution.messages
)
if len(json.dumps(output_rows, default=str).encode("utf-8")) > MAX_RESULT_BYTES:
raise PipelineExecutionError(
"A preview node exceeded the one-megabyte result limit.",
@@ -461,6 +425,309 @@ def _derive_value(operation: str, values: list[Any], *, separator: str) -> Any:
raise ValueError(f"unknown derive operation {operation!r}")
def _expression_filter_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
*,
node_id: str,
) -> list[dict[str, Any]]:
parsed = parse_expression(str(config["expression"]))
result: list[dict[str, Any]] = []
for row in rows:
try:
if bool(evaluate_expression(parsed, row)):
result.append(dict(row))
except (ArithmeticError, TypeError, ValueError) as exc:
raise PipelineExecutionError(
f"Cannot evaluate filter expression: {exc}",
node_id=node_id,
) from exc
return result
def _expression_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
*,
node_id: str,
) -> list[dict[str, Any]]:
target = str(config["target_column"])
parsed = parse_expression(str(config["expression"]))
output: list[dict[str, Any]] = []
for row in rows:
item = dict(row)
try:
item[target] = evaluate_expression(parsed, row)
except (ArithmeticError, TypeError, ValueError) as exc:
raise PipelineExecutionError(
f"Cannot evaluate expression for {target!r}: {exc}",
node_id=node_id,
) from exc
output.append(item)
return output
def _convert_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
*,
node_id: str,
) -> list[dict[str, Any]]:
source = str(config["source_column"])
target = str(config["target_column"])
target_type = str(config["target_type"])
on_error = str(config.get("on_error", "fail"))
output: list[dict[str, Any]] = []
for row in rows:
item = dict(row)
try:
item[target] = convert_value(
row.get(source),
target_type,
on_error=on_error, # type: ignore[arg-type]
)
except (ArithmeticError, TypeError, ValueError) as exc:
raise PipelineExecutionError(
f"Cannot convert {source!r} to {target_type}: {exc}",
node_id=node_id,
) from exc
output.append(item)
return output
def _replace_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
) -> list[dict[str, Any]]:
source = str(config["source_column"])
target = str(config["target_column"])
mode = str(config.get("mode", "exact"))
find = config.get("find")
replacement = config.get("replacement")
output: list[dict[str, Any]] = []
for row in rows:
item = dict(row)
value = row.get(source)
if mode == "text" and value is not None:
item[target] = str(value).replace(str(find), str(replacement))
else:
item[target] = replacement if value == find else value
output.append(item)
return output
def _quality_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
*,
node_id: str,
) -> list[dict[str, Any]]:
rules = config.get("rules", [])
action = str(config.get("action", "annotate"))
unique_values: dict[str, set[object]] = defaultdict(set)
output: list[dict[str, Any]] = []
invalid_count = 0
for row_index, row in enumerate(rows, start=1):
failures: list[str] = []
for rule in rules:
if not isinstance(rule, dict):
continue
rule_id = str(rule.get("id") or rule.get("operator") or "rule")
column = str(rule.get("column") or "")
value = row.get(column)
operator = str(rule.get("operator") or "")
valid = _quality_rule_matches(
value,
operator=operator,
rule=rule,
seen=unique_values[rule_id],
)
if not valid:
failures.append(rule_id)
if failures:
invalid_count += 1
if action == "fail":
raise PipelineExecutionError(
f"Quality rules failed for row {row_index}: {', '.join(failures)}.",
node_id=node_id,
)
if action == "drop":
continue
item = dict(row)
if action == "annotate":
item["_quality_valid"] = not failures
item["_quality_errors"] = failures
output.append(item)
if invalid_count and action not in {"annotate", "drop", "fail"}:
raise PipelineExecutionError(
f"Unsupported quality action {action!r}.",
node_id=node_id,
)
return output
def _quality_rule_matches(
value: Any,
*,
operator: str,
rule: dict[str, Any],
seen: set[object],
) -> bool:
if operator == "not_null":
return value is not None and value != ""
if operator == "type":
expected = str(rule.get("value") or "")
mapping = {
"string": str,
"integer": int,
"number": (int, float, Decimal),
"boolean": bool,
}
expected_type = mapping.get(expected)
return expected_type is not None and isinstance(value, expected_type)
if operator == "min":
return value is not None and value >= rule.get("value")
if operator == "max":
return value is not None and value <= rule.get("value")
if operator == "allowed":
allowed = rule.get("values")
return isinstance(allowed, list) and value in allowed
if operator == "unique":
marker = _hashable(value)
if marker in seen:
return False
seen.add(marker)
return True
return False
def _reconcile_rows(
left_rows: list[dict[str, Any]],
right_rows: list[dict[str, Any]],
config: dict[str, Any],
*,
node_id: str,
) -> list[dict[str, Any]]:
left_keys = [str(item) for item in config["left_keys"]]
right_keys = [str(item) for item in config["right_keys"]]
right_prefix = str(config.get("right_prefix", "observed_"))
comparisons = _comparison_fields(config.get("compare_columns"))
left_index = _unique_row_index(left_rows, left_keys, node_id=node_id)
right_index = _unique_row_index(right_rows, right_keys, node_id=node_id)
output: list[dict[str, Any]] = []
for key in dict.fromkeys((*left_index, *right_index)):
left = left_index.get(key)
right = right_index.get(key)
item = dict(left or {})
if right is not None:
item.update({f"{right_prefix}{name}": value for name, value in right.items()})
if left is None:
status = "missing_expected"
differences: list[str] = []
elif right is None:
status = "missing_observed"
differences = []
else:
fields = comparisons or tuple((name, name) for name in left if name not in left_keys)
differences = [
left_name
for left_name, right_name in fields
if left.get(left_name) != right.get(right_name)
]
status = "changed" if differences else "match"
item["_reconciliation_status"] = status
item["_reconciliation_differences"] = differences
output.append(item)
return output
def _comparison_fields(value: object) -> tuple[tuple[str, str], ...]:
if not isinstance(value, list):
return ()
fields: list[tuple[str, str]] = []
for item in value:
if isinstance(item, str) and item:
fields.append((item, item))
elif isinstance(item, dict):
left = str(item.get("left") or item.get("column") or "")
right = str(item.get("right") or left)
if left and right:
fields.append((left, right))
return tuple(fields)
def _unique_row_index(
rows: list[dict[str, Any]],
columns: list[str],
*,
node_id: str,
) -> dict[tuple[Any, ...], dict[str, Any]]:
index: dict[tuple[Any, ...], dict[str, Any]] = {}
for row in rows:
key = tuple(_hashable(row.get(column)) for column in columns)
if key in index:
raise PipelineExecutionError(
f"Reconciliation keys are not unique: {key!r}.",
node_id=node_id,
)
index[key] = row
return index
def _execute_subflow(
rows: list[dict[str, Any]],
config: dict[str, Any],
*,
source_resolver: SourceResolver | None,
execution_depth: int,
) -> PipelineExecutionResult:
parameters = config.get("parameters")
graph_payload = substitute_parameters(
config.get("graph"),
parameters if isinstance(parameters, dict) else {},
)
graph = PipelineGraph.model_validate(graph_payload)
input_nodes = [
node
for node in graph.nodes
if node.type == "source.inline" and node.config.get("input_binding") is True
]
if len(input_nodes) != 1:
raise PipelineExecutionError(
"A reusable subflow needs exactly one inline source with input_binding=true."
)
input_node = input_nodes[0]
graph = graph.model_copy(
update={
"nodes": [
(
node.model_copy(
update={
"config": {
**node.config,
"rows": [dict(row) for row in rows],
}
},
deep=True,
)
if node.id == input_node.id
else node
)
for node in graph.nodes
]
},
deep=True,
)
return execute_preview(
graph,
row_limit=MAX_INTERMEDIATE_ROWS,
source_resolver=source_resolver,
_execution_depth=execution_depth + 1,
)
def _join_key(row: dict[str, Any], columns: list[str]) -> tuple[Any, ...] | None:
values = tuple(row.get(column) for column in columns)
if any(value is None for value in values):
@@ -634,6 +901,12 @@ def _sortable_value(value: Any) -> tuple[str, Any]:
return type(value).__name__, str(value)
def _hashable(value: Any) -> Any:
if isinstance(value, (dict, list, tuple, set)):
return json.dumps(value, sort_keys=True, default=str)
return value
def _rows_fingerprint(rows: list[dict[str, Any]]) -> str:
encoded = json.dumps(rows, sort_keys=True, separators=(",", ":"), default=str)
return hashlib.sha256(encoded.encode("utf-8")).hexdigest()
@@ -657,6 +930,193 @@ def _type_name(value: Any) -> str:
return type(value).__name__.lower()
def _execute_inline_source(
context: OperatorExecutionContext,
) -> OperatorExecutionResult:
rows = [dict(row) for row in context.node.config.get("rows", [])]
return OperatorExecutionResult(
rows=rows,
input_row_count=len(rows),
source_fingerprints=(
{
"node_id": context.node.id,
"source_name": context.node.config.get("source_name"),
"kind": "inline",
"fingerprint": _rows_fingerprint(rows),
"row_count": len(rows),
},
),
)
def _execute_reference_source(
context: OperatorExecutionContext,
) -> OperatorExecutionResult:
if context.source_resolver is None:
raise PipelineExecutionError(
"Datasource-backed preview requires the Datasources catalogue capability.",
node_id=context.node.id,
)
resolved = context.source_resolver(context.node, MAX_SOURCE_ROWS)
rows = [dict(row) for row in resolved.rows]
messages = (
(
f"Source preview used {len(rows):,} of "
f"{resolved.total_rows:,} rows."
),
) if resolved.truncated else ()
return OperatorExecutionResult(
rows=rows,
messages=messages,
input_row_count=len(rows),
source_fingerprints=(
{
"node_id": context.node.id,
"source_ref": resolved.source_ref,
"source_name": context.node.config.get("source_name"),
"kind": "datasource",
"provider": resolved.provider,
"fingerprint": resolved.fingerprint,
"row_count": resolved.total_rows,
"preview_rows": len(rows),
"truncated": resolved.truncated,
},
),
)
def _execute_subflow_node(
context: OperatorExecutionContext,
) -> OperatorExecutionResult:
nested = _execute_subflow(
context.input_rows,
context.node.config,
source_resolver=context.source_resolver,
execution_depth=context.execution_depth,
)
return OperatorExecutionResult(
rows=nested.rows,
messages=tuple(
f"Subflow {item.node_id}: {message}"
for item in nested.node_diagnostics
for message in item.messages
),
source_fingerprints=tuple(
{
**fingerprint,
"subflow_node_id": context.node.id,
}
for fingerprint in nested.source_fingerprints
),
)
def _register_executors() -> None:
executors = {
"source.inline": _execute_inline_source,
"source.reference": _execute_reference_source,
"combine.union": lambda context: OperatorExecutionResult(
rows=_union_rows(
[
context.outputs[source_id]
for source_id in context.inputs_by_port.get("input", ())
],
context.node.config,
)
),
"combine.join": lambda context: OperatorExecutionResult(
rows=_join_rows(
context.outputs[context.inputs_by_port["left"][0]],
context.outputs[context.inputs_by_port["right"][0]],
context.node.config,
node_id=context.node.id,
)
),
"filter": lambda context: OperatorExecutionResult(
rows=_filter_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"filter.expression": lambda context: OperatorExecutionResult(
rows=_expression_filter_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"distinct": lambda context: OperatorExecutionResult(
rows=_distinct_rows(context.input_rows, context.node.config)
),
"select": lambda context: OperatorExecutionResult(
rows=_select_rows(context.input_rows, context.node.config)
),
"derive": lambda context: OperatorExecutionResult(
rows=_derive_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"expression": lambda context: OperatorExecutionResult(
rows=_expression_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"convert": lambda context: OperatorExecutionResult(
rows=_convert_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"replace": lambda context: OperatorExecutionResult(
rows=_replace_rows(context.input_rows, context.node.config)
),
"aggregate": lambda context: OperatorExecutionResult(
rows=_aggregate_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"sort": lambda context: OperatorExecutionResult(
rows=_sort_rows(context.input_rows, context.node.config)
),
"limit": lambda context: OperatorExecutionResult(
rows=context.input_rows[: int(context.node.config["count"])]
),
"quality.rules": lambda context: OperatorExecutionResult(
rows=_quality_rows(
context.input_rows,
context.node.config,
node_id=context.node.id,
)
),
"reconcile.compare": lambda context: OperatorExecutionResult(
rows=_reconcile_rows(
context.outputs[context.inputs_by_port["left"][0]],
context.outputs[context.inputs_by_port["right"][0]],
context.node.config,
node_id=context.node.id,
)
),
"subflow": _execute_subflow_node,
"output": lambda context: OperatorExecutionResult(
rows=[dict(row) for row in context.input_rows]
),
}
for node_type, executor in executors.items():
if OPERATOR_REGISTRY.executor(node_type) is None:
OPERATOR_REGISTRY.register_executor(node_type, executor)
_register_executors()
__all__ = [
"EXECUTOR_VERSION",
"MAX_SOURCE_ROWS",
@@ -0,0 +1,426 @@
from __future__ import annotations
from dataclasses import dataclass
from decimal import Decimal
from typing import Any, Literal
import sqlglot
from sqlglot import exp
from sqlglot.errors import ParseError
ExpressionDataType = Literal[
"unknown",
"null",
"string",
"integer",
"number",
"boolean",
"date",
"datetime",
]
class ExpressionError(ValueError):
pass
@dataclass(frozen=True, slots=True)
class ParsedExpression:
source: str
expression: exp.Expression
columns: tuple[str, ...]
def sql(self) -> str:
return self.expression.sql(dialect="duckdb")
_LEAF_TYPES = (
exp.Column,
exp.Identifier,
exp.Literal,
exp.Null,
exp.Boolean,
exp.DataType,
)
_BINARY_TYPES = (
exp.Add,
exp.Sub,
exp.Mul,
exp.Div,
exp.Mod,
exp.EQ,
exp.NEQ,
exp.GT,
exp.GTE,
exp.LT,
exp.LTE,
exp.And,
exp.Or,
exp.Is,
)
_UNARY_TYPES = (exp.Not, exp.Neg, exp.Paren)
_FUNCTION_TYPES = (
exp.Lower,
exp.Upper,
exp.Trim,
exp.Coalesce,
exp.Concat,
exp.Length,
exp.Abs,
exp.Round,
exp.Replace,
exp.Substring,
exp.Cast,
)
_CONTROL_TYPES = (exp.Case, exp.If)
_ALLOWED_TYPES = (
*_LEAF_TYPES,
*_BINARY_TYPES,
*_UNARY_TYPES,
*_FUNCTION_TYPES,
*_CONTROL_TYPES,
)
def parse_expression(source: str) -> ParsedExpression:
normalized = str(source or "").strip()
if not normalized:
raise ExpressionError("Enter an expression.")
if len(normalized) > 10_000:
raise ExpressionError("Expressions are limited to 10,000 characters.")
try:
expression = sqlglot.parse_one(normalized, read="duckdb")
except ParseError as exc:
raise ExpressionError(f"Expression could not be parsed: {exc}") from exc
for item in expression.walk():
if not isinstance(item, _ALLOWED_TYPES):
raise ExpressionError(
f"Expression element {type(item).__name__} is not allowed."
)
if isinstance(item, exp.Column) and item.table:
raise ExpressionError("Expressions use column names without table qualifiers.")
if isinstance(item, exp.Cast):
_cast_type(item)
columns = tuple(
dict.fromkeys(
column.name
for column in expression.find_all(exp.Column)
if column.name
)
)
return ParsedExpression(
source=normalized,
expression=expression,
columns=columns,
)
def evaluate_expression(
parsed: ParsedExpression | str,
row: dict[str, Any],
) -> Any:
expression = (
parse_expression(parsed).expression
if isinstance(parsed, str)
else parsed.expression
)
return _evaluate(expression, row)
def infer_expression_type(
parsed: ParsedExpression | str,
schema: dict[str, ExpressionDataType] | None = None,
) -> ExpressionDataType:
expression = (
parse_expression(parsed).expression
if isinstance(parsed, str)
else parsed.expression
)
return _infer(expression, schema or {})
def _evaluate(expression: exp.Expression, row: dict[str, Any]) -> Any:
if isinstance(expression, exp.Paren):
return _evaluate(expression.this, row)
if isinstance(expression, exp.Column):
return row.get(expression.name)
if isinstance(expression, exp.Null):
return None
if isinstance(expression, exp.Boolean):
return bool(expression.this)
if isinstance(expression, exp.Literal):
return _literal(expression)
if isinstance(expression, exp.Neg):
value = _evaluate(expression.this, row)
return None if value is None else -value
if isinstance(expression, exp.Not):
return not bool(_evaluate(expression.this, row))
if isinstance(expression, exp.And):
return bool(_evaluate(expression.this, row)) and bool(
_evaluate(expression.expression, row)
)
if isinstance(expression, exp.Or):
return bool(_evaluate(expression.this, row)) or bool(
_evaluate(expression.expression, row)
)
if isinstance(expression, exp.Is):
left = _evaluate(expression.this, row)
right = _evaluate(expression.expression, row)
return left is right if right is None else left == right
if isinstance(expression, (exp.Add, exp.Sub, exp.Mul, exp.Div, exp.Mod)):
left = _evaluate(expression.this, row)
right = _evaluate(expression.expression, row)
if left is None or right is None:
return None
if isinstance(expression, exp.Add):
return left + right
if isinstance(expression, exp.Sub):
return left - right
if isinstance(expression, exp.Mul):
return left * right
if isinstance(expression, exp.Div):
return left / right
return left % right
if isinstance(expression, (exp.EQ, exp.NEQ, exp.GT, exp.GTE, exp.LT, exp.LTE)):
left = _evaluate(expression.this, row)
right = _evaluate(expression.expression, row)
if isinstance(expression, exp.EQ):
return left == right
if isinstance(expression, exp.NEQ):
return left != right
if left is None or right is None:
return False
if isinstance(expression, exp.GT):
return left > right
if isinstance(expression, exp.GTE):
return left >= right
if isinstance(expression, exp.LT):
return left < right
return left <= right
if isinstance(expression, exp.Lower):
return _string_call(expression.this, row, str.lower)
if isinstance(expression, exp.Upper):
return _string_call(expression.this, row, str.upper)
if isinstance(expression, exp.Trim):
return _string_call(expression.this, row, str.strip)
if isinstance(expression, exp.Length):
value = _evaluate(expression.this, row)
return None if value is None else len(value)
if isinstance(expression, exp.Abs):
value = _evaluate(expression.this, row)
return None if value is None else abs(value)
if isinstance(expression, exp.Round):
value = _evaluate(expression.this, row)
decimals = expression.args.get("decimals")
places = int(_evaluate(decimals, row)) if decimals is not None else 0
return None if value is None else round(value, places)
if isinstance(expression, exp.Coalesce):
arguments = [expression.this, *expression.expressions]
return next(
(
value
for argument in arguments
if (value := _evaluate(argument, row)) is not None
),
None,
)
if isinstance(expression, exp.Concat):
return "".join(
"" if (value := _evaluate(item, row)) is None else str(value)
for item in expression.expressions
)
if isinstance(expression, exp.Replace):
value = _evaluate(expression.this, row)
old = _evaluate(expression.expression, row)
new = _evaluate(expression.args["replacement"], row)
return None if value is None else str(value).replace(str(old), str(new))
if isinstance(expression, exp.Substring):
value = _evaluate(expression.this, row)
if value is None:
return None
start = int(_evaluate(expression.args["start"], row)) - 1
length = expression.args.get("length")
return (
str(value)[start:]
if length is None
else str(value)[start : start + int(_evaluate(length, row))]
)
if isinstance(expression, exp.Cast):
return convert_value(
_evaluate(expression.this, row),
_cast_type(expression),
on_error="fail",
)
if isinstance(expression, exp.Case):
base = _evaluate(expression.this, row) if expression.this is not None else None
for condition in expression.args.get("ifs") or []:
if not isinstance(condition, exp.If):
continue
candidate = _evaluate(condition.this, row)
matched = candidate == base if expression.this is not None else bool(candidate)
if matched:
return _evaluate(condition.args["true"], row)
default = expression.args.get("default")
return _evaluate(default, row) if default is not None else None
raise ExpressionError(
f"Expression element {type(expression).__name__} cannot be evaluated."
)
def convert_value(
value: Any,
target_type: str,
*,
on_error: Literal["fail", "null", "keep"] = "fail",
) -> Any:
if value is None:
return None
try:
if target_type == "string":
return str(value)
if target_type == "integer":
if isinstance(value, bool):
return int(value)
return int(str(value).strip())
if target_type == "number":
return Decimal(str(value).strip())
if target_type == "boolean":
if isinstance(value, bool):
return value
normalized = str(value).strip().casefold()
if normalized in {"true", "yes", "y", "1", "on"}:
return True
if normalized in {"false", "no", "n", "0", "off"}:
return False
raise ValueError(f"{value!r} is not a boolean")
if target_type in {"date", "datetime"}:
from datetime import datetime
parsed = datetime.fromisoformat(str(value).strip().replace("Z", "+00:00"))
return parsed.date() if target_type == "date" else parsed
raise ValueError(f"Unsupported conversion type {target_type!r}")
except (ArithmeticError, TypeError, ValueError):
if on_error == "null":
return None
if on_error == "keep":
return value
raise
def _infer(
expression: exp.Expression,
schema: dict[str, ExpressionDataType],
) -> ExpressionDataType:
if isinstance(expression, exp.Column):
return schema.get(expression.name, "unknown")
if isinstance(expression, exp.Null):
return "null"
if isinstance(expression, exp.Boolean):
return "boolean"
if isinstance(expression, exp.Literal):
if expression.is_string:
return "string"
return "number" if "." in str(expression.this) else "integer"
if isinstance(
expression,
(
exp.EQ,
exp.NEQ,
exp.GT,
exp.GTE,
exp.LT,
exp.LTE,
exp.And,
exp.Or,
exp.Is,
exp.Not,
),
):
return "boolean"
if isinstance(expression, (exp.Length,)):
return "integer"
if isinstance(expression, (exp.Lower, exp.Upper, exp.Trim, exp.Concat, exp.Replace, exp.Substring)):
return "string"
if isinstance(expression, exp.Cast):
return _cast_type(expression) # type: ignore[return-value]
if isinstance(expression, (exp.Add, exp.Sub, exp.Mul, exp.Div, exp.Mod, exp.Abs, exp.Round)):
child_types = {
_infer(item, schema)
for item in expression.iter_expressions()
}
return "number" if "number" in child_types or isinstance(expression, exp.Div) else "integer"
if isinstance(expression, exp.Coalesce):
types = [
_infer(item, schema)
for item in (expression.this, *expression.expressions)
]
return next((item for item in types if item not in {"null", "unknown"}), "unknown")
if isinstance(expression, exp.Case):
types = [
_infer(item.args["true"], schema)
for item in expression.args.get("ifs") or []
if isinstance(item, exp.If)
]
default = expression.args.get("default")
if default is not None:
types.append(_infer(default, schema))
concrete = {item for item in types if item not in {"null", "unknown"}}
return concrete.pop() if len(concrete) == 1 else "unknown"
if isinstance(expression, (exp.Paren, exp.Neg)):
return _infer(expression.this, schema)
return "unknown"
def _literal(expression: exp.Literal) -> Any:
if expression.is_string:
return str(expression.this)
source = str(expression.this)
try:
return int(source)
except ValueError:
return Decimal(source)
def _string_call(
expression: exp.Expression,
row: dict[str, Any],
operation,
) -> str | None:
value = _evaluate(expression, row)
return None if value is None else operation(str(value))
def _cast_type(expression: exp.Cast) -> str:
data_type = expression.args.get("to")
name = str(data_type.this).casefold() if isinstance(data_type, exp.DataType) else ""
mapping = {
"dtype.int": "integer",
"dtype.bigint": "integer",
"dtype.smallint": "integer",
"dtype.float": "number",
"dtype.double": "number",
"dtype.decimal": "number",
"dtype.boolean": "boolean",
"dtype.text": "string",
"dtype.varchar": "string",
"dtype.date": "date",
"dtype.datetime": "datetime",
"dtype.timestamp": "datetime",
"dtype.timestamptz": "datetime",
}
target = mapping.get(name)
if target is None:
raise ExpressionError(f"CAST target {data_type} is not supported.")
return target
__all__ = [
"ExpressionDataType",
"ExpressionError",
"ParsedExpression",
"convert_value",
"evaluate_expression",
"infer_expression_type",
"parse_expression",
]
+709 -20
View File
@@ -4,7 +4,7 @@ import hashlib
import json
import re
from collections import deque
from dataclasses import dataclass
from dataclasses import dataclass, field
from typing import Any
from govoplan_core.core.definition_graphs import (
@@ -17,6 +17,16 @@ from govoplan_dataflow.backend.node_library import (
NODE_TYPES,
node_definition,
)
from govoplan_dataflow.backend.expressions import (
ExpressionError,
infer_expression_type,
parse_expression,
)
from govoplan_dataflow.backend.operator_registry import OPERATOR_REGISTRY
from govoplan_dataflow.backend.subflows import (
SubflowParameterError,
substitute_parameters,
)
from govoplan_dataflow.backend.schemas import DataflowDiagnostic, GraphEdge, GraphNode, PipelineGraph
@@ -40,6 +50,12 @@ DERIVE_OPERATIONS = frozenset(
}
)
JOIN_TYPES = frozenset({"inner", "left", "right", "full"})
DATA_TYPES = frozenset(
{"string", "integer", "number", "boolean", "date", "datetime"}
)
QUALITY_OPERATORS = frozenset(
{"not_null", "type", "min", "max", "allowed", "unique"}
)
def canonical_graph_payload(graph: PipelineGraph) -> dict[str, Any]:
@@ -191,7 +207,17 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
for node in graph.nodes:
if node.type not in SUPPORTED_NODE_TYPES:
continue
diagnostics.extend(_validate_node_config(node))
validator = OPERATOR_REGISTRY.config_validator(node.type)
if validator is None:
diagnostics.append(
_error(
"node.validator_missing",
f"Node type {node.type!r} has no registered validator.",
node_id=node.id,
)
)
continue
diagnostics.extend(validator(node))
ordered, cyclic = topological_order(graph)
if not cyclic and source_nodes and len(output_nodes) == 1:
@@ -341,10 +367,14 @@ def _reachable_from(start: str, adjacency: dict[str, list[str]]) -> set[str]:
class _SchemaState:
columns: frozenset[str]
open: bool = False
types: dict[str, str] = field(default_factory=dict)
def knows(self, column: str) -> bool:
return self.open or column in self.columns
def type_of(self, column: str) -> str:
return self.types.get(column, "unknown")
def _validate_graph_schemas(
graph: PipelineGraph,
@@ -367,21 +397,11 @@ def _validate_graph_schemas(
input_state = input_states[0] if input_states else _SchemaState(frozenset(), open=True)
if node.type == "source.inline":
rows = node.config.get("rows")
columns = {
str(column)
for row in rows if isinstance(rows, list) and isinstance(row, dict)
for column in row
} if isinstance(rows, list) else set()
schemas[node.id] = _SchemaState(
frozenset(columns),
open=not columns,
)
schemas[node.id] = _inline_schema(rows)
continue
if node.type == "source.reference":
columns = _configured_source_columns(node.config.get("source_columns"))
schemas[node.id] = _SchemaState(
frozenset(columns),
open=not columns,
schemas[node.id] = _configured_source_schema(
node.config.get("source_columns")
)
continue
if node.type == "combine.union":
@@ -402,6 +422,7 @@ def _validate_graph_schemas(
schemas[node.id] = _SchemaState(
frozenset().union(*(state.columns for state in input_states)),
open=any(state.open for state in input_states),
types=_merged_schema_types(input_states),
)
continue
if node.type == "combine.join":
@@ -436,6 +457,13 @@ def _validate_graph_schemas(
schemas[node.id] = _SchemaState(
frozenset(left_state.columns | prefixed_right),
open=left_state.open or right_state.open,
types={
**left_state.types,
**{
f"{prefix}{column}": right_state.type_of(column)
for column in right_state.columns
},
},
)
continue
if node.type == "filter":
@@ -448,6 +476,22 @@ def _validate_graph_schemas(
)
schemas[node.id] = input_state
continue
if node.type == "filter.expression":
parsed = _parsed_node_expression(
diagnostics,
node=node,
field_name="expression",
)
if parsed is not None:
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=list(parsed.columns),
field="expression",
)
schemas[node.id] = input_state
continue
if node.type == "distinct":
_validate_columns(
diagnostics,
@@ -487,7 +531,17 @@ def _validate_graph_schemas(
columns=output_columns,
field="fields",
)
schemas[node.id] = _SchemaState(frozenset(output_columns))
schemas[node.id] = _SchemaState(
frozenset(output_columns),
types={
output: input_state.type_of(source)
for source, output in zip(
selected_columns,
output_columns,
strict=False,
)
},
)
continue
if node.type == "derive":
_validate_columns(
@@ -511,10 +565,95 @@ def _validate_graph_schemas(
schemas[node.id] = _SchemaState(
input_state.columns | frozenset((target,)),
open=input_state.open,
types={
**input_state.types,
target: _derive_result_type(
str(node.config.get("operation") or ""),
[
input_state.type_of(column)
for column in _text_items(
node.config.get("source_columns")
)
],
),
},
)
else:
schemas[node.id] = input_state
continue
if node.type == "expression":
parsed = _parsed_node_expression(
diagnostics,
node=node,
field_name="expression",
)
target = str(node.config.get("target_column") or "")
result_type = "unknown"
if parsed is not None:
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=list(parsed.columns),
field="expression",
)
result_type = infer_expression_type(
parsed,
{
name: input_state.type_of(name) # type: ignore[dict-item]
for name in input_state.columns
},
)
expected = str(node.config.get("result_type") or "unknown")
if (
expected != "unknown"
and result_type not in {"unknown", "null", expected}
):
diagnostics.append(
_warning(
"expression.type_mismatch",
f"Expression infers {result_type}, not {expected}.",
node_id=node.id,
field="result_type",
)
)
schemas[node.id] = _SchemaState(
input_state.columns | ({target} if target else set()),
open=input_state.open,
types={
**input_state.types,
**(
{target: expected if expected != "unknown" else result_type}
if target
else {}
),
},
)
continue
if node.type in {"convert", "replace"}:
source = str(node.config.get("source_column") or "")
target = str(node.config.get("target_column") or "")
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=[source],
field="source_column",
)
target_type = (
str(node.config.get("target_type") or "unknown")
if node.type == "convert"
else input_state.type_of(source)
)
schemas[node.id] = _SchemaState(
input_state.columns | ({target} if target else set()),
open=input_state.open,
types={
**input_state.types,
**({target: target_type} if target else {}),
},
)
continue
if node.type == "aggregate":
group_by = _text_items(node.config.get("group_by"))
aggregates = node.config.get("aggregates")
@@ -546,7 +685,27 @@ def _validate_graph_schemas(
columns=output_columns,
field="aggregates",
)
schemas[node.id] = _SchemaState(frozenset(output_columns))
aggregate_types = {
str(item.get("alias")): (
"integer"
if item.get("function") == "count"
else input_state.type_of(str(item.get("column") or ""))
)
for item in aggregates
if isinstance(aggregates, list)
and isinstance(item, dict)
and item.get("alias")
} if isinstance(aggregates, list) else {}
schemas[node.id] = _SchemaState(
frozenset(output_columns),
types={
**{
column: input_state.type_of(column)
for column in group_by
},
**aggregate_types,
},
)
continue
if node.type == "sort":
fields = node.config.get("fields")
@@ -566,14 +725,124 @@ def _validate_graph_schemas(
)
schemas[node.id] = input_state
continue
if node.type == "quality.rules":
rules = node.config.get("rules")
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=[
str(rule.get("column") or "")
for rule in rules
if isinstance(rules, list) and isinstance(rule, dict)
] if isinstance(rules, list) else [],
field="rules",
)
if node.config.get("action", "annotate") == "annotate":
schemas[node.id] = _SchemaState(
input_state.columns
| frozenset(("_quality_valid", "_quality_errors")),
open=input_state.open,
types={
**input_state.types,
"_quality_valid": "boolean",
"_quality_errors": "array",
},
)
else:
schemas[node.id] = input_state
continue
if node.type == "reconcile.compare":
left_state = _port_schema(node_inputs, schemas, "left")
right_state = _port_schema(node_inputs, schemas, "right")
_validate_columns(
diagnostics,
node=node,
state=left_state,
columns=node.config.get("left_keys"),
field="left_keys",
)
_validate_columns(
diagnostics,
node=node,
state=right_state,
columns=node.config.get("right_keys"),
field="right_keys",
)
comparison_columns = node.config.get("compare_columns")
if isinstance(comparison_columns, list):
left_columns: list[str] = []
right_columns: list[str] = []
for item in comparison_columns:
if isinstance(item, str):
left_columns.append(item)
right_columns.append(item)
elif isinstance(item, dict):
left_columns.append(
str(item.get("left") or item.get("column") or "")
)
right_columns.append(
str(
item.get("right")
or item.get("left")
or item.get("column")
or ""
)
)
_validate_columns(
diagnostics,
node=node,
state=left_state,
columns=left_columns,
field="compare_columns",
)
_validate_columns(
diagnostics,
node=node,
state=right_state,
columns=right_columns,
field="compare_columns",
)
prefix = str(node.config.get("right_prefix") or "observed_")
schemas[node.id] = _SchemaState(
left_state.columns
| frozenset(f"{prefix}{column}" for column in right_state.columns)
| frozenset(
(
"_reconciliation_status",
"_reconciliation_differences",
)
),
open=left_state.open or right_state.open,
types={
**left_state.types,
**{
f"{prefix}{column}": right_state.type_of(column)
for column in right_state.columns
},
"_reconciliation_status": "string",
"_reconciliation_differences": "array",
},
)
continue
if node.type == "subflow":
output_schema = _configured_source_schema(
node.config.get("output_schema")
)
schemas[node.id] = (
output_schema
if output_schema.columns
else _SchemaState(frozenset(), open=True)
)
continue
schemas[node.id] = input_state
return diagnostics
def _configured_source_columns(value: object) -> set[str]:
def _configured_source_schema(value: object) -> _SchemaState:
if not isinstance(value, list):
return set()
return {
return _SchemaState(frozenset(), open=True)
columns = {
item if isinstance(item, str) else str(item.get("name"))
for item in value
if (
@@ -581,6 +850,104 @@ def _configured_source_columns(value: object) -> set[str]:
or isinstance(item, dict) and item.get("name")
)
}
types = {
str(item.get("name")): str(
item.get("data_type") or item.get("type") or "unknown"
)
for item in value
if isinstance(item, dict) and item.get("name")
}
return _SchemaState(
frozenset(columns),
open=not columns,
types=types,
)
def _inline_schema(value: object) -> _SchemaState:
if not isinstance(value, list):
return _SchemaState(frozenset(), open=True)
columns = {
str(column)
for row in value
if isinstance(row, dict)
for column in row
}
types: dict[str, str] = {}
for column in columns:
observed = {
_schema_value_type(row.get(column))
for row in value
if isinstance(row, dict) and row.get(column) is not None
}
types[column] = observed.pop() if len(observed) == 1 else "unknown"
return _SchemaState(
frozenset(columns),
open=not columns,
types=types,
)
def _schema_value_type(value: object) -> str:
if isinstance(value, bool):
return "boolean"
if isinstance(value, int):
return "integer"
if isinstance(value, float):
return "number"
if isinstance(value, str):
return "string"
if isinstance(value, list):
return "array"
if isinstance(value, dict):
return "object"
return "unknown"
def _merged_schema_types(states: list[_SchemaState]) -> dict[str, str]:
columns = frozenset().union(*(state.columns for state in states))
types: dict[str, str] = {}
for column in columns:
observed = {
state.type_of(column)
for state in states
if column in state.columns and state.type_of(column) != "unknown"
}
types[column] = observed.pop() if len(observed) == 1 else "unknown"
return types
def _derive_result_type(operation: str, source_types: list[str]) -> str:
if operation in {"upper", "lower", "trim", "concat"}:
return "string"
if operation in {"add", "subtract", "multiply", "divide"}:
return "number" if "number" in source_types or operation == "divide" else "integer"
if operation in {"copy", "coalesce"}:
concrete = {
item for item in source_types if item not in {"unknown", "null"}
}
return concrete.pop() if len(concrete) == 1 else "unknown"
return "unknown"
def _parsed_node_expression(
diagnostics: list[DataflowDiagnostic],
*,
node: GraphNode,
field_name: str,
):
try:
return parse_expression(str(node.config.get(field_name) or ""))
except ExpressionError as exc:
diagnostics.append(
_error(
"expression.invalid",
str(exc),
node_id=node.id,
field=field_name,
)
)
return None
def _port_schema(
@@ -708,6 +1075,16 @@ def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
)
if operator not in {"is_null", "not_null"} and "value" not in config:
diagnostics.append(_node_field_error(node, "filter.value", "Enter a comparison value.", "value"))
elif node.type == "filter.expression":
if not _non_empty_text(config.get("expression")):
diagnostics.append(
_node_field_error(
node,
"expression.required",
"Enter a filter expression.",
"expression",
)
)
elif node.type == "distinct":
columns = config.get("columns", [])
if not isinstance(columns, list) or any(not _non_empty_text(item) for item in columns):
@@ -878,6 +1255,92 @@ def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
"source_columns",
)
)
elif node.type == "expression":
if not _non_empty_text(config.get("target_column")):
diagnostics.append(
_node_field_error(
node,
"expression.target_column",
"Choose an output column.",
"target_column",
)
)
if not _non_empty_text(config.get("expression")):
diagnostics.append(
_node_field_error(
node,
"expression.required",
"Enter an expression.",
"expression",
)
)
if str(config.get("result_type") or "unknown") not in {
"unknown",
*DATA_TYPES,
}:
diagnostics.append(
_node_field_error(
node,
"expression.result_type",
"Choose a supported expression result type.",
"result_type",
)
)
elif node.type == "convert":
for field_name, message in (
("source_column", "Choose a source column."),
("target_column", "Choose an output column."),
):
if not _non_empty_text(config.get(field_name)):
diagnostics.append(
_node_field_error(
node,
f"convert.{field_name}",
message,
field_name,
)
)
if config.get("target_type") not in DATA_TYPES:
diagnostics.append(
_node_field_error(
node,
"convert.target_type",
"Choose a supported conversion type.",
"target_type",
)
)
if config.get("on_error", "fail") not in {"fail", "null", "keep"}:
diagnostics.append(
_node_field_error(
node,
"convert.on_error",
"Choose how conversion errors are handled.",
"on_error",
)
)
elif node.type == "replace":
for field_name, message in (
("source_column", "Choose a source column."),
("target_column", "Choose an output column."),
):
if not _non_empty_text(config.get(field_name)):
diagnostics.append(
_node_field_error(
node,
f"replace.{field_name}",
message,
field_name,
)
)
if config.get("mode", "exact") not in {"exact", "text"}:
diagnostics.append(
_node_field_error(
node,
"replace.mode",
"Choose exact-value or text replacement.",
"mode",
)
)
elif node.type == "sort":
fields = config.get("fields")
if not isinstance(fields, list) or not fields:
@@ -897,6 +1360,220 @@ def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
diagnostics.append(
_node_field_error(node, "limit.count", "Limit must be between 1 and 100,000.", "count")
)
elif node.type == "quality.rules":
rules = config.get("rules")
if not isinstance(rules, list) or not rules or len(rules) > 100:
diagnostics.append(
_node_field_error(
node,
"quality.rules",
"Add between one and 100 quality rules.",
"rules",
)
)
else:
rule_ids: list[str] = []
for rule in rules:
if not isinstance(rule, dict):
diagnostics.append(
_node_field_error(
node,
"quality.rule_shape",
"Every quality rule must be an object.",
"rules",
)
)
break
rule_id = str(rule.get("id") or "").strip()
rule_ids.append(rule_id)
operator = rule.get("operator")
if not rule_id or not _non_empty_text(rule.get("column")):
diagnostics.append(
_node_field_error(
node,
"quality.rule_identity",
"Quality rules need an ID and column.",
"rules",
)
)
break
if operator not in QUALITY_OPERATORS:
diagnostics.append(
_node_field_error(
node,
"quality.operator",
"Choose a supported quality operator.",
"rules",
)
)
break
if operator in {"type", "min", "max"} and "value" not in rule:
diagnostics.append(
_node_field_error(
node,
"quality.value",
"This quality rule requires a value.",
"rules",
)
)
break
if operator == "allowed" and not isinstance(rule.get("values"), list):
diagnostics.append(
_node_field_error(
node,
"quality.values",
"Allowed-value rules require a list of values.",
"rules",
)
)
break
duplicates = sorted(
{rule_id for rule_id in rule_ids if rule_ids.count(rule_id) > 1}
)
if duplicates:
diagnostics.append(
_node_field_error(
node,
"quality.duplicate_id",
f"Quality rule IDs must be unique: {', '.join(duplicates)}.",
"rules",
)
)
if config.get("action", "annotate") not in {"annotate", "drop", "fail"}:
diagnostics.append(
_node_field_error(
node,
"quality.action",
"Choose how invalid rows are handled.",
"action",
)
)
elif node.type == "reconcile.compare":
left_keys = config.get("left_keys")
right_keys = config.get("right_keys")
if (
not isinstance(left_keys, list)
or not left_keys
or any(not _non_empty_text(item) for item in left_keys)
):
diagnostics.append(
_node_field_error(
node,
"reconcile.left_keys",
"Add at least one expected-table key.",
"left_keys",
)
)
if (
not isinstance(right_keys, list)
or not right_keys
or any(not _non_empty_text(item) for item in right_keys)
):
diagnostics.append(
_node_field_error(
node,
"reconcile.right_keys",
"Add at least one observed-table key.",
"right_keys",
)
)
if (
isinstance(left_keys, list)
and isinstance(right_keys, list)
and len(left_keys) != len(right_keys)
):
diagnostics.append(
_node_field_error(
node,
"reconcile.key_count",
"Expected and observed keys must have the same length.",
"right_keys",
)
)
prefix = config.get("right_prefix", "observed_")
if (
not _non_empty_text(prefix)
or re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*_", str(prefix)) is None
):
diagnostics.append(
_node_field_error(
node,
"reconcile.right_prefix",
"Enter an identifier prefix ending in an underscore.",
"right_prefix",
)
)
elif node.type == "subflow":
if not _non_empty_text(config.get("template_ref")):
diagnostics.append(
_node_field_error(
node,
"subflow.template_ref",
"Choose a reusable template.",
"template_ref",
)
)
if not _non_empty_text(config.get("template_version")):
diagnostics.append(
_node_field_error(
node,
"subflow.template_version",
"Pin a template version.",
"template_version",
)
)
parameters = config.get("parameters")
if not isinstance(parameters, dict) or len(parameters) > 100:
diagnostics.append(
_node_field_error(
node,
"subflow.parameters",
"Subflow parameters must be an object with at most 100 entries.",
"parameters",
)
)
try:
nested_payload = substitute_parameters(
config.get("graph"),
parameters if isinstance(parameters, dict) else {},
)
nested = PipelineGraph.model_validate(nested_payload)
except (SubflowParameterError, ValueError) as exc:
diagnostics.append(
_node_field_error(
node,
"subflow.graph",
f"Pinned subflow graph is invalid: {exc}",
"graph",
)
)
else:
bound_inputs = [
item
for item in nested.nodes
if item.type == "source.inline"
and item.config.get("input_binding") is True
]
if len(bound_inputs) != 1:
diagnostics.append(
_node_field_error(
node,
"subflow.input_binding",
"Pinned subflows need exactly one inline input binding.",
"graph",
)
)
for nested_diagnostic in validate_graph(nested):
if nested_diagnostic.severity == "error":
diagnostics.append(
_node_field_error(
node,
"subflow.graph_validation",
f"Pinned subflow: {nested_diagnostic.message}",
"graph",
)
)
break
return diagnostics
@@ -940,6 +1617,18 @@ def _warning(
)
def _register_config_validators() -> None:
for node_type in SUPPORTED_NODE_TYPES:
if OPERATOR_REGISTRY.config_validator(node_type) is None:
OPERATOR_REGISTRY.register_config_validator(
node_type,
_validate_node_config,
)
_register_config_validators()
__all__ = [
"AGGREGATE_FUNCTIONS",
"DERIVE_OPERATIONS",
+2 -1
View File
@@ -164,7 +164,8 @@ DOCUMENTATION = (
metadata={
"first_slice": (
"Inline and governed datasources, union, join, filter, deduplication, select, "
"derived columns, aggregate, sort, limit, output, revisioning, and bounded preview."
"typed expressions, conversion, quality and reconciliation, reusable subflows, "
"aggregate, sort, limit, output, revisioning, and bounded preview."
),
"sql_safety": "Constrained AST compilation only; no pass-through execution.",
},
+221 -1
View File
@@ -13,7 +13,14 @@ from govoplan_core.core.definition_graphs import (
)
NodeCategory = Literal["load", "combine", "filter", "transform", "output"]
NodeCategory = Literal[
"load",
"combine",
"filter",
"transform",
"quality",
"output",
]
SqlSupport = Literal["full", "partial", "none"]
NodePortDefinition = DefinitionPort
NodeConfigField = DefinitionConfigField
@@ -154,6 +161,24 @@ _NODE_TYPES = (
),
default_config={"column": "", "operator": "eq", "value": ""},
),
NodeTypeDefinition(
type="filter.expression",
category="filter",
label="Expression filter",
description="Keep rows for which a safe typed expression is true.",
icon="list-checks",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(
id="expression",
label="Expression",
kind="expression",
required=True,
),
),
default_config={"expression": "true"},
sql_support="partial",
),
NodeTypeDefinition(
type="distinct",
category="filter",
@@ -222,6 +247,119 @@ _NODE_TYPES = (
},
sql_support="partial",
),
NodeTypeDefinition(
type="expression",
category="transform",
label="Expression",
description="Create or replace a column using a safe typed expression.",
icon="function-square",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(
id="target_column",
label="Output column",
kind="text",
required=True,
),
NodeConfigField(
id="expression",
label="Expression",
kind="expression",
required=True,
),
NodeConfigField(
id="result_type",
label="Expected type",
kind="select",
options=(
("unknown", "Infer"),
("string", "Text"),
("integer", "Integer"),
("number", "Number"),
("boolean", "Boolean"),
("date", "Date"),
("datetime", "Date and time"),
),
),
),
default_config={
"target_column": "",
"expression": "",
"result_type": "unknown",
},
sql_support="partial",
),
NodeTypeDefinition(
type="convert",
category="transform",
label="Convert type",
description="Convert a column to a declared data type with explicit error handling.",
icon="replace-all",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="source_column", label="Source column", kind="column", required=True),
NodeConfigField(id="target_column", label="Output column", kind="text", required=True),
NodeConfigField(
id="target_type",
label="Data type",
kind="select",
required=True,
options=(
("string", "Text"),
("integer", "Integer"),
("number", "Number"),
("boolean", "Boolean"),
("date", "Date"),
("datetime", "Date and time"),
),
),
NodeConfigField(
id="on_error",
label="Conversion error",
kind="select",
options=(
("fail", "Stop"),
("null", "Use null"),
("keep", "Keep original"),
),
),
),
default_config={
"source_column": "",
"target_column": "",
"target_type": "string",
"on_error": "fail",
},
sql_support="partial",
),
NodeTypeDefinition(
type="replace",
category="transform",
label="Replace values",
description="Replace exact values or text fragments in a column.",
icon="replace",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="source_column", label="Source column", kind="column", required=True),
NodeConfigField(id="target_column", label="Output column", kind="text", required=True),
NodeConfigField(
id="mode",
label="Mode",
kind="select",
options=(("exact", "Exact value"), ("text", "Text fragment")),
),
NodeConfigField(id="find", label="Find", kind="value"),
NodeConfigField(id="replacement", label="Replacement", kind="value"),
),
default_config={
"source_column": "",
"target_column": "",
"mode": "exact",
"find": "",
"replacement": "",
},
sql_support="none",
),
NodeTypeDefinition(
type="aggregate",
category="transform",
@@ -258,6 +396,87 @@ _NODE_TYPES = (
config_fields=(NodeConfigField(id="count", label="Rows", kind="number", required=True),),
default_config={"count": 100},
),
NodeTypeDefinition(
type="quality.rules",
category="quality",
label="Quality rules",
description="Validate rows and annotate, drop, or reject invalid data.",
icon="badge-check",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="rules", label="Rules", kind="quality_rules", required=True),
NodeConfigField(
id="action",
label="Invalid rows",
kind="select",
options=(
("annotate", "Annotate"),
("drop", "Drop"),
("fail", "Stop"),
),
),
),
default_config={
"rules": [
{
"id": "required",
"column": "",
"operator": "not_null",
}
],
"action": "annotate",
},
sql_support="none",
),
NodeTypeDefinition(
type="reconcile.compare",
category="quality",
label="Reconcile tables",
description="Compare keyed rows and expose missing, matching, and changed records.",
icon="scan-search",
input_ports=(
NodePortDefinition(id="left", label="Expected"),
NodePortDefinition(id="right", label="Observed"),
),
config_fields=(
NodeConfigField(id="left_keys", label="Expected keys", kind="column_list", required=True),
NodeConfigField(id="right_keys", label="Observed keys", kind="column_list", required=True),
NodeConfigField(id="compare_columns", label="Compared columns", kind="field_mapping"),
NodeConfigField(id="right_prefix", label="Observed prefix", kind="text", required=True),
),
default_config={
"left_keys": [""],
"right_keys": [""],
"compare_columns": [],
"right_prefix": "observed_",
},
sql_support="none",
),
NodeTypeDefinition(
type="subflow",
category="transform",
label="Reusable subflow",
description="Run a pinned parameterized template snapshot as one node.",
icon="boxes",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="template_ref", label="Template reference", kind="text", required=True),
NodeConfigField(id="template_version", label="Template version", kind="text", required=True),
NodeConfigField(id="parameters", label="Parameters", kind="json", required=True),
NodeConfigField(id="graph", label="Pinned graph", kind="json", required=True),
),
default_config={
"template_ref": "",
"template_version": "",
"parameters": {},
"graph": {
"schema_version": 1,
"nodes": [],
"edges": [],
},
},
sql_support="none",
),
NodeTypeDefinition(
type="output",
category="output",
@@ -275,6 +494,7 @@ CATEGORY_LABELS = {
"combine": "Combine",
"filter": "Filter",
"transform": "Transform",
"quality": "Quality",
"output": "Output",
}
DATAFLOW_GRAPH_LIBRARY = DefinitionGraphLibrary(
@@ -0,0 +1,156 @@
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Any
from govoplan_dataflow.backend.node_library import (
NODE_LIBRARY,
NodeTypeDefinition,
)
ConfigValidator = Callable[[Any], Sequence[Any]]
SchemaPropagator = Callable[..., Any]
NodeExecutor = Callable[["OperatorExecutionContext"], "OperatorExecutionResult"]
SqlRenderer = Callable[..., Any]
@dataclass(frozen=True, slots=True)
class OperatorExecutionContext:
node: Any
inputs_by_port: Mapping[str, Sequence[str]]
outputs: Mapping[str, list[dict[str, Any]]]
input_sets: Sequence[list[dict[str, Any]]]
input_rows: list[dict[str, Any]]
source_resolver: Callable[..., Any] | None
execution_depth: int
@dataclass(frozen=True, slots=True)
class OperatorExecutionResult:
rows: list[dict[str, Any]]
messages: tuple[str, ...] = ()
source_fingerprints: tuple[dict[str, Any], ...] = ()
input_row_count: int = 0
@dataclass(slots=True)
class NodeOperatorRuntime:
definition: NodeTypeDefinition
config_validator: ConfigValidator | None = None
schema_propagator: SchemaPropagator | None = None
executor: NodeExecutor | None = None
sql_renderer: SqlRenderer | None = None
class NodeOperatorRegistry:
def __init__(self, definitions: Sequence[NodeTypeDefinition] = ()) -> None:
self._operators = {
definition.type: NodeOperatorRuntime(definition=definition)
for definition in definitions
}
def register_definition(self, definition: NodeTypeDefinition) -> None:
if definition.type in self._operators:
raise ValueError(f"Duplicate dataflow node type {definition.type!r}.")
self._operators[definition.type] = NodeOperatorRuntime(
definition=definition
)
def definition(self, node_type: str) -> NodeTypeDefinition | None:
runtime = self._operators.get(node_type)
return runtime.definition if runtime is not None else None
def register_config_validator(
self,
node_type: str,
validator: ConfigValidator,
) -> None:
self._register_facet(node_type, "config_validator", validator)
def register_schema_propagator(
self,
node_type: str,
propagator: SchemaPropagator,
) -> None:
self._register_facet(node_type, "schema_propagator", propagator)
def register_executor(
self,
node_type: str,
executor: NodeExecutor,
) -> None:
self._register_facet(node_type, "executor", executor)
def register_sql_renderer(
self,
node_type: str,
renderer: SqlRenderer,
) -> None:
self._register_facet(node_type, "sql_renderer", renderer)
def config_validator(self, node_type: str) -> ConfigValidator | None:
runtime = self._operators.get(node_type)
return runtime.config_validator if runtime is not None else None
def schema_propagator(self, node_type: str) -> SchemaPropagator | None:
runtime = self._operators.get(node_type)
return runtime.schema_propagator if runtime is not None else None
def executor(self, node_type: str) -> NodeExecutor | None:
runtime = self._operators.get(node_type)
return runtime.executor if runtime is not None else None
def sql_renderer(self, node_type: str) -> SqlRenderer | None:
runtime = self._operators.get(node_type)
return runtime.sql_renderer if runtime is not None else None
def coverage(self) -> Mapping[str, frozenset[str]]:
return {
node_type: frozenset(
facet
for facet in (
"config_validator",
"schema_propagator",
"executor",
"sql_renderer",
)
if getattr(runtime, facet) is not None
)
for node_type, runtime in self._operators.items()
}
def _register_facet(
self,
node_type: str,
facet: str,
callback: object,
) -> None:
runtime = self._operators.get(node_type)
if runtime is None:
raise ValueError(
f"Register the dataflow node definition before its {facet}: "
f"{node_type!r}."
)
if getattr(runtime, facet) is not None:
raise ValueError(
f"Dataflow node {node_type!r} already has a {facet}."
)
setattr(runtime, facet, callback)
OPERATOR_REGISTRY = NodeOperatorRegistry(NODE_LIBRARY)
__all__ = [
"ConfigValidator",
"NodeExecutor",
"NodeOperatorRegistry",
"NodeOperatorRuntime",
"OPERATOR_REGISTRY",
"OperatorExecutionContext",
"OperatorExecutionResult",
"SchemaPropagator",
"SqlRenderer",
]
+408 -113
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Callable, Iterable
import sqlglot
@@ -7,6 +8,8 @@ from sqlglot import exp
from sqlglot.errors import ParseError
from govoplan_dataflow.backend.graph import graph_inputs_by_port, topological_order, validate_graph
from govoplan_dataflow.backend.expressions import parse_expression
from govoplan_dataflow.backend.operator_registry import OPERATOR_REGISTRY
from govoplan_dataflow.backend.schemas import (
DataflowDiagnostic,
GraphEdge,
@@ -28,6 +31,19 @@ LAYOUT_LAYER_GAP = 240
LAYOUT_BRANCH_GAP = 120
@dataclass
class _SqlRenderState:
where_conditions: list[exp.Expression] = field(default_factory=list)
select_expressions: list[exp.Expression] = field(
default_factory=lambda: [exp.Star()]
)
group_by: list[exp.Expression] = field(default_factory=list)
order_by: list[exp.Expression] = field(default_factory=list)
limit: int | None = None
selected: bool = False
distinct: bool = False
def compile_sql(
sql_text: str,
*,
@@ -521,122 +537,31 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
return exp.column(name, table=left_source_name)
return exp.column(name)
where_conditions: list[exp.Expression] = []
select_expressions: list[exp.Expression] = [exp.Star()]
group_by: list[exp.Expression] = []
order_by: list[exp.Expression] = []
limit: int | None = None
selected = False
distinct = False
state = _SqlRenderState()
for node_id in ordered:
node = node_by_id[node_id]
if node.type.startswith("source.") or node.type in {"combine.join", "combine.union"}:
continue
if node.type == "filter":
if selected:
raise SqlCompilationError(
[_node_sql_error(node.id, "sql.filter_order", "Filters after projection or aggregation are not representable yet.")]
)
where_conditions.append(
_condition_expression(node.config, column_expression=column_expression)
)
elif node.type == "select":
if selected:
raise SqlCompilationError(
[_node_sql_error(node.id, "sql.multiple_select", "Only one select or aggregate transform is supported.")]
)
if distinct:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.distinct_order",
"Projection after deduplication is not representable as SELECT DISTINCT.",
)
]
)
select_expressions = []
for field in node.config["fields"]:
column = field if isinstance(field, str) else str(field["column"])
alias = column if isinstance(field, str) else str(field.get("alias") or column)
expression: exp.Expression = column_expression(column)
if alias != column:
expression = expression.as_(alias)
select_expressions.append(expression)
selected = True
elif node.type == "aggregate":
if selected:
raise SqlCompilationError(
[_node_sql_error(node.id, "sql.multiple_select", "Only one select or aggregate transform is supported.")]
)
if distinct:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.distinct_order",
"Aggregation after deduplication is not representable in the constrained dialect.",
)
]
)
select_expressions = [
column_expression(str(column))
for column in node.config.get("group_by", [])
]
group_by = [
column_expression(str(column))
for column in node.config.get("group_by", [])
]
for aggregate in node.config["aggregates"]:
function = str(aggregate["function"])
column = aggregate.get("column")
argument: exp.Expression = (
exp.Star()
if function == "count" and column in (None, "", "*")
else column_expression(str(column))
)
aggregate_expression = _aggregate_expression(function, argument)
select_expressions.append(aggregate_expression.as_(str(aggregate["alias"])))
selected = True
elif node.type == "distinct":
if node.config.get("columns"):
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.distinct_keys",
"Key-based deduplication is not representable as SELECT DISTINCT.",
)
]
)
if distinct:
raise SqlCompilationError(
[_node_sql_error(node.id, "sql.multiple_distinct", "Only one DISTINCT transform is supported.")]
)
distinct = True
elif node.type == "sort":
order_by = [
exp.Ordered(
this=column_expression(str(field["column"])),
desc=field.get("direction", "asc") == "desc",
nulls_first=False,
)
for field in node.config["fields"]
]
elif node.type == "limit":
limit = int(node.config["count"])
elif node.type != "output":
renderer = OPERATOR_REGISTRY.sql_renderer(node.type)
if renderer is None:
raise SqlCompilationError(
[_node_sql_error(node.id, "sql.node_not_representable", f"{node.type!r} cannot be rendered as SQL.")]
[
_node_sql_error(
node.id,
"sql.renderer_missing",
f"Node type {node.type!r} has no registered SQL renderer.",
)
]
)
renderer(node, state, column_expression)
if union_expression is not None:
query = exp.select(*select_expressions).from_(
query = exp.select(*state.select_expressions).from_(
union_expression.subquery("_unioned")
)
elif left_source_name:
query = exp.select(*select_expressions).from_(exp.to_table(left_source_name))
query = exp.select(*state.select_expressions).from_(exp.to_table(left_source_name))
else:
raise SqlCompilationError(
[_sql_error("sql.source_required", "SQL rendering needs a source.")]
@@ -658,16 +583,16 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
on=_combine_and(join_conditions),
join_type=str(join_node.config.get("join_type", "inner")),
)
if where_conditions:
query = query.where(_combine_and(where_conditions))
if group_by:
query = query.group_by(*group_by)
if order_by:
query = query.order_by(*order_by)
if distinct:
if state.where_conditions:
query = query.where(_combine_and(state.where_conditions))
if state.group_by:
query = query.group_by(*state.group_by)
if state.order_by:
query = query.order_by(*state.order_by)
if state.distinct:
query = query.distinct()
if limit is not None:
query = query.limit(limit)
if state.limit is not None:
query = query.limit(state.limit)
return query.sql(dialect="duckdb", pretty=True), diagnostics
@@ -1032,6 +957,376 @@ def _condition_expression(
raise SqlCompilationError([_sql_error("filter.operator", f"Unsupported filter operator {operator!r}.")])
def _qualified_expression(
source: str,
*,
column_expression: Callable[[str], exp.Column],
) -> exp.Expression:
parsed = parse_expression(source)
expression = parsed.expression.copy()
for column in list(expression.find_all(exp.Column)):
column.replace(column_expression(column.name))
return expression
def _sql_data_type(data_type: str) -> exp.DataType:
mapping = {
"string": exp.DataType.Type.TEXT,
"integer": exp.DataType.Type.BIGINT,
"number": exp.DataType.Type.DECIMAL,
"boolean": exp.DataType.Type.BOOLEAN,
"date": exp.DataType.Type.DATE,
"datetime": exp.DataType.Type.TIMESTAMPTZ,
}
target = mapping.get(data_type)
if target is None:
raise SqlCompilationError(
[_sql_error("convert.target_type", f"Unsupported conversion type {data_type!r}.")]
)
return exp.DataType.build(target)
def _render_filter(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
if state.selected:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.filter_order",
"Filters after projection or aggregation are not representable yet.",
)
]
)
state.where_conditions.append(
_condition_expression(node.config, column_expression=column_expression)
)
def _render_expression_filter(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
if state.selected:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.filter_order",
"Expression filters after projection are not representable yet.",
)
]
)
state.where_conditions.append(
_qualified_expression(
str(node.config["expression"]),
column_expression=column_expression,
)
)
def _render_select(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
_require_projection_slot(node, state)
if state.distinct:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.distinct_order",
"Projection after deduplication is not representable as SELECT DISTINCT.",
)
]
)
state.select_expressions = []
for field_config in node.config["fields"]:
column = (
field_config
if isinstance(field_config, str)
else str(field_config["column"])
)
alias = (
column
if isinstance(field_config, str)
else str(field_config.get("alias") or column)
)
expression: exp.Expression = column_expression(column)
if alias != column:
expression = expression.as_(alias)
state.select_expressions.append(expression)
state.selected = True
def _render_aggregate(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
_require_projection_slot(node, state)
if state.distinct:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.distinct_order",
"Aggregation after deduplication is not representable in the constrained dialect.",
)
]
)
state.select_expressions = [
column_expression(str(column))
for column in node.config.get("group_by", [])
]
state.group_by = [
column_expression(str(column))
for column in node.config.get("group_by", [])
]
for aggregate in node.config["aggregates"]:
function = str(aggregate["function"])
column = aggregate.get("column")
argument: exp.Expression = (
exp.Star()
if function == "count" and column in (None, "", "*")
else column_expression(str(column))
)
aggregate_expression = _aggregate_expression(function, argument)
state.select_expressions.append(
aggregate_expression.as_(str(aggregate["alias"]))
)
state.selected = True
def _render_expression(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
_require_projection_slot(node, state)
expression = _qualified_expression(
str(node.config["expression"]),
column_expression=column_expression,
)
state.select_expressions = [
exp.Star(),
expression.as_(str(node.config["target_column"])),
]
state.selected = True
def _render_derive(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
_require_projection_slot(node, state)
columns = [
column_expression(str(column))
for column in node.config["source_columns"]
]
operation = str(node.config["operation"])
if operation == "copy":
expression: exp.Expression = columns[0]
elif operation == "upper":
expression = exp.Upper(this=columns[0])
elif operation == "lower":
expression = exp.Lower(this=columns[0])
elif operation == "trim":
expression = exp.Trim(this=columns[0])
elif operation == "concat":
separator = exp.Literal.string(str(node.config.get("separator", " ")))
expressions: list[exp.Expression] = []
for index, column in enumerate(columns):
if index:
expressions.append(separator.copy())
expressions.append(column)
expression = exp.Concat(expressions=expressions)
elif operation == "coalesce":
expression = exp.Coalesce(
this=columns[0],
expressions=columns[1:],
)
elif operation in {"add", "subtract", "multiply", "divide"}:
operation_type = {
"add": exp.Add,
"subtract": exp.Sub,
"multiply": exp.Mul,
"divide": exp.Div,
}[operation]
expression = operation_type(this=columns[0], expression=columns[1])
else:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.derive_operation",
f"Derive operation {operation!r} cannot be rendered as SQL.",
)
]
)
state.select_expressions = [
exp.Star(),
expression.as_(str(node.config["target_column"])),
]
state.selected = True
def _render_convert(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
_require_projection_slot(node, state)
on_error = str(node.config.get("on_error", "fail"))
if on_error == "keep":
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.convert_keep",
"Keeping the original value after a conversion error is not representable in SQL view.",
)
]
)
cast_class = exp.TryCast if on_error == "null" else exp.Cast
expression = cast_class(
this=column_expression(str(node.config["source_column"])),
to=_sql_data_type(str(node.config["target_type"])),
)
state.select_expressions = [
exp.Star(),
expression.as_(str(node.config["target_column"])),
]
state.selected = True
def _render_distinct(
node: GraphNode,
state: _SqlRenderState,
_column_expression: Callable[[str], exp.Column],
) -> None:
if node.config.get("columns"):
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.distinct_keys",
"Key-based deduplication is not representable as SELECT DISTINCT.",
)
]
)
if state.distinct:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.multiple_distinct",
"Only one DISTINCT transform is supported.",
)
]
)
state.distinct = True
def _render_sort(
node: GraphNode,
state: _SqlRenderState,
column_expression: Callable[[str], exp.Column],
) -> None:
state.order_by = [
exp.Ordered(
this=column_expression(str(item["column"])),
desc=item.get("direction", "asc") == "desc",
nulls_first=False,
)
for item in node.config["fields"]
]
def _render_limit(
node: GraphNode,
state: _SqlRenderState,
_column_expression: Callable[[str], exp.Column],
) -> None:
state.limit = int(node.config["count"])
def _render_noop(
_node: GraphNode,
_state: _SqlRenderState,
_column_expression: Callable[[str], exp.Column],
) -> None:
return None
def _render_unsupported(
node: GraphNode,
_state: _SqlRenderState,
_column_expression: Callable[[str], exp.Column],
) -> None:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.node_not_representable",
f"{node.type!r} cannot be rendered as SQL.",
)
]
)
def _require_projection_slot(
node: GraphNode,
state: _SqlRenderState,
) -> None:
if state.selected:
raise SqlCompilationError(
[
_node_sql_error(
node.id,
"sql.multiple_select",
"Only one expression, conversion, derive, select, or aggregate transform is supported in SQL view.",
)
]
)
def _register_sql_renderers() -> None:
renderers = {
"source.inline": _render_noop,
"source.reference": _render_noop,
"combine.union": _render_noop,
"combine.join": _render_noop,
"filter": _render_filter,
"filter.expression": _render_expression_filter,
"distinct": _render_distinct,
"select": _render_select,
"derive": _render_derive,
"expression": _render_expression,
"convert": _render_convert,
"replace": _render_unsupported,
"aggregate": _render_aggregate,
"sort": _render_sort,
"limit": _render_limit,
"quality.rules": _render_unsupported,
"reconcile.compare": _render_unsupported,
"subflow": _render_unsupported,
"output": _render_noop,
}
for node_type, renderer in renderers.items():
if OPERATOR_REGISTRY.sql_renderer(node_type) is None:
OPERATOR_REGISTRY.register_sql_renderer(node_type, renderer)
_register_sql_renderers()
def _literal_value(expression: exp.Expression) -> Any:
if isinstance(expression, exp.Null):
return None
+44
View File
@@ -0,0 +1,44 @@
from __future__ import annotations
import json
import re
from typing import Any
class SubflowParameterError(ValueError):
pass
def substitute_parameters(value: Any, parameters: dict[str, Any]) -> Any:
if isinstance(value, dict):
if set(value) == {"$parameter"}:
key = str(value["$parameter"])
if key not in parameters:
raise SubflowParameterError(
f"Subflow parameter {key!r} was not provided."
)
return parameters[key]
return {
key: substitute_parameters(item, parameters)
for key, item in value.items()
}
if isinstance(value, list):
return [substitute_parameters(item, parameters) for item in value]
if isinstance(value, str):
def replacement(match: re.Match[str]) -> str:
key = match.group(1)
if key not in parameters:
raise SubflowParameterError(
f"Subflow parameter {key!r} was not provided."
)
return json.dumps(parameters[key], ensure_ascii=True, default=str)
return re.sub(
r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}",
replacement,
value,
)
return value
__all__ = ["SubflowParameterError", "substitute_parameters"]