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
+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