feat: add extensible operators and golden flows
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user