Expand governed Dataflow editor and node library

This commit is contained in:
2026-07-28 11:14:01 +02:00
parent df468a2bd8
commit dee8380631
22 changed files with 3564 additions and 291 deletions

View File

@@ -17,10 +17,40 @@ the previewed row contents.
- **Risk Compliance:** sanctions matching policy, review, dispositions, and - **Risk Compliance:** sanctions matching policy, review, dispositions, and
legal evidence. legal evidence.
The first implementation supports bounded inline tabular sources and the ## Node Library
`filter`, `select`, `aggregate`, `sort`, `limit`, and `output` transforms. The
connector-backed source contract is tracked separately so credentials never The canonical backend catalogue is exposed to the WebUI and groups executable
become part of a pipeline definition. nodes by purpose:
| Group | Nodes |
| --- | --- |
| Load | Inline data, connector source |
| Combine | Append rows, join tables |
| Filter | Filter rows, remove duplicates |
| Transform | Select columns, derive column, aggregate, sort, limit |
| Output | Preview output |
Join nodes have explicit left and right ports. Append nodes accept two or more
inputs. Derived columns use a constrained operation catalogue rather than
arbitrary code.
Connector sources are resolved through the versioned Core capability
`connectors.tabular_sources`; Dataflow does not import Connectors or store its
credentials. Source references pin an expected fingerprint so schema/content
drift fails visibly instead of silently changing a run.
## SQL And Preview Safety
The SQL workbench parses one `SELECT` statement into the canonical graph. The
dialect supports projection, aliases, filters, grouping, aggregate functions,
sorting, limits, `DISTINCT`, and one two-source equi-join. It rejects DDL, DML,
subqueries, arbitrary functions, file access, and unchecked pass-through
execution.
Preview reads at most 250 rows per source and enforces time, intermediate-row,
result-byte, graph-node, and response-row bounds. Saved previews record the
pipeline revision, executor version, source fingerprints, node diagnostics,
and output summary, but not source or result rows.
## Development ## Development

View File

@@ -6,26 +6,42 @@ import time
from collections import defaultdict from collections import defaultdict
from dataclasses import dataclass from dataclasses import dataclass
from decimal import Decimal from decimal import Decimal
from typing import Any from typing import Any, Callable
from govoplan_dataflow.backend.graph import graph_input_map, topological_order, validate_graph from govoplan_dataflow.backend.graph import graph_inputs_by_port, topological_order, validate_graph
from govoplan_dataflow.backend.schemas import ( from govoplan_dataflow.backend.schemas import (
DataflowDiagnostic, DataflowDiagnostic,
GraphNode,
NodePreviewDiagnostic, NodePreviewDiagnostic,
PipelineGraph, PipelineGraph,
PreviewColumn, PreviewColumn,
) )
EXECUTOR_VERSION = "dataflow-preview-v1" EXECUTOR_VERSION = "dataflow-preview-v2"
MAX_EXECUTION_SECONDS = 2.0 MAX_EXECUTION_SECONDS = 2.0
MAX_RESULT_BYTES = 1_000_000 MAX_RESULT_BYTES = 1_000_000
MAX_SOURCE_ROWS = 250
MAX_INTERMEDIATE_ROWS = 10_000
class PipelineExecutionError(RuntimeError): class PipelineExecutionError(RuntimeError):
def __init__(self, message: str, *, node_id: str | None = None) -> None: def __init__(
self,
message: str,
*,
node_id: str | None = None,
node_diagnostics: tuple[NodePreviewDiagnostic, ...] = (),
source_fingerprints: tuple[dict[str, Any], ...] = (),
input_row_count: int = 0,
diagnostics: tuple[DataflowDiagnostic, ...] = (),
) -> None:
super().__init__(message) super().__init__(message)
self.node_id = node_id self.node_id = node_id
self.node_diagnostics = node_diagnostics
self.source_fingerprints = source_fingerprints
self.input_row_count = input_row_count
self.diagnostics = diagnostics
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -40,14 +56,32 @@ class PipelineExecutionResult:
input_row_count: int input_row_count: int
def execute_preview(graph: PipelineGraph, *, row_limit: int) -> PipelineExecutionResult: @dataclass(frozen=True)
class ResolvedSource:
rows: tuple[dict[str, Any], ...]
source_ref: str
provider: str
fingerprint: str
total_rows: int
truncated: bool = False
SourceResolver = Callable[[GraphNode, int], ResolvedSource]
def execute_preview(
graph: PipelineGraph,
*,
row_limit: int,
source_resolver: SourceResolver | None = None,
) -> PipelineExecutionResult:
validation = validate_graph(graph) validation = validate_graph(graph)
errors = [item for item in validation if item.severity == "error"] errors = [item for item in validation if item.severity == "error"]
if errors: if errors:
raise PipelineExecutionError(errors[0].message, node_id=errors[0].node_id) raise PipelineExecutionError(errors[0].message, node_id=errors[0].node_id)
node_by_id = {node.id: node for node in graph.nodes} node_by_id = {node.id: node for node in graph.nodes}
inputs = graph_input_map(graph) inputs = graph_inputs_by_port(graph)
ordered, cyclic = topological_order(graph) ordered, cyclic = topological_order(graph)
if cyclic: if cyclic:
raise PipelineExecutionError("Pipeline graph contains a cycle") raise PipelineExecutionError("Pipeline graph contains a cycle")
@@ -59,11 +93,34 @@ def execute_preview(graph: PipelineGraph, *, row_limit: int) -> PipelineExecutio
input_row_count = 0 input_row_count = 0
for node_id in ordered: for node_id in ordered:
if time.monotonic() - started > MAX_EXECUTION_SECONDS:
raise PipelineExecutionError("Preview exceeded the two-second execution limit", node_id=node_id)
node = node_by_id[node_id] node = node_by_id[node_id]
node_started = time.monotonic() node_started = time.monotonic()
input_rows = [] if node.type.startswith("source.") else outputs[inputs[node_id]] node_inputs = inputs.get(node_id, {})
input_sets = [
outputs[source_id]
for port_sources in node_inputs.values()
for source_id in port_sources
]
input_rows = input_sets[0] if len(input_sets) == 1 else []
node_messages: list[str] = []
if time.monotonic() - started > MAX_EXECUTION_SECONDS:
failed = NodePreviewDiagnostic(
node_id=node.id,
status="failed",
input_rows=sum(len(rows) for rows in input_sets),
output_rows=0,
duration_ms=0,
columns=[],
messages=["Preview exceeded the two-second execution limit"],
)
raise PipelineExecutionError(
"Preview exceeded the two-second execution limit",
node_id=node.id,
node_diagnostics=tuple([*node_diagnostics, failed]),
source_fingerprints=tuple(source_fingerprints),
input_row_count=input_row_count,
diagnostics=tuple(item for item in validation if item.severity != "error"),
)
try: try:
if node.type == "source.inline": if node.type == "source.inline":
output_rows = [dict(row) for row in node.config.get("rows", [])] output_rows = [dict(row) for row in node.config.get("rows", [])]
@@ -78,14 +135,64 @@ def execute_preview(graph: PipelineGraph, *, row_limit: int) -> PipelineExecutio
} }
) )
elif node.type == "source.reference": elif node.type == "source.reference":
if source_resolver is None:
raise PipelineExecutionError( raise PipelineExecutionError(
"This connector-backed source is not available to the local preview executor.", "Connector-backed preview requires the Connectors tabular-source 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": "connector",
"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,
node_id=node.id, node_id=node.id,
) )
elif node.type == "filter": elif node.type == "filter":
output_rows = _filter_rows(input_rows, node.config, node_id=node.id) 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": elif node.type == "select":
output_rows = _select_rows(input_rows, node.config) 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": elif node.type == "aggregate":
output_rows = _aggregate_rows(input_rows, node.config, node_id=node.id) output_rows = _aggregate_rows(input_rows, node.config, node_id=node.id)
elif node.type == "sort": elif node.type == "sort":
@@ -96,25 +203,45 @@ def execute_preview(graph: PipelineGraph, *, row_limit: int) -> PipelineExecutio
output_rows = [dict(row) for row in input_rows] output_rows = [dict(row) for row in input_rows]
else: else:
raise PipelineExecutionError(f"Unsupported node type: {node.type}", node_id=node.id) raise PipelineExecutionError(f"Unsupported node type: {node.type}", node_id=node.id)
except PipelineExecutionError:
raise
except (ArithmeticError, TypeError, ValueError) as exc:
raise PipelineExecutionError(str(exc), node_id=node.id) from exc
if len(json.dumps(output_rows, default=str).encode("utf-8")) > MAX_RESULT_BYTES: if len(json.dumps(output_rows, default=str).encode("utf-8")) > MAX_RESULT_BYTES:
raise PipelineExecutionError( raise PipelineExecutionError(
"A preview node exceeded the one-megabyte result limit.", "A preview node exceeded the one-megabyte result limit.",
node_id=node.id, node_id=node.id,
) )
except (PipelineExecutionError, ArithmeticError, KeyError, TypeError, ValueError) as exc:
execution_error = (
exc
if isinstance(exc, PipelineExecutionError)
else PipelineExecutionError(str(exc), node_id=node.id)
)
failed = NodePreviewDiagnostic(
node_id=execution_error.node_id or node.id,
status="failed",
input_rows=sum(len(rows) for rows in input_sets),
output_rows=0,
duration_ms=round((time.monotonic() - node_started) * 1000, 3),
columns=[],
messages=[str(execution_error)],
)
raise PipelineExecutionError(
str(execution_error),
node_id=execution_error.node_id or node.id,
node_diagnostics=tuple([*node_diagnostics, failed]),
source_fingerprints=tuple(source_fingerprints),
input_row_count=input_row_count,
diagnostics=tuple(item for item in validation if item.severity != "error"),
) from exc
outputs[node.id] = output_rows outputs[node.id] = output_rows
node_diagnostics.append( node_diagnostics.append(
NodePreviewDiagnostic( NodePreviewDiagnostic(
node_id=node.id, node_id=node.id,
status="succeeded", status="succeeded",
input_rows=len(input_rows), input_rows=sum(len(rows) for rows in input_sets),
output_rows=len(output_rows), output_rows=len(output_rows),
duration_ms=round((time.monotonic() - node_started) * 1000, 3), duration_ms=round((time.monotonic() - node_started) * 1000, 3),
columns=infer_columns(output_rows), columns=infer_columns(output_rows),
messages=node_messages,
) )
) )
@@ -126,7 +253,7 @@ def execute_preview(graph: PipelineGraph, *, row_limit: int) -> PipelineExecutio
total_rows=len(all_rows), total_rows=len(all_rows),
truncated=len(rows) < len(all_rows), truncated=len(rows) < len(all_rows),
columns=infer_columns(all_rows), columns=infer_columns(all_rows),
diagnostics=[], diagnostics=[item for item in validation if item.severity != "error"],
node_diagnostics=node_diagnostics, node_diagnostics=node_diagnostics,
source_fingerprints=source_fingerprints, source_fingerprints=source_fingerprints,
input_row_count=input_row_count, input_row_count=input_row_count,
@@ -156,6 +283,202 @@ def infer_columns(rows: list[dict[str, Any]]) -> list[PreviewColumn]:
return columns return columns
def _union_rows(
inputs: list[list[dict[str, Any]]],
config: dict[str, Any],
) -> list[dict[str, Any]]:
rows = [dict(row) for input_rows in inputs for row in input_rows]
if config.get("mode", "all") == "distinct":
return _distinct_rows(rows, {})
return rows
def _join_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"]]
join_type = str(config.get("join_type", "inner"))
right_prefix = str(config.get("right_prefix", "right_"))
left_columns = _ordered_columns(left_rows)
right_columns = _ordered_columns(right_rows)
right_index: dict[tuple[Any, ...], list[tuple[int, dict[str, Any]]]] = defaultdict(list)
for index, row in enumerate(right_rows):
key = _join_key(row, right_keys)
if key is not None:
right_index[key].append((index, row))
output: list[dict[str, Any]] = []
matched_right: set[int] = set()
for left_row in left_rows:
key = _join_key(left_row, left_keys)
matches = right_index.get(key, ()) if key is not None else ()
if matches:
for right_index_value, right_row in matches:
matched_right.add(right_index_value)
output.append(
_merge_join_rows(
left_row,
right_row,
left_columns=left_columns,
right_columns=right_columns,
right_prefix=right_prefix,
)
)
_guard_intermediate_size(output, node_id=node_id)
elif join_type in {"left", "full"}:
output.append(
_merge_join_rows(
left_row,
None,
left_columns=left_columns,
right_columns=right_columns,
right_prefix=right_prefix,
)
)
if join_type in {"right", "full"}:
for index, right_row in enumerate(right_rows):
if index in matched_right:
continue
output.append(
_merge_join_rows(
None,
right_row,
left_columns=left_columns,
right_columns=right_columns,
right_prefix=right_prefix,
)
)
_guard_intermediate_size(output, node_id=node_id)
return output
def _distinct_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
) -> list[dict[str, Any]]:
columns = [str(item) for item in config.get("columns", [])]
seen: set[str] = set()
output: list[dict[str, Any]] = []
for row in rows:
value = {column: row.get(column) for column in columns} if columns else row
identity = json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
if identity in seen:
continue
seen.add(identity)
output.append(dict(row))
return output
def _derive_rows(
rows: list[dict[str, Any]],
config: dict[str, Any],
*,
node_id: str,
) -> list[dict[str, Any]]:
target = str(config["target_column"])
operation = str(config["operation"])
columns = [str(item) for item in config["source_columns"]]
separator = str(config.get("separator", " "))
output: list[dict[str, Any]] = []
for row in rows:
values = [row.get(column) for column in columns]
try:
derived = _derive_value(operation, values, separator=separator)
except (ArithmeticError, TypeError, ValueError) as exc:
raise PipelineExecutionError(
f"Cannot derive {target!r} with {operation!r}: {exc}",
node_id=node_id,
) from exc
result = dict(row)
result[target] = derived
output.append(result)
return output
def _derive_value(operation: str, values: list[Any], *, separator: str) -> Any:
if operation == "copy":
return values[0]
if operation == "upper":
return None if values[0] is None else str(values[0]).upper()
if operation == "lower":
return None if values[0] is None else str(values[0]).lower()
if operation == "trim":
return None if values[0] is None else str(values[0]).strip()
if operation == "concat":
return separator.join(str(value) for value in values if value is not None)
if operation == "coalesce":
return next((value for value in values if value not in (None, "")), None)
if any(value is None for value in values):
return None
left, right = values
if isinstance(left, bool) or isinstance(right, bool):
raise TypeError("boolean values are not numeric inputs")
if operation == "add":
return left + right
if operation == "subtract":
return left - right
if operation == "multiply":
return left * right
if operation == "divide":
return left / right
raise ValueError(f"unknown derive operation {operation!r}")
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):
return None
try:
hash(values)
except TypeError:
return tuple(
json.dumps(value, sort_keys=True, separators=(",", ":"), default=str)
for value in values
)
return values
def _ordered_columns(rows: list[dict[str, Any]]) -> list[str]:
columns: list[str] = []
for row in rows:
for column in row:
if column not in columns:
columns.append(column)
return columns
def _merge_join_rows(
left_row: dict[str, Any] | None,
right_row: dict[str, Any] | None,
*,
left_columns: list[str],
right_columns: list[str],
right_prefix: str,
) -> dict[str, Any]:
output = {
column: left_row.get(column) if left_row is not None else None
for column in left_columns
}
for column in right_columns:
output_column = f"{right_prefix}{column}"
output[output_column] = right_row.get(column) if right_row is not None else None
return output
def _guard_intermediate_size(rows: list[dict[str, Any]], *, node_id: str) -> None:
if len(rows) > MAX_INTERMEDIATE_ROWS:
raise PipelineExecutionError(
f"Preview join exceeded the {MAX_INTERMEDIATE_ROWS:,}-row intermediate limit.",
node_id=node_id,
)
def _filter_rows( def _filter_rows(
rows: list[dict[str, Any]], rows: list[dict[str, Any]],
config: dict[str, Any], config: dict[str, Any],
@@ -305,8 +628,11 @@ def _type_name(value: Any) -> str:
__all__ = [ __all__ = [
"EXECUTOR_VERSION", "EXECUTOR_VERSION",
"MAX_SOURCE_ROWS",
"PipelineExecutionError", "PipelineExecutionError",
"PipelineExecutionResult", "PipelineExecutionResult",
"ResolvedSource",
"SourceResolver",
"execute_preview", "execute_preview",
"infer_columns", "infer_columns",
] ]

View File

@@ -2,28 +2,35 @@ from __future__ import annotations
import hashlib import hashlib
import json import json
import re
from collections import deque from collections import deque
from dataclasses import dataclass
from typing import Any from typing import Any
from govoplan_dataflow.backend.schemas import DataflowDiagnostic, GraphNode, PipelineGraph from govoplan_dataflow.backend.node_library import NODE_TYPES, node_definition
from govoplan_dataflow.backend.schemas import DataflowDiagnostic, GraphEdge, GraphNode, PipelineGraph
SUPPORTED_NODE_TYPES = frozenset( SUPPORTED_NODE_TYPES = frozenset(NODE_TYPES)
{
"source.inline",
"source.reference",
"filter",
"select",
"aggregate",
"sort",
"limit",
"output",
}
)
FILTER_OPERATORS = frozenset( FILTER_OPERATORS = frozenset(
{"eq", "ne", "gt", "gte", "lt", "lte", "contains", "is_null", "not_null"} {"eq", "ne", "gt", "gte", "lt", "lte", "contains", "is_null", "not_null"}
) )
AGGREGATE_FUNCTIONS = frozenset({"count", "sum", "avg", "min", "max"}) AGGREGATE_FUNCTIONS = frozenset({"count", "sum", "avg", "min", "max"})
DERIVE_OPERATIONS = frozenset(
{
"copy",
"upper",
"lower",
"trim",
"concat",
"coalesce",
"add",
"subtract",
"multiply",
"divide",
}
)
JOIN_TYPES = frozenset({"inner", "left", "right", "full"})
def canonical_graph_payload(graph: PipelineGraph) -> dict[str, Any]: def canonical_graph_payload(graph: PipelineGraph) -> dict[str, Any]:
@@ -49,7 +56,7 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
if len(edge_ids) != len(graph.edges): if len(edge_ids) != len(graph.edges):
diagnostics.append(_error("graph.duplicate_edge", "Edge identifiers must be unique.")) diagnostics.append(_error("graph.duplicate_edge", "Edge identifiers must be unique."))
incoming: dict[str, list[str]] = {node_id: [] for node_id in nodes} incoming: dict[str, list[GraphEdge]] = {node_id: [] for node_id in nodes}
outgoing: dict[str, list[str]] = {node_id: [] for node_id in nodes} outgoing: dict[str, list[str]] = {node_id: [] for node_id in nodes}
for edge in graph.edges: for edge in graph.edges:
if edge.source not in nodes: if edge.source not in nodes:
@@ -67,8 +74,32 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
_error("edge.self_reference", "A node cannot connect to itself.", node_id=edge.source) _error("edge.self_reference", "A node cannot connect to itself.", node_id=edge.source)
) )
continue continue
source_definition = node_definition(nodes[edge.source].type)
target_definition = node_definition(nodes[edge.target].type)
if source_definition and edge.source_port not in {
port.id for port in source_definition.output_ports
}:
diagnostics.append(
_error(
"edge.unknown_source_port",
f"Node {edge.source!r} has no output port {edge.source_port!r}.",
node_id=edge.source,
)
)
continue
if target_definition and edge.target_port not in {
port.id for port in target_definition.input_ports
}:
diagnostics.append(
_error(
"edge.unknown_target_port",
f"Node {edge.target!r} has no input port {edge.target_port!r}.",
node_id=edge.target,
)
)
continue
outgoing[edge.source].append(edge.target) outgoing[edge.source].append(edge.target)
incoming[edge.target].append(edge.source) incoming[edge.target].append(edge)
if not graph.nodes: if not graph.nodes:
diagnostics.append(_error("graph.empty", "Add a source and an output before saving the pipeline.")) diagnostics.append(_error("graph.empty", "Add a source and an output before saving the pipeline."))
@@ -76,11 +107,34 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
source_nodes = [node for node in graph.nodes if node.type.startswith("source.")] source_nodes = [node for node in graph.nodes if node.type.startswith("source.")]
output_nodes = [node for node in graph.nodes if node.type == "output"] output_nodes = [node for node in graph.nodes if node.type == "output"]
if len(source_nodes) != 1: if not source_nodes:
diagnostics.append( diagnostics.append(
_error( _error(
"graph.source_count", "graph.source_count",
"The first release supports exactly one source node per pipeline.", "A pipeline needs at least one source node.",
)
)
elif len(source_nodes) > 10:
diagnostics.append(_error("graph.source_limit", "Pipelines are limited to ten sources."))
source_names = [
str(node.config.get("source_name", "")).strip()
for node in source_nodes
if _non_empty_text(node.config.get("source_name"))
]
duplicate_source_names = sorted(
{
name
for name in source_names
if sum(candidate.casefold() == name.casefold() for candidate in source_names) > 1
},
key=str.casefold,
)
if duplicate_source_names:
diagnostics.append(
_error(
"source.duplicate_name",
"Logical source names must be unique: "
f"{', '.join(duplicate_source_names)}.",
) )
) )
if len(output_nodes) != 1: if len(output_nodes) != 1:
@@ -99,16 +153,29 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
) )
) )
continue continue
if node.type.startswith("source."): definition = node_definition(node.type)
if incoming.get(node.id): node_edges = incoming.get(node.id, [])
diagnostics.append( if definition is not None:
_error("node.source_has_input", "Source nodes cannot have incoming edges.", node_id=node.id) for port in definition.input_ports:
) connections = [
elif len(incoming.get(node.id, [])) != 1: edge
for edge in node_edges
if edge.target_port == port.id
]
minimum = port.minimum_connections if port.required else 0
if len(connections) < minimum:
diagnostics.append( diagnostics.append(
_error( _error(
"node.input_count", "node.input_required",
"This transform requires exactly one incoming edge.", f"{definition.label} requires {port.label.lower()} input.",
node_id=node.id,
)
)
if not port.multiple and len(connections) > 1:
diagnostics.append(
_error(
"node.input_multiple",
f"{port.label} accepts only one connection.",
node_id=node.id, node_id=node.id,
) )
) )
@@ -118,12 +185,18 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
if cyclic: if cyclic:
diagnostics.append(_error("graph.cycle", "Pipeline edges must form an acyclic graph.")) diagnostics.append(_error("graph.cycle", "Pipeline edges must form an acyclic graph."))
elif source_nodes and output_nodes: elif source_nodes and output_nodes:
reachable = _reachable_from(source_nodes[0].id, outgoing) reachable: set[str] = set()
for source in source_nodes:
reachable.update(_reachable_from(source.id, outgoing))
if len(reachable) != len(nodes): if len(reachable) != len(nodes):
diagnostics.append( diagnostics.append(
_error("graph.disconnected", "Every node must be connected to the pipeline source.") _error("graph.disconnected", "Every node must be connected to a pipeline source.")
) )
reaches_output = _reachable_from(output_nodes[0].id, incoming) reverse_adjacency = {
node_id: [edge.source for edge in edges]
for node_id, edges in incoming.items()
}
reaches_output = _reachable_from(output_nodes[0].id, reverse_adjacency)
if len(reaches_output) != len(nodes): if len(reaches_output) != len(nodes):
diagnostics.append( diagnostics.append(
_error("graph.dead_end", "Every node must lead to the pipeline output.") _error("graph.dead_end", "Every node must lead to the pipeline output.")
@@ -132,6 +205,7 @@ def validate_graph(graph: PipelineGraph) -> list[DataflowDiagnostic]:
diagnostics.append( diagnostics.append(
_error("graph.output_not_terminal", "The output node must be the terminal transform.") _error("graph.output_not_terminal", "The output node must be the terminal transform.")
) )
diagnostics.extend(_validate_graph_schemas(graph, ordered=ordered))
return diagnostics return diagnostics
@@ -156,8 +230,11 @@ def topological_order(graph: PipelineGraph) -> tuple[list[str], bool]:
return ordered, len(ordered) != len(node_ids) return ordered, len(ordered) != len(node_ids)
def graph_input_map(graph: PipelineGraph) -> dict[str, str]: def graph_inputs_by_port(graph: PipelineGraph) -> dict[str, dict[str, list[str]]]:
return {edge.target: edge.source for edge in graph.edges} result: dict[str, dict[str, list[str]]] = {}
for edge in graph.edges:
result.setdefault(edge.target, {}).setdefault(edge.target_port, []).append(edge.source)
return result
def _reachable_from(start: str, adjacency: dict[str, list[str]]) -> set[str]: def _reachable_from(start: str, adjacency: dict[str, list[str]]) -> set[str]:
@@ -172,6 +249,306 @@ def _reachable_from(start: str, adjacency: dict[str, list[str]]) -> set[str]:
return seen return seen
@dataclass(frozen=True)
class _SchemaState:
columns: frozenset[str]
open: bool = False
def knows(self, column: str) -> bool:
return self.open or column in self.columns
def _validate_graph_schemas(
graph: PipelineGraph,
*,
ordered: list[str],
) -> list[DataflowDiagnostic]:
diagnostics: list[DataflowDiagnostic] = []
node_by_id = {node.id: node for node in graph.nodes}
inputs = graph_inputs_by_port(graph)
schemas: dict[str, _SchemaState] = {}
for node_id in ordered:
node = node_by_id[node_id]
node_inputs = inputs.get(node.id, {})
input_states = [
schemas[source_id]
for port_sources in node_inputs.values()
for source_id in port_sources
if source_id in 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,
)
continue
if node.type == "source.reference":
columns = _configured_source_columns(node.config.get("source_columns"))
schemas[node.id] = _SchemaState(
frozenset(columns),
open=not columns,
)
continue
if node.type == "combine.union":
if len(input_states) > 1:
closed_shapes = {
state.columns
for state in input_states
if not state.open
}
if len(closed_shapes) > 1:
diagnostics.append(
_warning(
"union.schema_mismatch",
"Appended inputs use different columns; missing values will be null.",
node_id=node.id,
)
)
schemas[node.id] = _SchemaState(
frozenset().union(*(state.columns for state in input_states)),
open=any(state.open for state in input_states),
)
continue
if node.type == "combine.join":
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",
)
prefix = str(node.config.get("right_prefix", "right_"))
prefixed_right = {f"{prefix}{column}" for column in right_state.columns}
collisions = left_state.columns & prefixed_right
if collisions:
diagnostics.append(
_error(
"join.output_collision",
f"Join output columns collide: {', '.join(sorted(collisions))}.",
node_id=node.id,
field="right_prefix",
)
)
schemas[node.id] = _SchemaState(
frozenset(left_state.columns | prefixed_right),
open=left_state.open or right_state.open,
)
continue
if node.type == "filter":
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=[node.config.get("column")],
field="column",
)
schemas[node.id] = input_state
continue
if node.type == "distinct":
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=node.config.get("columns"),
field="columns",
)
schemas[node.id] = input_state
continue
if node.type == "select":
fields = node.config.get("fields")
selected_columns: list[str] = []
output_columns: list[str] = []
if isinstance(fields, list):
for field in fields:
if isinstance(field, str):
selected_columns.append(field)
output_columns.append(field)
elif isinstance(field, dict):
column = field.get("column")
alias = field.get("alias") or column
if isinstance(column, str):
selected_columns.append(column)
if isinstance(alias, str):
output_columns.append(alias)
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=selected_columns,
field="fields",
)
_validate_output_names(
diagnostics,
node=node,
columns=output_columns,
field="fields",
)
schemas[node.id] = _SchemaState(frozenset(output_columns))
continue
if node.type == "derive":
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=node.config.get("source_columns"),
field="source_columns",
)
target = node.config.get("target_column")
if isinstance(target, str) and target:
if target in input_state.columns:
diagnostics.append(
_warning(
"derive.overwrites_column",
f"Derived column {target!r} replaces an existing value.",
node_id=node.id,
field="target_column",
)
)
schemas[node.id] = _SchemaState(
input_state.columns | frozenset((target,)),
open=input_state.open,
)
else:
schemas[node.id] = input_state
continue
if node.type == "aggregate":
group_by = _text_items(node.config.get("group_by"))
aggregates = node.config.get("aggregates")
aggregate_columns = [
str(item.get("column"))
for item in aggregates
if isinstance(aggregates, list)
and isinstance(item, dict)
and item.get("column") not in (None, "", "*")
] if isinstance(aggregates, list) else []
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=[*group_by, *aggregate_columns],
field="aggregates",
)
aliases = [
str(item.get("alias"))
for item in aggregates
if isinstance(aggregates, list)
and isinstance(item, dict)
and item.get("alias")
] if isinstance(aggregates, list) else []
output_columns = [*group_by, *aliases]
_validate_output_names(
diagnostics,
node=node,
columns=output_columns,
field="aggregates",
)
schemas[node.id] = _SchemaState(frozenset(output_columns))
continue
if node.type == "sort":
fields = node.config.get("fields")
columns = [
str(item.get("column"))
for item in fields
if isinstance(fields, list)
and isinstance(item, dict)
and item.get("column")
] if isinstance(fields, list) else []
_validate_columns(
diagnostics,
node=node,
state=input_state,
columns=columns,
field="fields",
)
schemas[node.id] = input_state
continue
schemas[node.id] = input_state
return diagnostics
def _configured_source_columns(value: object) -> set[str]:
if not isinstance(value, list):
return set()
return {
item if isinstance(item, str) else str(item.get("name"))
for item in value
if (
isinstance(item, str) and item
or isinstance(item, dict) and item.get("name")
)
}
def _port_schema(
inputs: dict[str, list[str]],
schemas: dict[str, _SchemaState],
port: str,
) -> _SchemaState:
source_ids = inputs.get(port, [])
return schemas.get(source_ids[0], _SchemaState(frozenset(), open=True)) if source_ids else _SchemaState(frozenset(), open=True)
def _validate_columns(
diagnostics: list[DataflowDiagnostic],
*,
node: GraphNode,
state: _SchemaState,
columns: object,
field: str,
) -> None:
for column in _text_items(columns):
if not state.knows(column):
diagnostics.append(
_error(
"schema.unknown_column",
f"Column {column!r} is not available at this node.",
node_id=node.id,
field=field,
)
)
def _validate_output_names(
diagnostics: list[DataflowDiagnostic],
*,
node: GraphNode,
columns: list[str],
field: str,
) -> None:
duplicates = sorted({column for column in columns if columns.count(column) > 1})
if duplicates:
diagnostics.append(
_error(
"schema.duplicate_output",
f"Output column names must be unique: {', '.join(duplicates)}.",
node_id=node.id,
field=field,
)
)
def _text_items(value: object) -> list[str]:
if not isinstance(value, list):
return []
return [item for item in value if isinstance(item, str) and item]
def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]: def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
config = node.config config = node.config
diagnostics: list[DataflowDiagnostic] = [] diagnostics: list[DataflowDiagnostic] = []
@@ -186,7 +563,35 @@ def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
field="source_name", field="source_name",
) )
) )
elif re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", source_name.strip()) is None:
diagnostics.append(
_error(
"source.name_invalid",
"Logical source names must be SQL identifiers, such as monthly_cases.",
node_id=node.id,
field="source_name",
)
)
if node.type == "source.reference": if node.type == "source.reference":
if not _non_empty_text(config.get("source_ref")):
diagnostics.append(
_error(
"source.reference_required",
"Choose a connector source.",
node_id=node.id,
field="source_ref",
)
)
expected_fingerprint = config.get("expected_fingerprint")
if expected_fingerprint is not None and not isinstance(expected_fingerprint, str):
diagnostics.append(
_error(
"source.fingerprint",
"The expected source fingerprint must be text.",
node_id=node.id,
field="expected_fingerprint",
)
)
return diagnostics return diagnostics
rows = config.get("rows") rows = config.get("rows")
if not isinstance(rows, list): if not isinstance(rows, list):
@@ -215,6 +620,72 @@ def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
) )
if operator not in {"is_null", "not_null"} and "value" not in config: 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")) diagnostics.append(_node_field_error(node, "filter.value", "Enter a comparison value.", "value"))
elif node.type == "distinct":
columns = config.get("columns", [])
if not isinstance(columns, list) or any(not _non_empty_text(item) for item in columns):
diagnostics.append(
_node_field_error(
node,
"distinct.columns",
"Distinct key columns must be named.",
"columns",
)
)
elif node.type == "combine.union":
if config.get("mode", "all") not in {"all", "distinct"}:
diagnostics.append(
_node_field_error(
node,
"union.mode",
"Choose whether duplicate rows are kept or removed.",
"mode",
)
)
elif node.type == "combine.join":
if config.get("join_type", "inner") not in JOIN_TYPES:
diagnostics.append(
_node_field_error(node, "join.type", "Choose a supported join type.", "join_type")
)
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, "join.left_keys", "Add at least one left 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, "join.right_keys", "Add at least one right 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,
"join.key_count",
"Left and right joins need the same number of keys.",
"right_keys",
)
)
right_prefix = config.get("right_prefix", "right_")
if (
not _non_empty_text(right_prefix)
or re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*_", str(right_prefix)) is None
):
diagnostics.append(
_node_field_error(
node,
"join.right_prefix",
"Enter an identifier prefix ending in an underscore, such as right_.",
"right_prefix",
)
)
elif node.type == "select": elif node.type == "select":
fields = config.get("fields") fields = config.get("fields")
if not isinstance(fields, list) or not fields: if not isinstance(fields, list) or not fields:
@@ -267,6 +738,58 @@ def _validate_node_config(node: GraphNode) -> list[DataflowDiagnostic]:
_node_field_error(node, "aggregate.alias", "Every aggregate needs an alias.", "aggregates") _node_field_error(node, "aggregate.alias", "Every aggregate needs an alias.", "aggregates")
) )
break break
elif node.type == "derive":
if not _non_empty_text(config.get("target_column")):
diagnostics.append(
_node_field_error(
node,
"derive.target_column",
"Choose an output column.",
"target_column",
)
)
operation = config.get("operation")
if operation not in DERIVE_OPERATIONS:
diagnostics.append(
_node_field_error(
node,
"derive.operation",
"Choose a supported derive operation.",
"operation",
)
)
source_columns = config.get("source_columns")
if (
not isinstance(source_columns, list)
or not source_columns
or any(not _non_empty_text(item) for item in source_columns)
):
diagnostics.append(
_node_field_error(
node,
"derive.source_columns",
"Choose at least one source column.",
"source_columns",
)
)
if operation in {"copy", "upper", "lower", "trim"} and isinstance(source_columns, list) and len(source_columns) != 1:
diagnostics.append(
_node_field_error(
node,
"derive.source_count",
"This operation requires exactly one source column.",
"source_columns",
)
)
if operation in {"add", "subtract", "multiply", "divide"} and isinstance(source_columns, list) and len(source_columns) != 2:
diagnostics.append(
_node_field_error(
node,
"derive.numeric_source_count",
"Numeric operations require exactly two source columns.",
"source_columns",
)
)
elif node.type == "sort": elif node.type == "sort":
fields = config.get("fields") fields = config.get("fields")
if not isinstance(fields, list) or not fields: if not isinstance(fields, list) or not fields:
@@ -313,13 +836,31 @@ def _error(
) )
def _warning(
code: str,
message: str,
*,
node_id: str | None = None,
field: str | None = None,
) -> DataflowDiagnostic:
return DataflowDiagnostic(
severity="warning",
code=code,
message=message,
node_id=node_id,
field=field,
)
__all__ = [ __all__ = [
"AGGREGATE_FUNCTIONS", "AGGREGATE_FUNCTIONS",
"DERIVE_OPERATIONS",
"FILTER_OPERATORS", "FILTER_OPERATORS",
"JOIN_TYPES",
"SUPPORTED_NODE_TYPES", "SUPPORTED_NODE_TYPES",
"canonical_graph_payload", "canonical_graph_payload",
"definition_hash", "definition_hash",
"graph_input_map", "graph_inputs_by_port",
"topological_order", "topological_order",
"validate_graph", "validate_graph",
] ]

View File

@@ -11,11 +11,17 @@ from govoplan_core.core.modules import (
FrontendModule, FrontendModule,
MigrationSpec, MigrationSpec,
ModuleInterfaceProvider, ModuleInterfaceProvider,
ModuleInterfaceRequirement,
ModuleContext,
ModuleManifest, ModuleManifest,
NavItem, NavItem,
PermissionDefinition, PermissionDefinition,
RoleTemplate, RoleTemplate,
) )
from govoplan_core.core.tabular_sources import (
CAPABILITY_CONNECTORS_TABULAR_SNAPSHOT_WRITER,
CAPABILITY_CONNECTORS_TABULAR_SOURCES,
)
from govoplan_core.db.base import Base from govoplan_core.db.base import Base
from govoplan_dataflow.backend.db import models as dataflow_models from govoplan_dataflow.backend.db import models as dataflow_models
@@ -116,14 +122,20 @@ DOCUMENTATION = (
"audit", "audit",
), ),
metadata={ metadata={
"first_slice": "Inline source, filter, select, aggregate, sort, limit, output, revisioning, and bounded preview.", "first_slice": (
"Inline and connector sources, union, join, filter, deduplication, select, "
"derived columns, aggregate, sort, limit, output, revisioning, and bounded preview."
),
"sql_safety": "Constrained AST compilation only; no pass-through execution.", "sql_safety": "Constrained AST compilation only; no pass-through execution.",
}, },
), ),
) )
def _dataflow_router(_context): def _dataflow_router(context: ModuleContext):
from govoplan_dataflow.backend.runtime import configure_runtime
configure_runtime(registry=context.registry, settings=context.settings)
from govoplan_dataflow.backend.router import router from govoplan_dataflow.backend.router import router
return router return router
@@ -163,12 +175,30 @@ manifest = ModuleManifest(
"risk_compliance", "risk_compliance",
"workflow", "workflow",
), ),
optional_capabilities=(
CAPABILITY_CONNECTORS_TABULAR_SOURCES,
CAPABILITY_CONNECTORS_TABULAR_SNAPSHOT_WRITER,
),
provides_interfaces=( provides_interfaces=(
ModuleInterfaceProvider(name="dataflow.pipeline_catalog", version=MODULE_VERSION), ModuleInterfaceProvider(name="dataflow.pipeline_catalog", version=MODULE_VERSION),
ModuleInterfaceProvider(name="dataflow.pipeline_preview", version=MODULE_VERSION), ModuleInterfaceProvider(name="dataflow.pipeline_preview", version=MODULE_VERSION),
ModuleInterfaceProvider(name="dataflow.run_lifecycle", version=MODULE_VERSION), ModuleInterfaceProvider(name="dataflow.run_lifecycle", version=MODULE_VERSION),
ModuleInterfaceProvider(name="dataflow.dataset_output", version=MODULE_VERSION), ModuleInterfaceProvider(name="dataflow.dataset_output", version=MODULE_VERSION),
), ),
requires_interfaces=(
ModuleInterfaceRequirement(
name="connectors.tabular_sources",
version_min="0.1.0",
version_max_exclusive="1.0.0",
optional=True,
),
ModuleInterfaceRequirement(
name="connectors.tabular_snapshot_writer",
version_min="0.1.0",
version_max_exclusive="1.0.0",
optional=True,
),
),
permissions=PERMISSIONS, permissions=PERMISSIONS,
role_templates=ROLE_TEMPLATES, role_templates=ROLE_TEMPLATES,
nav_items=( nav_items=(

View File

@@ -0,0 +1,305 @@
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
NodeCategory = Literal["load", "combine", "filter", "transform", "output"]
SqlSupport = Literal["full", "partial", "none"]
@dataclass(frozen=True, slots=True)
class NodePortDefinition:
id: str
label: str
required: bool = True
multiple: bool = False
minimum_connections: int = 1
@dataclass(frozen=True, slots=True)
class NodeConfigField:
id: str
label: str
kind: str
required: bool = False
description: str | None = None
options: tuple[tuple[str, str], ...] = ()
@dataclass(frozen=True, slots=True)
class NodeTypeDefinition:
type: str
category: NodeCategory
label: str
description: str
icon: str
input_ports: tuple[NodePortDefinition, ...] = ()
output_ports: tuple[NodePortDefinition, ...] = (
NodePortDefinition(id="output", label="Output"),
)
config_fields: tuple[NodeConfigField, ...] = ()
default_config: dict[str, Any] = field(default_factory=dict)
sql_support: SqlSupport = "full"
NODE_LIBRARY = (
NodeTypeDefinition(
type="source.inline",
category="load",
label="Inline data",
description="Enter a small JSON table directly in the pipeline.",
icon="braces",
config_fields=(
NodeConfigField(id="source_name", label="SQL source name", kind="text", required=True),
NodeConfigField(id="rows", label="Rows", kind="json_rows", required=True),
),
default_config={"source_name": "inline_source", "rows": []},
),
NodeTypeDefinition(
type="source.reference",
category="load",
label="Connector source",
description="Load a bounded, fingerprinted table exposed by Connectors.",
icon="database",
config_fields=(
NodeConfigField(id="source_ref", label="Source", kind="tabular_source", required=True),
NodeConfigField(id="source_name", label="SQL source name", kind="text", required=True),
NodeConfigField(id="expected_fingerprint", label="Expected fingerprint", kind="readonly"),
),
default_config={
"source_ref": "",
"source_name": "connector_source",
"expected_fingerprint": "",
},
),
NodeTypeDefinition(
type="combine.union",
category="combine",
label="Append rows",
description="Append two or more inputs by column name.",
icon="combine",
input_ports=(
NodePortDefinition(
id="input",
label="Inputs",
multiple=True,
minimum_connections=2,
),
),
config_fields=(
NodeConfigField(
id="mode",
label="Duplicates",
kind="select",
options=(("all", "Keep all"), ("distinct", "Remove duplicates")),
),
),
default_config={"mode": "all"},
sql_support="partial",
),
NodeTypeDefinition(
type="combine.join",
category="combine",
label="Join tables",
description="Match two inputs using one or more key columns.",
icon="git-merge",
input_ports=(
NodePortDefinition(id="left", label="Left"),
NodePortDefinition(id="right", label="Right"),
),
config_fields=(
NodeConfigField(
id="join_type",
label="Join type",
kind="select",
options=(
("inner", "Matching rows"),
("left", "All left rows"),
("right", "All right rows"),
("full", "All rows"),
),
),
NodeConfigField(id="left_keys", label="Left keys", kind="column_list", required=True),
NodeConfigField(id="right_keys", label="Right keys", kind="column_list", required=True),
NodeConfigField(id="right_prefix", label="Right-column prefix", kind="text", required=True),
),
default_config={
"join_type": "inner",
"left_keys": [""],
"right_keys": [""],
"right_prefix": "right_",
},
sql_support="partial",
),
NodeTypeDefinition(
type="filter",
category="filter",
label="Filter rows",
description="Keep rows that satisfy a comparison.",
icon="filter",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="column", label="Column", kind="column", required=True),
NodeConfigField(
id="operator",
label="Operator",
kind="select",
required=True,
options=(
("eq", "Equals"),
("ne", "Does not equal"),
("gt", "Greater than"),
("gte", "Greater than or equal"),
("lt", "Less than"),
("lte", "Less than or equal"),
("contains", "Contains"),
("is_null", "Is empty"),
("not_null", "Is not empty"),
),
),
NodeConfigField(id="value", label="Value", kind="value"),
),
default_config={"column": "", "operator": "eq", "value": ""},
),
NodeTypeDefinition(
type="distinct",
category="filter",
label="Remove duplicates",
description="Keep the first row for each selected key.",
icon="list-filter",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(
id="columns",
label="Key columns",
kind="column_list",
description="Leave empty to compare complete rows.",
),
),
default_config={"columns": []},
sql_support="partial",
),
NodeTypeDefinition(
type="select",
category="transform",
label="Select columns",
description="Choose, order, and rename output columns.",
icon="columns-3",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="fields", label="Fields", kind="field_mapping", required=True),
),
default_config={"fields": [{"column": "", "alias": ""}]},
),
NodeTypeDefinition(
type="derive",
category="transform",
label="Derive column",
description="Create a column through a constrained reusable operation.",
icon="variable",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="target_column", label="Output column", kind="text", required=True),
NodeConfigField(
id="operation",
label="Operation",
kind="select",
required=True,
options=(
("copy", "Copy"),
("upper", "Uppercase"),
("lower", "Lowercase"),
("trim", "Trim whitespace"),
("concat", "Concatenate"),
("coalesce", "First non-empty"),
("add", "Add"),
("subtract", "Subtract"),
("multiply", "Multiply"),
("divide", "Divide"),
),
),
NodeConfigField(id="source_columns", label="Source columns", kind="column_list", required=True),
NodeConfigField(id="separator", label="Separator", kind="text"),
),
default_config={
"target_column": "",
"operation": "copy",
"source_columns": [""],
"separator": " ",
},
sql_support="partial",
),
NodeTypeDefinition(
type="aggregate",
category="transform",
label="Aggregate",
description="Group rows and calculate counts or numeric summaries.",
icon="sigma",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(
NodeConfigField(id="group_by", label="Group by", kind="column_list"),
NodeConfigField(id="aggregates", label="Calculations", kind="aggregates", required=True),
),
default_config={
"group_by": [],
"aggregates": [{"function": "count", "column": "*", "alias": "row_count"}],
},
),
NodeTypeDefinition(
type="sort",
category="transform",
label="Sort rows",
description="Order rows by one or more columns.",
icon="arrow-up-down",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(NodeConfigField(id="fields", label="Sort fields", kind="sort_fields", required=True),),
default_config={"fields": [{"column": "", "direction": "asc"}]},
),
NodeTypeDefinition(
type="limit",
category="transform",
label="Limit rows",
description="Keep only the first number of rows.",
icon="list-end",
input_ports=(NodePortDefinition(id="input", label="Input"),),
config_fields=(NodeConfigField(id="count", label="Rows", kind="number", required=True),),
default_config={"count": 100},
),
NodeTypeDefinition(
type="output",
category="output",
label="Preview output",
description="Expose the terminal table for preview or publication.",
icon="panel-top-open",
input_ports=(NodePortDefinition(id="input", label="Input"),),
output_ports=(),
default_config={},
),
)
NODE_TYPES = {definition.type: definition for definition in NODE_LIBRARY}
CATEGORY_LABELS = {
"load": "Load",
"combine": "Combine",
"filter": "Filter",
"transform": "Transform",
"output": "Output",
}
def node_definition(node_type: str) -> NodeTypeDefinition | None:
return NODE_TYPES.get(node_type)
__all__ = [
"CATEGORY_LABELS",
"NODE_LIBRARY",
"NODE_TYPES",
"NodeCategory",
"NodeConfigField",
"NodePortDefinition",
"NodeTypeDefinition",
"SqlSupport",
"node_definition",
]

View File

@@ -5,9 +5,24 @@ from sqlalchemy.orm import Session
from govoplan_core.audit.logging import audit_event from govoplan_core.audit.logging import audit_event
from govoplan_core.auth import ApiPrincipal, get_api_principal, has_scope from govoplan_core.auth import ApiPrincipal, get_api_principal, has_scope
from govoplan_core.core.tabular_sources import (
TabularSnapshotInput,
TabularSource,
TabularSourceAccessError,
TabularSourceError,
TabularSourceNotFoundError,
TabularSourceValidationError,
parse_tabular_csv,
tabular_snapshot_writer,
tabular_source_provider,
)
from govoplan_core.db.session import get_session from govoplan_core.db.session import get_session
from govoplan_dataflow.backend.manifest import ADMIN_SCOPE, READ_SCOPE, RUN_SCOPE, WRITE_SCOPE from govoplan_dataflow.backend.manifest import ADMIN_SCOPE, READ_SCOPE, RUN_SCOPE, WRITE_SCOPE
from govoplan_dataflow.backend.node_library import CATEGORY_LABELS, NODE_LIBRARY
from govoplan_dataflow.backend.runtime import get_registry
from govoplan_dataflow.backend.schemas import ( from govoplan_dataflow.backend.schemas import (
NodeLibraryResponse,
NodeTypeDefinitionResponse,
PipelineCreateRequest, PipelineCreateRequest,
PipelineDeleteResponse, PipelineDeleteResponse,
PipelineDraftRequest, PipelineDraftRequest,
@@ -18,6 +33,10 @@ from govoplan_dataflow.backend.schemas import (
PipelineSqlResponse, PipelineSqlResponse,
PipelineUpdateRequest, PipelineUpdateRequest,
PipelineValidationResponse, PipelineValidationResponse,
TabularSnapshotCreateRequest,
TabularSourceColumnResponse,
TabularSourceListResponse,
TabularSourceResponse,
) )
from govoplan_dataflow.backend.service import ( from govoplan_dataflow.backend.service import (
DataflowConflictError, DataflowConflictError,
@@ -74,6 +93,191 @@ def _http_error(exc: DataflowError) -> HTTPException:
return HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) return HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc))
def _source_http_error(exc: TabularSourceError) -> HTTPException:
if isinstance(exc, TabularSourceAccessError):
return HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=str(exc))
if isinstance(exc, TabularSourceNotFoundError):
return HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc))
if isinstance(exc, TabularSourceValidationError):
return HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
detail=str(exc),
)
return HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc))
def _node_library_response() -> NodeLibraryResponse:
return NodeLibraryResponse(
nodes=[
NodeTypeDefinitionResponse(
type=definition.type,
category=definition.category,
category_label=CATEGORY_LABELS[definition.category],
label=definition.label,
description=definition.description,
icon=definition.icon,
input_ports=[
{
"id": port.id,
"label": port.label,
"required": port.required,
"multiple": port.multiple,
"minimum_connections": port.minimum_connections,
}
for port in definition.input_ports
],
output_ports=[
{
"id": port.id,
"label": port.label,
"required": port.required,
"multiple": port.multiple,
"minimum_connections": port.minimum_connections,
}
for port in definition.output_ports
],
config_fields=[
{
"id": field.id,
"label": field.label,
"kind": field.kind,
"required": field.required,
"description": field.description,
"options": list(field.options),
}
for field in definition.config_fields
],
default_config=dict(definition.default_config),
sql_support=definition.sql_support,
)
for definition in NODE_LIBRARY
]
)
def _source_response(source: TabularSource) -> TabularSourceResponse:
return TabularSourceResponse(
ref=source.ref,
provider=source.provider,
source_name=source.source_name,
name=source.name,
description=source.description,
columns=[
TabularSourceColumnResponse(
name=column.name,
data_type=column.data_type,
nullable=column.nullable,
)
for column in source.schema
],
schema_version=source.schema_version,
fingerprint=source.fingerprint,
row_count=source.row_count,
byte_count=source.byte_count,
updated_at=source.updated_at,
capabilities=list(source.capabilities),
)
@router.get("/node-types", response_model=NodeLibraryResponse)
def api_node_types(
principal: ApiPrincipal = Depends(get_api_principal),
) -> NodeLibraryResponse:
_require_any_scope(principal, READ_SCOPE, WRITE_SCOPE, RUN_SCOPE, ADMIN_SCOPE)
return _node_library_response()
@router.get("/sources", response_model=TabularSourceListResponse)
def api_list_sources(
query: str = "",
session: Session = Depends(get_session),
principal: ApiPrincipal = Depends(get_api_principal),
) -> TabularSourceListResponse:
_require_any_scope(principal, READ_SCOPE, WRITE_SCOPE, RUN_SCOPE, ADMIN_SCOPE)
registry = get_registry()
provider = tabular_source_provider(registry)
writer = tabular_snapshot_writer(registry)
if provider is None:
return TabularSourceListResponse(available=False, writable=False, sources=[])
try:
sources = provider.list_sources(
session,
principal,
query=query,
limit=100,
)
except TabularSourceError as exc:
raise _source_http_error(exc) from exc
return TabularSourceListResponse(
available=True,
writable=writer is not None,
sources=[_source_response(source) for source in sources],
)
@router.post(
"/sources/snapshots",
response_model=TabularSourceResponse,
status_code=status.HTTP_201_CREATED,
)
def api_create_source_snapshot(
payload: TabularSnapshotCreateRequest,
session: Session = Depends(get_session),
principal: ApiPrincipal = Depends(get_api_principal),
) -> TabularSourceResponse:
_require_any_scope(principal, WRITE_SCOPE, ADMIN_SCOPE)
writer = tabular_snapshot_writer(get_registry())
if writer is None:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail="No tabular snapshot writer is available.",
)
try:
rows = (
parse_tabular_csv(
payload.csv_text or "",
delimiter=payload.delimiter,
max_rows=10_000,
)
if payload.format == "csv"
else tuple(payload.rows or ())
)
source = writer.create_snapshot(
session,
principal,
snapshot=TabularSnapshotInput(
name=payload.name,
source_name=payload.source_name,
description=payload.description,
rows=rows,
metadata={
"created_via": "dataflow",
"source_format": payload.format,
},
),
)
except TabularSourceError as exc:
raise _source_http_error(exc) from exc
audit_event(
session,
tenant_id=principal.tenant_id,
user_id=getattr(principal.user, "id", None),
api_key_id=principal.api_key_id,
action="dataflow.source_snapshot.created",
object_type="tabular_source",
object_id=source.ref,
details={
"provider": source.provider,
"source_name": source.source_name,
"row_count": source.row_count,
"fingerprint": source.fingerprint,
},
)
response = _source_response(source)
session.commit()
return response
@router.get("/pipelines", response_model=PipelineListResponse) @router.get("/pipelines", response_model=PipelineListResponse)
def api_list_pipelines( def api_list_pipelines(
session: Session = Depends(get_session), session: Session = Depends(get_session),
@@ -237,6 +441,8 @@ def api_preview_pipeline(
session, session,
tenant_id=principal.tenant_id, tenant_id=principal.tenant_id,
actor_id=_actor_id(principal), actor_id=_actor_id(principal),
principal=principal,
registry=get_registry(),
payload=payload, payload=payload,
) )
except DataflowError as exc: except DataflowError as exc:

View File

@@ -0,0 +1,11 @@
from __future__ import annotations
from govoplan_core.core.runtime import ModuleRuntimeState
_runtime = ModuleRuntimeState("Dataflow")
configure_runtime = _runtime.configure_runtime
get_registry = _runtime.get_registry
get_settings = _runtime.get_settings
settings = _runtime.settings

View File

@@ -4,7 +4,7 @@ import math
from datetime import datetime from datetime import datetime
from typing import Any, Literal from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
PipelineStatus = Literal["draft", "active", "archived"] PipelineStatus = Literal["draft", "active", "archived"]
@@ -152,6 +152,8 @@ class PipelinePreviewResponse(BaseModel):
truncated: bool truncated: bool
diagnostics: list[DataflowDiagnostic] diagnostics: list[DataflowDiagnostic]
node_diagnostics: list[NodePreviewDiagnostic] node_diagnostics: list[NodePreviewDiagnostic]
source_fingerprints: list[dict[str, Any]]
input_row_count: int
definition_hash: str definition_hash: str
executor_version: str executor_version: str
@@ -159,3 +161,93 @@ class PipelinePreviewResponse(BaseModel):
class PipelineDeleteResponse(BaseModel): class PipelineDeleteResponse(BaseModel):
deleted: bool deleted: bool
pipeline_id: str pipeline_id: str
class NodePortDefinitionResponse(BaseModel):
id: str
label: str
required: bool
multiple: bool
minimum_connections: int
class NodeConfigFieldResponse(BaseModel):
id: str
label: str
kind: str
required: bool
description: str | None
options: list[tuple[str, str]]
class NodeTypeDefinitionResponse(BaseModel):
type: str
category: str
category_label: str
label: str
description: str
icon: str
input_ports: list[NodePortDefinitionResponse]
output_ports: list[NodePortDefinitionResponse]
config_fields: list[NodeConfigFieldResponse]
default_config: dict[str, Any]
sql_support: str
class NodeLibraryResponse(BaseModel):
nodes: list[NodeTypeDefinitionResponse]
class TabularSourceColumnResponse(BaseModel):
name: str
data_type: str
nullable: bool
class TabularSourceResponse(BaseModel):
ref: str
provider: str
source_name: str
name: str
description: str | None
columns: list[TabularSourceColumnResponse]
schema_version: str
fingerprint: str
row_count: int | None
byte_count: int | None
updated_at: datetime | None
capabilities: list[str]
class TabularSourceListResponse(BaseModel):
available: bool
writable: bool
sources: list[TabularSourceResponse]
class TabularSnapshotCreateRequest(BaseModel):
name: str = Field(min_length=1, max_length=300)
source_name: str = Field(
min_length=1,
max_length=120,
pattern=r"^[A-Za-z_][A-Za-z0-9_]*$",
)
description: str | None = Field(default=None, max_length=4000)
format: Literal["json", "csv"] = "json"
rows: list[dict[str, Any]] | None = Field(default=None, max_length=10_000)
csv_text: str | None = Field(default=None, max_length=5_000_000)
delimiter: str = Field(default=",", min_length=1, max_length=1)
@model_validator(mode="after")
def validate_format_payload(self) -> TabularSnapshotCreateRequest:
if self.format == "json":
if self.rows is None:
raise ValueError("JSON snapshots require rows.")
if self.csv_text is not None:
raise ValueError("JSON snapshots cannot include CSV text.")
else:
if not self.csv_text:
raise ValueError("CSV snapshots require CSV text.")
if self.rows is not None:
raise ValueError("CSV snapshots cannot include JSON rows.")
return self

View File

@@ -5,6 +5,12 @@ from dataclasses import dataclass
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from govoplan_core.auth import ApiPrincipal
from govoplan_core.core.tabular_sources import (
TabularReadRequest,
TabularSourceError,
tabular_source_provider,
)
from govoplan_core.db.base import utcnow from govoplan_core.db.base import utcnow
from govoplan_dataflow.backend.db.models import ( from govoplan_dataflow.backend.db.models import (
DataflowPipeline, DataflowPipeline,
@@ -14,6 +20,7 @@ from govoplan_dataflow.backend.db.models import (
from govoplan_dataflow.backend.executor import ( from govoplan_dataflow.backend.executor import (
EXECUTOR_VERSION, EXECUTOR_VERSION,
PipelineExecutionError, PipelineExecutionError,
ResolvedSource,
execute_preview, execute_preview,
) )
from govoplan_dataflow.backend.graph import canonical_graph_payload, definition_hash, validate_graph from govoplan_dataflow.backend.graph import canonical_graph_payload, definition_hash, validate_graph
@@ -333,6 +340,8 @@ def preview_pipeline(
tenant_id: str, tenant_id: str,
actor_id: str | None, actor_id: str | None,
payload: PipelinePreviewRequest, payload: PipelinePreviewRequest,
principal: ApiPrincipal | None = None,
registry: object | None = None,
) -> PipelinePreviewResponse: ) -> PipelinePreviewResponse:
pipeline: DataflowPipeline | None = None pipeline: DataflowPipeline | None = None
revision: DataflowPipelineRevision | None = None revision: DataflowPipelineRevision | None = None
@@ -360,6 +369,8 @@ def preview_pipeline(
truncated=False, truncated=False,
diagnostics=validated.diagnostics, diagnostics=validated.diagnostics,
node_diagnostics=[], node_diagnostics=[],
source_fingerprints=[],
input_row_count=0,
definition_hash="", definition_hash="",
executor_version=EXECUTOR_VERSION, executor_version=EXECUTOR_VERSION,
) )
@@ -370,7 +381,47 @@ def preview_pipeline(
started_at = utcnow() started_at = utcnow()
run: DataflowRun | None = None run: DataflowRun | None = None
try: try:
result = execute_preview(graph, row_limit=payload.row_limit) provider = tabular_source_provider(registry)
def resolve_source(node: GraphNode, limit: int) -> ResolvedSource:
if provider is None:
raise PipelineExecutionError(
"Connector-backed preview requires the Connectors tabular-source capability.",
node_id=node.id,
)
if principal is None:
raise PipelineExecutionError(
"Connector-backed preview requires a tenant API principal.",
node_id=node.id,
)
try:
resolved = provider.read_source(
session,
principal,
request=TabularReadRequest(
source_ref=str(node.config["source_ref"]),
limit=limit,
expected_fingerprint=_clean_optional(
node.config.get("expected_fingerprint")
),
),
)
except TabularSourceError as exc:
raise PipelineExecutionError(str(exc), node_id=node.id) from exc
return ResolvedSource(
rows=tuple(dict(row) for row in resolved.rows),
source_ref=resolved.source.ref,
provider=resolved.source.provider,
fingerprint=resolved.source.fingerprint,
total_rows=resolved.total_rows,
truncated=resolved.truncated,
)
result = execute_preview(
graph,
row_limit=payload.row_limit,
source_resolver=resolve_source,
)
status = "succeeded" status = "succeeded"
error = None error = None
diagnostics = result.diagnostics diagnostics = result.diagnostics
@@ -385,6 +436,7 @@ def preview_pipeline(
status = "failed" status = "failed"
error = str(exc) error = str(exc)
diagnostics = [ diagnostics = [
*exc.diagnostics,
DataflowDiagnostic( DataflowDiagnostic(
severity="error", severity="error",
code="preview.execution", code="preview.execution",
@@ -396,9 +448,9 @@ def preview_pipeline(
rows = [] rows = []
total_rows = 0 total_rows = 0
truncated = False truncated = False
node_diagnostics = [] node_diagnostics = list(exc.node_diagnostics)
source_fingerprints = [] source_fingerprints = list(exc.source_fingerprints)
input_row_count = 0 input_row_count = exc.input_row_count
if pipeline is not None and revision is not None: if pipeline is not None and revision is not None:
run = DataflowRun( run = DataflowRun(
@@ -433,6 +485,8 @@ def preview_pipeline(
truncated=truncated, truncated=truncated,
diagnostics=diagnostics, diagnostics=diagnostics,
node_diagnostics=node_diagnostics, node_diagnostics=node_diagnostics,
source_fingerprints=source_fingerprints,
input_row_count=input_row_count,
definition_hash=graph_hash, definition_hash=graph_hash,
executor_version=EXECUTOR_VERSION, executor_version=EXECUTOR_VERSION,
) )

View File

@@ -1,12 +1,12 @@
from __future__ import annotations from __future__ import annotations
from typing import Any, Iterable from typing import Any, Callable, Iterable
import sqlglot import sqlglot
from sqlglot import exp from sqlglot import exp
from sqlglot.errors import ParseError from sqlglot.errors import ParseError
from govoplan_dataflow.backend.graph import topological_order, validate_graph from govoplan_dataflow.backend.graph import graph_inputs_by_port, topological_order, validate_graph
from govoplan_dataflow.backend.schemas import ( from govoplan_dataflow.backend.schemas import (
DataflowDiagnostic, DataflowDiagnostic,
GraphEdge, GraphEdge,
@@ -47,49 +47,121 @@ def compile_sql(
query = statements[0] query = statements[0]
_reject_unsupported_query_shape(query) _reject_unsupported_query_shape(query)
tables = list(query.find_all(exp.Table)) from_clause = query.args.get("from_")
if len(tables) != 1: if from_clause is None or not isinstance(from_clause.this, exp.Table):
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.source_count", "The first release supports exactly one logical source.")] [_sql_error("sql.source_required", "SELECT requires a logical tabular source.")]
) )
table = tables[0] joins = list(query.args.get("joins") or [])
if len(joins) > 1:
raise SqlCompilationError(
[_sql_error("sql.join_count", "Dataflow SQL currently supports one two-source join.")]
)
left_table = from_clause.this
right_table = joins[0].this if joins else None
if right_table is not None and not isinstance(right_table, exp.Table):
raise SqlCompilationError(
[_sql_error("sql.join_source", "JOIN requires a logical tabular source.")]
)
tables = [left_table, *([right_table] if isinstance(right_table, exp.Table) else [])]
for table in tables:
if table.catalog or table.db: if table.catalog or table.db:
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.qualified_source", "Use the logical source name without a catalog or schema.")] [_sql_error("sql.qualified_source", "Use logical source names without a catalog or schema.")]
)
source_name = table.name
preserved_source = next(
(
node.model_copy(deep=True)
for node in source_nodes
if node.type.startswith("source.")
and str(node.config.get("source_name", "")).casefold() == source_name.casefold()
),
None,
)
source_node = preserved_source or GraphNode(
id="source",
type="source.reference",
label=source_name,
position=GraphPosition(x=80, y=180),
config={"source_name": source_name},
) )
nodes = [source_node] source_node_list = list(source_nodes)
left_source = _source_node(
left_table,
source_node_list,
fallback_id="source-left" if right_table is not None else "source",
position=GraphPosition(x=60, y=120 if right_table is not None else 180),
)
nodes = [left_source]
edges: list[GraphEdge] = []
previous_node_id = left_source.id
qualifier_prefixes: dict[str, str] | None = None
if isinstance(right_table, exp.Table):
right_source = _source_node(
right_table,
source_node_list,
fallback_id="source-right",
position=GraphPosition(x=60, y=280),
)
if right_source.id == left_source.id:
raise SqlCompilationError(
[_sql_error("sql.source_identity", "Joined sources must use different graph nodes.")]
)
right_prefix = f"{right_table.alias_or_name}_"
join_config = _join_config(joins[0], left_table=left_table, right_table=right_table)
join_config["right_prefix"] = right_prefix
join_node = GraphNode(
id="join",
type="combine.join",
label=f"Join {left_table.name} and {right_table.name}",
position=GraphPosition(x=300, y=200),
config=join_config,
)
nodes.extend((right_source, join_node))
edges.extend(
(
GraphEdge(
id=f"edge-{left_source.id}-{join_node.id}-left",
source=left_source.id,
target=join_node.id,
target_port="left",
),
GraphEdge(
id=f"edge-{right_source.id}-{join_node.id}-right",
source=right_source.id,
target=join_node.id,
target_port="right",
),
)
)
previous_node_id = join_node.id
qualifier_prefixes = _join_qualifier_prefixes(
left_table,
right_table,
right_prefix=right_prefix,
)
def append_transform(node: GraphNode) -> None:
nonlocal previous_node_id
nodes.append(node)
edges.append(
GraphEdge(
id=f"edge-{previous_node_id}-{node.id}",
source=previous_node_id,
target=node.id,
)
)
previous_node_id = node.id
conditions = _flatten_and(query.args.get("where").this) if query.args.get("where") else [] conditions = _flatten_and(query.args.get("where").this) if query.args.get("where") else []
for index, condition in enumerate(conditions, start=1): for index, condition in enumerate(conditions, start=1):
nodes.append( append_transform(
GraphNode( GraphNode(
id=f"filter-{index}", id=f"filter-{index}",
type="filter", type="filter",
label=f"Filter {index}", label=f"Filter {index}",
position=_position(len(nodes)), position=_position(len(nodes)),
config=_condition_config(condition), config=_condition_config(
condition,
qualifier_prefixes=qualifier_prefixes,
),
) )
) )
group = query.args.get("group") group = query.args.get("group")
group_by = [_column_name(item, context="GROUP BY") for item in group.expressions] if group else [] group_by = [
_column_name(
item,
context="GROUP BY",
qualifier_prefixes=qualifier_prefixes,
)
for item in group.expressions
] if group else []
aggregate_specs: list[dict[str, Any]] = [] aggregate_specs: list[dict[str, Any]] = []
projection_fields: list[dict[str, str]] = [] projection_fields: list[dict[str, str]] = []
saw_star = False saw_star = False
@@ -97,7 +169,11 @@ def compile_sql(
for item in query.expressions: for item in query.expressions:
inner = item.this if isinstance(item, exp.Alias) else item inner = item.this if isinstance(item, exp.Alias) else item
alias = item.alias if isinstance(item, exp.Alias) else "" alias = item.alias if isinstance(item, exp.Alias) else ""
aggregate = _aggregate_config(inner, alias=alias) aggregate = _aggregate_config(
inner,
alias=alias,
qualifier_prefixes=qualifier_prefixes,
)
if aggregate is not None: if aggregate is not None:
saw_aggregate = True saw_aggregate = True
aggregate_specs.append(aggregate) aggregate_specs.append(aggregate)
@@ -109,7 +185,11 @@ def compile_sql(
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.select_expression", "SELECT supports columns and COUNT/SUM/AVG/MIN/MAX only.")] [_sql_error("sql.select_expression", "SELECT supports columns and COUNT/SUM/AVG/MIN/MAX only.")]
) )
column = _column_name(inner, context="SELECT") column = _column_name(
inner,
context="SELECT",
qualifier_prefixes=qualifier_prefixes,
)
projection_fields.append({"column": column, "alias": alias or column}) projection_fields.append({"column": column, "alias": alias or column})
if saw_star and len(query.expressions) != 1: if saw_star and len(query.expressions) != 1:
@@ -132,7 +212,7 @@ def compile_sql(
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.aggregate_required", "GROUP BY requires at least one aggregate in this dialect.")] [_sql_error("sql.aggregate_required", "GROUP BY requires at least one aggregate in this dialect.")]
) )
nodes.append( append_transform(
GraphNode( GraphNode(
id="aggregate", id="aggregate",
type="aggregate", type="aggregate",
@@ -142,7 +222,7 @@ def compile_sql(
) )
) )
elif not saw_star: elif not saw_star:
nodes.append( append_transform(
GraphNode( GraphNode(
id="select", id="select",
type="select", type="select",
@@ -152,6 +232,21 @@ def compile_sql(
) )
) )
if query.args.get("distinct"):
if saw_aggregate or group_by:
raise SqlCompilationError(
[_sql_error("sql.distinct_group", "DISTINCT cannot be combined with aggregation yet.")]
)
append_transform(
GraphNode(
id="distinct",
type="distinct",
label="Remove duplicates",
position=_position(len(nodes)),
config={"columns": []},
)
)
order = query.args.get("order") order = query.args.get("order")
if order: if order:
fields: list[dict[str, str]] = [] fields: list[dict[str, str]] = []
@@ -160,11 +255,15 @@ def compile_sql(
raise SqlCompilationError([_sql_error("sql.order", "Unsupported ORDER BY expression.")]) raise SqlCompilationError([_sql_error("sql.order", "Unsupported ORDER BY expression.")])
fields.append( fields.append(
{ {
"column": _column_name(item.this, context="ORDER BY"), "column": _column_name(
item.this,
context="ORDER BY",
qualifier_prefixes=qualifier_prefixes,
),
"direction": "desc" if item.args.get("desc") else "asc", "direction": "desc" if item.args.get("desc") else "asc",
} }
) )
nodes.append( append_transform(
GraphNode( GraphNode(
id="sort", id="sort",
type="sort", type="sort",
@@ -187,7 +286,7 @@ def compile_sql(
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.limit_range", "LIMIT must be between 1 and 100,000.")] [_sql_error("sql.limit_range", "LIMIT must be between 1 and 100,000.")]
) )
nodes.append( append_transform(
GraphNode( GraphNode(
id="limit", id="limit",
type="limit", type="limit",
@@ -197,7 +296,7 @@ def compile_sql(
) )
) )
nodes.append( append_transform(
GraphNode( GraphNode(
id="output", id="output",
type="output", type="output",
@@ -206,14 +305,6 @@ def compile_sql(
config={}, config={},
) )
) )
edges = [
GraphEdge(
id=f"edge-{source.id}-{target.id}",
source=source.id,
target=target.id,
)
for source, target in zip(nodes, nodes[1:])
]
graph = PipelineGraph(nodes=nodes, edges=edges) graph = PipelineGraph(nodes=nodes, edges=edges)
diagnostics = validate_graph(graph) diagnostics = validate_graph(graph)
if any(item.severity == "error" for item in diagnostics): if any(item.severity == "error" for item in diagnostics):
@@ -229,10 +320,57 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
if cyclic: if cyclic:
raise SqlCompilationError([_sql_error("graph.cycle", "A cyclic graph cannot be rendered as SQL.")]) raise SqlCompilationError([_sql_error("graph.cycle", "A cyclic graph cannot be rendered as SQL.")])
node_by_id = {node.id: node for node in graph.nodes} node_by_id = {node.id: node for node in graph.nodes}
source = node_by_id[ordered[0]] source_nodes = [node for node in graph.nodes if node.type.startswith("source.")]
source_name = str(source.config.get("source_name", "")).strip() join_nodes = [node for node in graph.nodes if node.type == "combine.join"]
if not source_name: if len(join_nodes) > 1:
raise SqlCompilationError([_sql_error("source.name_required", "The source needs a logical SQL name.")]) raise SqlCompilationError(
[_sql_error("sql.join_count", "Only one two-source join can be rendered as SQL.")]
)
left_source: GraphNode
right_source: GraphNode | None = None
join_node = join_nodes[0] if join_nodes else None
right_prefix: str | None = None
right_alias: str | None = None
if join_node is not None:
inputs = graph_inputs_by_port(graph).get(join_node.id, {})
left_source = node_by_id[inputs["left"][0]]
right_source = node_by_id[inputs["right"][0]]
right_prefix = str(join_node.config["right_prefix"])
right_alias = right_prefix[:-1]
elif len(source_nodes) == 1:
left_source = source_nodes[0]
else:
raise SqlCompilationError(
[_sql_error("sql.source_count", "SQL rendering needs one source or one two-source join.")]
)
left_source_name = str(left_source.config.get("source_name", "")).strip()
right_source_name = (
str(right_source.config.get("source_name", "")).strip()
if right_source is not None
else None
)
if not left_source_name or (right_source is not None and not right_source_name):
raise SqlCompilationError(
[_sql_error("source.name_required", "Every source needs a logical SQL name.")]
)
if right_alias and right_alias.casefold() == left_source_name.casefold():
raise SqlCompilationError(
[_sql_error("sql.join_alias", "The right-column prefix conflicts with the left source name.")]
)
def column_expression(name: str) -> exp.Column:
if right_prefix and right_alias and name.startswith(right_prefix):
right_name = name[len(right_prefix) :]
if not right_name:
raise SqlCompilationError(
[_sql_error("sql.column", "A right-side column name is missing.")]
)
return exp.column(right_name, table=right_alias)
if right_source is not None:
return exp.column(name, table=left_source_name)
return exp.column(name)
where_conditions: list[exp.Expression] = [] where_conditions: list[exp.Expression] = []
select_expressions: list[exp.Expression] = [exp.Star()] select_expressions: list[exp.Expression] = [exp.Star()]
@@ -240,25 +378,40 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
order_by: list[exp.Expression] = [] order_by: list[exp.Expression] = []
limit: int | None = None limit: int | None = None
selected = False selected = False
distinct = False
for node_id in ordered[1:]: for node_id in ordered:
node = node_by_id[node_id] node = node_by_id[node_id]
if node.type.startswith("source.") or node.type == "combine.join":
continue
if node.type == "filter": if node.type == "filter":
if selected: if selected:
raise SqlCompilationError( raise SqlCompilationError(
[_node_sql_error(node.id, "sql.filter_order", "Filters after projection or aggregation are not representable yet.")] [_node_sql_error(node.id, "sql.filter_order", "Filters after projection or aggregation are not representable yet.")]
) )
where_conditions.append(_condition_expression(node.config)) where_conditions.append(
_condition_expression(node.config, column_expression=column_expression)
)
elif node.type == "select": elif node.type == "select":
if selected: if selected:
raise SqlCompilationError( raise SqlCompilationError(
[_node_sql_error(node.id, "sql.multiple_select", "Only one select or aggregate transform is supported.")] [_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 = [] select_expressions = []
for field in node.config["fields"]: for field in node.config["fields"]:
column = field if isinstance(field, str) else str(field["column"]) column = field if isinstance(field, str) else str(field["column"])
alias = column if isinstance(field, str) else str(field.get("alias") or column) alias = column if isinstance(field, str) else str(field.get("alias") or column)
expression: exp.Expression = exp.column(column) expression: exp.Expression = column_expression(column)
if alias != column: if alias != column:
expression = expression.as_(alias) expression = expression.as_(alias)
select_expressions.append(expression) select_expressions.append(expression)
@@ -268,19 +421,55 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
raise SqlCompilationError( raise SqlCompilationError(
[_node_sql_error(node.id, "sql.multiple_select", "Only one select or aggregate transform is supported.")] [_node_sql_error(node.id, "sql.multiple_select", "Only one select or aggregate transform is supported.")]
) )
select_expressions = [exp.column(column) for column in node.config.get("group_by", [])] if distinct:
group_by = [exp.column(column) for column in node.config.get("group_by", [])] 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"]: for aggregate in node.config["aggregates"]:
function = str(aggregate["function"]) function = str(aggregate["function"])
column = aggregate.get("column") column = aggregate.get("column")
argument: exp.Expression = exp.Star() if function == "count" and column in (None, "", "*") else exp.column(str(column)) argument: exp.Expression = (
exp.Star()
if function == "count" and column in (None, "", "*")
else column_expression(str(column))
)
aggregate_expression = _aggregate_expression(function, argument) aggregate_expression = _aggregate_expression(function, argument)
select_expressions.append(aggregate_expression.as_(str(aggregate["alias"]))) select_expressions.append(aggregate_expression.as_(str(aggregate["alias"])))
selected = True 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": elif node.type == "sort":
order_by = [ order_by = [
exp.Ordered( exp.Ordered(
this=exp.column(str(field["column"])), this=column_expression(str(field["column"])),
desc=field.get("direction", "asc") == "desc", desc=field.get("direction", "asc") == "desc",
nulls_first=False, nulls_first=False,
) )
@@ -293,13 +482,32 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
[_node_sql_error(node.id, "sql.node_not_representable", f"{node.type!r} cannot be rendered as SQL.")] [_node_sql_error(node.id, "sql.node_not_representable", f"{node.type!r} cannot be rendered as SQL.")]
) )
query = exp.select(*select_expressions).from_(exp.to_table(source_name)) query = exp.select(*select_expressions).from_(exp.to_table(left_source_name))
if join_node is not None and right_source_name and right_alias:
join_conditions = [
exp.EQ(
this=exp.column(str(left_key), table=left_source_name),
expression=exp.column(str(right_key), table=right_alias),
)
for left_key, right_key in zip(
join_node.config["left_keys"],
join_node.config["right_keys"],
strict=True,
)
]
query = query.join(
exp.to_table(right_source_name).as_(right_alias),
on=_combine_and(join_conditions),
join_type=str(join_node.config.get("join_type", "inner")),
)
if where_conditions: if where_conditions:
query = query.where(_combine_and(where_conditions)) query = query.where(_combine_and(where_conditions))
if group_by: if group_by:
query = query.group_by(*group_by) query = query.group_by(*group_by)
if order_by: if order_by:
query = query.order_by(*order_by) query = query.order_by(*order_by)
if distinct:
query = query.distinct()
if limit is not None: if limit is not None:
query = query.limit(limit) query = query.limit(limit)
return query.sql(dialect="duckdb", pretty=True), diagnostics return query.sql(dialect="duckdb", pretty=True), diagnostics
@@ -308,7 +516,6 @@ def render_sql(graph: PipelineGraph) -> tuple[str, list[DataflowDiagnostic]]:
def _reject_unsupported_query_shape(query: exp.Select) -> None: def _reject_unsupported_query_shape(query: exp.Select) -> None:
unsupported_args = { unsupported_args = {
"with_": "WITH queries", "with_": "WITH queries",
"distinct": "DISTINCT",
"having": "HAVING", "having": "HAVING",
"qualify": "QUALIFY", "qualify": "QUALIFY",
"offset": "OFFSET", "offset": "OFFSET",
@@ -320,12 +527,126 @@ def _reject_unsupported_query_shape(query: exp.Select) -> None:
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.unsupported_clause", f"{label} are not supported by the first Dataflow dialect.")] [_sql_error("sql.unsupported_clause", f"{label} are not supported by the first Dataflow dialect.")]
) )
if any(True for _ in query.find_all(exp.Join)):
raise SqlCompilationError([_sql_error("sql.join", "JOIN support belongs to the comparison/reconciliation slice.")])
if any(True for _ in query.find_all(exp.Subquery)): if any(True for _ in query.find_all(exp.Subquery)):
raise SqlCompilationError([_sql_error("sql.subquery", "Subqueries are not supported by the first Dataflow dialect.")]) raise SqlCompilationError([_sql_error("sql.subquery", "Subqueries are not supported by the first Dataflow dialect.")])
def _source_node(
table: exp.Table,
source_nodes: list[GraphNode],
*,
fallback_id: str,
position: GraphPosition,
) -> GraphNode:
source_name = table.name
preserved = next(
(
node.model_copy(deep=True)
for node in source_nodes
if node.type.startswith("source.")
and str(node.config.get("source_name", "")).casefold() == source_name.casefold()
),
None,
)
return preserved or GraphNode(
id=fallback_id,
type="source.reference",
label=source_name,
position=position,
config={
"source_ref": "",
"source_name": source_name,
"expected_fingerprint": "",
},
)
def _join_config(
join: exp.Join,
*,
left_table: exp.Table,
right_table: exp.Table,
) -> dict[str, Any]:
kind = str(join.args.get("kind") or "").casefold()
side = str(join.args.get("side") or "").casefold()
if kind not in {"", "inner", "outer"} or side not in {"", "left", "right", "full"}:
raise SqlCompilationError(
[_sql_error("sql.join_type", "JOIN supports INNER, LEFT, RIGHT, or FULL joins only.")]
)
join_type = side or ("inner" if kind in {"", "inner"} else "")
if not join_type:
raise SqlCompilationError(
[_sql_error("sql.join_type", "OUTER JOIN requires LEFT, RIGHT, or FULL.")]
)
on_expression = join.args.get("on")
if on_expression is None:
raise SqlCompilationError(
[_sql_error("sql.join_condition", "JOIN requires an ON key comparison.")]
)
left_qualifiers = _table_qualifiers(left_table)
right_qualifiers = _table_qualifiers(right_table)
if left_qualifiers & right_qualifiers:
raise SqlCompilationError(
[_sql_error("sql.join_alias", "Joined sources need distinct names or aliases.")]
)
left_keys: list[str] = []
right_keys: list[str] = []
for condition in _flatten_and(on_expression):
if (
not isinstance(condition, exp.EQ)
or not isinstance(condition.this, exp.Column)
or not isinstance(condition.expression, exp.Column)
):
raise SqlCompilationError(
[_sql_error("sql.join_condition", "JOIN ON accepts equality comparisons between source columns.")]
)
first = condition.this
second = condition.expression
first_qualifier = first.table.casefold()
second_qualifier = second.table.casefold()
if first_qualifier in left_qualifiers and second_qualifier in right_qualifiers:
left_keys.append(first.name)
right_keys.append(second.name)
elif first_qualifier in right_qualifiers and second_qualifier in left_qualifiers:
left_keys.append(second.name)
right_keys.append(first.name)
else:
raise SqlCompilationError(
[
_sql_error(
"sql.join_qualification",
"Every JOIN key must qualify one left and one right source column.",
)
]
)
return {
"join_type": join_type,
"left_keys": left_keys,
"right_keys": right_keys,
}
def _join_qualifier_prefixes(
left_table: exp.Table,
right_table: exp.Table,
*,
right_prefix: str,
) -> dict[str, str]:
return {
**{qualifier: "" for qualifier in _table_qualifiers(left_table)},
**{qualifier: right_prefix for qualifier in _table_qualifiers(right_table)},
}
def _table_qualifiers(table: exp.Table) -> set[str]:
return {
qualifier.casefold()
for qualifier in (table.name, table.alias_or_name)
if qualifier
}
def _flatten_and(expression: exp.Expression) -> list[exp.Expression]: def _flatten_and(expression: exp.Expression) -> list[exp.Expression]:
if isinstance(expression, exp.And): if isinstance(expression, exp.And):
return [*_flatten_and(expression.this), *_flatten_and(expression.expression)] return [*_flatten_and(expression.this), *_flatten_and(expression.expression)]
@@ -334,14 +655,32 @@ def _flatten_and(expression: exp.Expression) -> list[exp.Expression]:
return [expression] return [expression]
def _condition_config(expression: exp.Expression) -> dict[str, Any]: def _condition_config(
expression: exp.Expression,
*,
qualifier_prefixes: dict[str, str] | None = None,
) -> dict[str, Any]:
if isinstance(expression, exp.Not) and isinstance(expression.this, exp.Is): if isinstance(expression, exp.Not) and isinstance(expression.this, exp.Is):
inner = expression.this inner = expression.this
if isinstance(inner.this, exp.Column) and isinstance(inner.expression, exp.Null): if isinstance(inner.this, exp.Column) and isinstance(inner.expression, exp.Null):
return {"column": _column_name(inner.this, context="WHERE"), "operator": "not_null"} return {
"column": _column_name(
inner.this,
context="WHERE",
qualifier_prefixes=qualifier_prefixes,
),
"operator": "not_null",
}
if isinstance(expression, exp.Is): if isinstance(expression, exp.Is):
if isinstance(expression.this, exp.Column) and isinstance(expression.expression, exp.Null): if isinstance(expression.this, exp.Column) and isinstance(expression.expression, exp.Null):
return {"column": _column_name(expression.this, context="WHERE"), "operator": "is_null"} return {
"column": _column_name(
expression.this,
context="WHERE",
qualifier_prefixes=qualifier_prefixes,
),
"operator": "is_null",
}
mapping: tuple[tuple[type[exp.Expression], str], ...] = ( mapping: tuple[tuple[type[exp.Expression], str], ...] = (
(exp.EQ, "eq"), (exp.EQ, "eq"),
(exp.NEQ, "ne"), (exp.NEQ, "ne"),
@@ -355,7 +694,11 @@ def _condition_config(expression: exp.Expression) -> dict[str, Any]:
if not isinstance(expression.this, exp.Column): if not isinstance(expression.this, exp.Column):
break break
return { return {
"column": _column_name(expression.this, context="WHERE"), "column": _column_name(
expression.this,
context="WHERE",
qualifier_prefixes=qualifier_prefixes,
),
"operator": operator, "operator": operator,
"value": _literal_value(expression.expression), "value": _literal_value(expression.expression),
} }
@@ -366,7 +709,11 @@ def _condition_config(expression: exp.Expression) -> dict[str, Any]:
[_sql_error("sql.like", "LIKE is supported only as a contains pattern: LIKE '%value%'.")] [_sql_error("sql.like", "LIKE is supported only as a contains pattern: LIKE '%value%'.")]
) )
return { return {
"column": _column_name(expression.this, context="WHERE"), "column": _column_name(
expression.this,
context="WHERE",
qualifier_prefixes=qualifier_prefixes,
),
"operator": "contains", "operator": "contains",
"value": value[1:-1], "value": value[1:-1],
} }
@@ -375,8 +722,12 @@ def _condition_config(expression: exp.Expression) -> dict[str, Any]:
) )
def _condition_expression(config: dict[str, Any]) -> exp.Expression: def _condition_expression(
column = exp.column(str(config["column"])) config: dict[str, Any],
*,
column_expression: Callable[[str], exp.Column] = exp.column,
) -> exp.Expression:
column = column_expression(str(config["column"]))
operator = str(config["operator"]) operator = str(config["operator"])
if operator == "is_null": if operator == "is_null":
return exp.Is(this=column, expression=exp.Null()) return exp.Is(this=column, expression=exp.Null())
@@ -417,7 +768,12 @@ def _literal_value(expression: exp.Expression) -> Any:
raise SqlCompilationError([_sql_error("sql.literal", f"Unsupported literal {text!r}.")]) from exc raise SqlCompilationError([_sql_error("sql.literal", f"Unsupported literal {text!r}.")]) from exc
def _aggregate_config(expression: exp.Expression, *, alias: str) -> dict[str, Any] | None: def _aggregate_config(
expression: exp.Expression,
*,
alias: str,
qualifier_prefixes: dict[str, str] | None = None,
) -> dict[str, Any] | None:
mapping: tuple[tuple[type[exp.Expression], str], ...] = ( mapping: tuple[tuple[type[exp.Expression], str], ...] = (
(exp.Count, "count"), (exp.Count, "count"),
(exp.Sum, "sum"), (exp.Sum, "sum"),
@@ -432,7 +788,11 @@ def _aggregate_config(expression: exp.Expression, *, alias: str) -> dict[str, An
if isinstance(argument, exp.Star): if isinstance(argument, exp.Star):
column = "*" column = "*"
elif isinstance(argument, exp.Column): elif isinstance(argument, exp.Column):
column = _column_name(argument, context=function.upper()) column = _column_name(
argument,
context=function.upper(),
qualifier_prefixes=qualifier_prefixes,
)
else: else:
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.aggregate_argument", f"{function.upper()} requires a column or * argument.")] [_sql_error("sql.aggregate_argument", f"{function.upper()} requires a column or * argument.")]
@@ -453,12 +813,37 @@ def _aggregate_expression(function: str, argument: exp.Expression) -> exp.Expres
return mapping[function](this=argument) return mapping[function](this=argument)
def _column_name(expression: exp.Expression, *, context: str) -> str: def _column_name(
if not isinstance(expression, exp.Column) or expression.table: expression: exp.Expression,
*,
context: str,
qualifier_prefixes: dict[str, str] | None = None,
) -> str:
if not isinstance(expression, exp.Column):
raise SqlCompilationError(
[_sql_error("sql.column", f"{context} accepts column names only.")]
)
if not expression.table and qualifier_prefixes is not None:
raise SqlCompilationError(
[
_sql_error(
"sql.join_column_qualification",
f"{context} columns must be source-qualified when a JOIN is present.",
)
]
)
if not expression.table:
return expression.name
if qualifier_prefixes is None:
raise SqlCompilationError( raise SqlCompilationError(
[_sql_error("sql.column", f"{context} accepts unqualified column names only.")] [_sql_error("sql.column", f"{context} accepts unqualified column names only.")]
) )
return expression.name prefix = qualifier_prefixes.get(expression.table.casefold())
if prefix is None:
raise SqlCompilationError(
[_sql_error("sql.column_source", f"{context} references an unknown source qualifier.")]
)
return f"{prefix}{expression.name}"
def _position(index: int) -> GraphPosition: def _position(index: int) -> GraphPosition:

View File

@@ -2,7 +2,13 @@ from __future__ import annotations
import unittest import unittest
from govoplan_dataflow.backend.executor import PipelineExecutionError, execute_preview from govoplan_dataflow.backend.executor import (
MAX_SOURCE_ROWS,
PipelineExecutionError,
ResolvedSource,
execute_preview,
)
from govoplan_dataflow.backend.graph import validate_graph
from govoplan_dataflow.backend.schemas import ( from govoplan_dataflow.backend.schemas import (
GraphEdge, GraphEdge,
GraphNode, GraphNode,
@@ -29,7 +35,66 @@ def inline_source() -> GraphNode:
) )
def lookup_source() -> GraphNode:
return GraphNode(
id="lookup",
type="source.inline",
label="Department lookup",
position=GraphPosition(x=40, y=300),
config={
"source_name": "department_lookup",
"rows": [
{"department": "A", "label": "Administration"},
{"department": "B", "label": "Building services"},
],
},
)
class DataflowGraphAndSqlTests(unittest.TestCase): class DataflowGraphAndSqlTests(unittest.TestCase):
def test_rejects_invalid_or_duplicate_logical_source_names(self) -> None:
invalid_source = inline_source().model_copy(deep=True)
invalid_source.config["source_name"] = "monthly cases"
duplicate_source = lookup_source().model_copy(deep=True)
duplicate_source.config["source_name"] = "MONTHLY_FILES"
graph = PipelineGraph(
nodes=[
invalid_source,
duplicate_source,
GraphNode(
id="union",
type="combine.union",
label="Append rows",
position=GraphPosition(x=250, y=160),
config={"mode": "all"},
),
GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=450, y=160),
config={},
),
],
edges=[
GraphEdge(id="edge-1", source="source", target="union"),
GraphEdge(id="edge-2", source="lookup", target="union"),
GraphEdge(id="edge-3", source="union", target="output"),
],
)
diagnostics = validate_graph(graph)
self.assertTrue(any(item.code == "source.name_invalid" for item in diagnostics))
invalid_source.config["source_name"] = "monthly_files"
diagnostics = validate_graph(
graph.model_copy(
update={"nodes": [invalid_source, *graph.nodes[1:]]},
deep=True,
)
)
self.assertTrue(any(item.code == "source.duplicate_name" for item in diagnostics))
def test_compiles_renders_and_executes_grouped_query(self) -> None: def test_compiles_renders_and_executes_grouped_query(self) -> None:
sql = """ sql = """
SELECT department, COUNT(*) AS records, SUM(amount) AS total SELECT department, COUNT(*) AS records, SUM(amount) AS total
@@ -71,13 +136,207 @@ class DataflowGraphAndSqlTests(unittest.TestCase):
with self.subTest(sql=sql), self.assertRaises(SqlCompilationError): with self.subTest(sql=sql), self.assertRaises(SqlCompilationError):
compile_sql(sql, source_nodes=[inline_source()]) compile_sql(sql, source_nodes=[inline_source()])
def test_rejects_join_until_comparison_slice(self) -> None: def test_compiles_renders_and_executes_two_source_join(self) -> None:
with self.assertRaisesRegex(SqlCompilationError, "JOIN support"): sql = """
compile_sql( SELECT monthly_files.department, lookup.label
"SELECT * FROM monthly_files JOIN other ON monthly_files.id = other.id", FROM monthly_files
source_nodes=[inline_source()], LEFT JOIN department_lookup AS lookup
ON monthly_files.department = lookup.department
ORDER BY monthly_files.department
"""
graph, _, diagnostics = compile_sql(
sql,
source_nodes=[inline_source(), lookup_source()],
) )
self.assertEqual([], diagnostics)
self.assertEqual(
[
"source.inline",
"source.inline",
"combine.join",
"select",
"sort",
"output",
],
[node.type for node in graph.nodes],
)
self.assertEqual("lookup_", graph.nodes[2].config["right_prefix"])
rendered, render_diagnostics = render_sql(graph)
self.assertEqual([], render_diagnostics)
roundtrip, _, _ = compile_sql(
rendered,
source_nodes=[inline_source(), lookup_source()],
)
self.assertEqual(
graph.model_dump(mode="json"),
roundtrip.model_dump(mode="json"),
)
result = execute_preview(graph, row_limit=100)
self.assertEqual(
[
{"department": "A", "lookup_label": "Administration"},
{"department": "A", "lookup_label": "Administration"},
{"department": "B", "lookup_label": "Building services"},
],
result.rows,
)
with self.assertRaisesRegex(SqlCompilationError, "source-qualified"):
compile_sql(
"""
SELECT label
FROM monthly_files
JOIN department_lookup AS lookup
ON monthly_files.department = lookup.department
""",
source_nodes=[inline_source(), lookup_source()],
)
def test_distinct_and_derived_columns_are_deterministic(self) -> None:
source = inline_source().model_copy(
update={
"config": {
"source_name": "monthly_files",
"rows": [
{"first": " Ada ", "last": "Lovelace"},
{"first": " Ada ", "last": "Lovelace"},
],
}
}
)
derive = GraphNode(
id="derive",
type="derive",
label="Display name",
position=GraphPosition(x=260, y=160),
config={
"target_column": "display_name",
"operation": "concat",
"source_columns": ["first", "last"],
"separator": " ",
},
)
distinct = GraphNode(
id="distinct",
type="distinct",
label="Unique people",
position=GraphPosition(x=480, y=160),
config={"columns": ["display_name"]},
)
output = GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=700, y=160),
config={},
)
graph = PipelineGraph(
nodes=[source, derive, distinct, output],
edges=[
GraphEdge(id="source-derive", source=source.id, target=derive.id),
GraphEdge(id="derive-distinct", source=derive.id, target=distinct.id),
GraphEdge(id="distinct-output", source=distinct.id, target=output.id),
],
)
result = execute_preview(graph, row_limit=100)
self.assertEqual(
[{"first": " Ada ", "last": "Lovelace", "display_name": " Ada Lovelace"}],
result.rows,
)
def test_union_can_keep_or_remove_duplicate_rows(self) -> None:
first = inline_source().model_copy(
update={
"id": "first",
"config": {"source_name": "first", "rows": [{"id": 1}, {"id": 2}]},
}
)
second = inline_source().model_copy(
update={
"id": "second",
"config": {"source_name": "second", "rows": [{"id": 2}, {"id": 3}]},
}
)
union = GraphNode(
id="union",
type="combine.union",
label="Append",
position=GraphPosition(x=300, y=200),
config={"mode": "distinct"},
)
output = GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=520, y=200),
config={},
)
graph = PipelineGraph(
nodes=[first, second, union, output],
edges=[
GraphEdge(id="first-union", source=first.id, target=union.id),
GraphEdge(id="second-union", source=second.id, target=union.id),
GraphEdge(id="union-output", source=union.id, target=output.id),
],
)
result = execute_preview(graph, row_limit=100)
self.assertEqual([{"id": 1}, {"id": 2}, {"id": 3}], result.rows)
def test_connector_source_uses_bounded_resolver_and_records_lineage(self) -> None:
source = GraphNode(
id="connector",
type="source.reference",
label="Monthly import",
position=GraphPosition(x=40, y=160),
config={
"source_ref": "snapshot:source-1",
"source_name": "monthly_import",
"expected_fingerprint": "sha256:expected",
},
)
output = GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=280, y=160),
config={},
)
graph = PipelineGraph(
nodes=[source, output],
edges=[GraphEdge(id="source-output", source=source.id, target=output.id)],
)
calls: list[tuple[str, int]] = []
def resolve(node: GraphNode, limit: int) -> ResolvedSource:
calls.append((node.id, limit))
return ResolvedSource(
rows=({"id": 1}, {"id": 2}),
source_ref="snapshot:source-1",
provider="snapshot",
fingerprint="sha256:expected",
total_rows=20,
truncated=True,
)
result = execute_preview(
graph,
row_limit=100,
source_resolver=resolve,
)
self.assertEqual([("connector", MAX_SOURCE_ROWS)], calls)
self.assertEqual([{"id": 1}, {"id": 2}], result.rows)
self.assertEqual(20, result.source_fingerprints[0]["row_count"])
self.assertTrue(result.source_fingerprints[0]["truncated"])
self.assertEqual("source.preview_truncated", result.diagnostics[0].code)
def test_graph_validation_rejects_cycle(self) -> None: def test_graph_validation_rejects_cycle(self) -> None:
source = inline_source() source = inline_source()
output = GraphNode( output = GraphNode(
@@ -95,11 +354,78 @@ class DataflowGraphAndSqlTests(unittest.TestCase):
], ],
) )
from govoplan_dataflow.backend.graph import validate_graph
codes = {item.code for item in validate_graph(graph)} codes = {item.code for item in validate_graph(graph)}
self.assertIn("graph.cycle", codes) self.assertIn("graph.cycle", codes)
def test_schema_validation_reports_unknown_and_colliding_columns(self) -> None:
source = inline_source().model_copy(
update={
"config": {
"source_name": "monthly_files",
"rows": [{"id": 1, "right_id": "reserved"}],
}
}
)
lookup = lookup_source().model_copy(
update={
"config": {
"source_name": "department_lookup",
"rows": [{"id": 1}],
}
}
)
join = GraphNode(
id="join",
type="combine.join",
label="Join",
position=GraphPosition(x=280, y=200),
config={
"join_type": "inner",
"left_keys": ["missing_id"],
"right_keys": ["id"],
"right_prefix": "right_",
},
)
output = GraphNode(
id="output",
type="output",
label="Output",
position=GraphPosition(x=500, y=200),
config={},
)
graph = PipelineGraph(
nodes=[source, lookup, join, output],
edges=[
GraphEdge(
id="source-join",
source=source.id,
target=join.id,
target_port="left",
),
GraphEdge(
id="lookup-join",
source=lookup.id,
target=join.id,
target_port="right",
),
GraphEdge(id="join-output", source=join.id, target=output.id),
],
)
diagnostics = validate_graph(graph)
self.assertTrue(
any(
item.code == "schema.unknown_column"
and item.node_id == join.id
and item.field == "left_keys"
for item in diagnostics
)
)
self.assertTrue(
any(item.code == "join.output_collision" for item in diagnostics)
)
def test_preview_truncates_response_without_changing_pipeline_limit(self) -> None: def test_preview_truncates_response_without_changing_pipeline_limit(self) -> None:
source = inline_source() source = inline_source()
output = GraphNode( output = GraphNode(
@@ -202,6 +528,14 @@ class DataflowGraphAndSqlTests(unittest.TestCase):
self.assertEqual("filter-1", raised.exception.node_id) self.assertEqual("filter-1", raised.exception.node_id)
self.assertIn("Cannot apply", str(raised.exception)) self.assertIn("Cannot apply", str(raised.exception))
self.assertEqual(
[("source", "succeeded"), ("filter-1", "failed")],
[
(diagnostic.node_id, diagnostic.status)
for diagnostic in raised.exception.node_diagnostics
],
)
self.assertEqual(1, len(raised.exception.source_fingerprints))
if __name__ == "__main__": if __name__ == "__main__":

View File

@@ -0,0 +1,41 @@
from __future__ import annotations
import unittest
from govoplan_dataflow.backend.node_library import CATEGORY_LABELS, NODE_LIBRARY, NODE_TYPES
class DataflowNodeLibraryTests(unittest.TestCase):
def test_library_has_unique_executable_nodes_in_every_category(self) -> None:
self.assertEqual(len(NODE_LIBRARY), len(NODE_TYPES))
self.assertEqual(set(CATEGORY_LABELS), {item.category for item in NODE_LIBRARY})
self.assertEqual(
{
"source.inline",
"source.reference",
"combine.union",
"combine.join",
"filter",
"distinct",
"select",
"derive",
"aggregate",
"sort",
"limit",
"output",
},
set(NODE_TYPES),
)
def test_join_and_union_publish_their_connection_contract(self) -> None:
join = NODE_TYPES["combine.join"]
union = NODE_TYPES["combine.union"]
self.assertEqual(["left", "right"], [port.id for port in join.input_ports])
self.assertTrue(union.input_ports[0].multiple)
self.assertEqual(2, union.input_ports[0].minimum_connections)
self.assertEqual([], list(NODE_TYPES["output"].output_ports))
if __name__ == "__main__":
unittest.main()

43
tests/test_schemas.py Normal file
View File

@@ -0,0 +1,43 @@
from __future__ import annotations
import unittest
from govoplan_dataflow.backend.schemas import TabularSnapshotCreateRequest
class DataflowSchemaTests(unittest.TestCase):
def test_snapshot_request_requires_the_selected_format_payload(self) -> None:
json_request = TabularSnapshotCreateRequest(
name="Monthly cases",
source_name="monthly_cases",
format="json",
rows=[{"case_id": "0012"}],
)
csv_request = TabularSnapshotCreateRequest(
name="Monthly cases",
source_name="monthly_cases_csv",
format="csv",
csv_text="case_id;amount\n0012;7.5\n",
delimiter=";",
)
self.assertEqual([{"case_id": "0012"}], json_request.rows)
self.assertEqual(";", csv_request.delimiter)
with self.assertRaises(ValueError):
TabularSnapshotCreateRequest(
name="Missing rows",
source_name="missing_rows",
format="json",
)
with self.assertRaises(ValueError):
TabularSnapshotCreateRequest(
name="Ambiguous",
source_name="ambiguous",
format="csv",
rows=[],
csv_text="id\n1\n",
)
if __name__ == "__main__":
unittest.main()

View File

@@ -220,12 +220,49 @@ class DataflowServiceTests(unittest.TestCase):
self.assertEqual(2, run.output_row_count) self.assertEqual(2, run.output_row_count)
self.assertEqual(3, run.input_row_count) self.assertEqual(3, run.input_row_count)
self.assertEqual(1, len(run.source_fingerprints)) self.assertEqual(1, len(run.source_fingerprints))
self.assertEqual(3, response.input_row_count)
self.assertEqual(run.source_fingerprints, response.source_fingerprints)
self.assertFalse(hasattr(run, "result_rows")) self.assertFalse(hasattr(run, "result_rows"))
self.assertEqual( self.assertEqual(
1, 1,
self.session.scalar(select(func.count()).select_from(DataflowRun)), self.session.scalar(select(func.count()).select_from(DataflowRun)),
) )
def test_failed_preview_keeps_upstream_lineage_and_failed_node_diagnostic(self) -> None:
graph = sample_graph()
graph.nodes[0].config["rows"] = [{"id": 1, "amount": "not-a-number"}]
pipeline = create_pipeline(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
payload=PipelineCreateRequest(
name="Invalid runtime value",
graph=graph,
editor_mode="graph",
),
)
self.session.commit()
response = preview_pipeline(
self.session,
tenant_id="tenant-1",
actor_id="user-1",
payload=PipelinePreviewRequest(pipeline_id=pipeline.id),
)
self.session.commit()
run = self.session.scalar(select(DataflowRun).where(DataflowRun.id == response.run_id))
self.assertEqual("failed", response.status)
self.assertEqual(
[("source", "succeeded"), ("filter", "failed")],
[(item.node_id, item.status) for item in response.node_diagnostics],
)
self.assertEqual(1, run.input_row_count)
self.assertEqual(1, len(run.source_fingerprints))
self.assertEqual(1, response.input_row_count)
self.assertEqual(run.source_fingerprints, response.source_fingerprints)
self.assertEqual("preview.execution", run.diagnostics[-1]["code"])
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()

View File

@@ -2,6 +2,7 @@ import { apiFetch, type ApiSettings } from "@govoplan/core-webui";
export type PipelineStatus = "draft" | "active" | "archived"; export type PipelineStatus = "draft" | "active" | "archived";
export type EditorMode = "graph" | "sql"; export type EditorMode = "graph" | "sql";
export type NodeCategory = "load" | "combine" | "filter" | "transform" | "output";
export type GraphPosition = { x: number; y: number }; export type GraphPosition = { x: number; y: number };
export type PipelineGraphNode = { export type PipelineGraphNode = {
@@ -101,10 +102,109 @@ export type PipelinePreview = {
truncated: boolean; truncated: boolean;
diagnostics: DataflowDiagnostic[]; diagnostics: DataflowDiagnostic[];
node_diagnostics: NodePreviewDiagnostic[]; node_diagnostics: NodePreviewDiagnostic[];
source_fingerprints: Record<string, unknown>[];
input_row_count: number;
definition_hash: string; definition_hash: string;
executor_version: string; executor_version: string;
}; };
export type NodePortDefinition = {
id: string;
label: string;
required: boolean;
multiple: boolean;
minimum_connections: number;
};
export type NodeConfigFieldDefinition = {
id: string;
label: string;
kind: string;
required: boolean;
description?: string | null;
options: Array<[string, string]>;
};
export type NodeTypeDefinition = {
type: string;
category: NodeCategory;
category_label: string;
label: string;
description: string;
icon: string;
input_ports: NodePortDefinition[];
output_ports: NodePortDefinition[];
config_fields: NodeConfigFieldDefinition[];
default_config: Record<string, unknown>;
sql_support: "full" | "partial" | "none";
};
export type TabularSourceColumn = {
name: string;
data_type: string;
nullable: boolean;
};
export type TabularSource = {
ref: string;
provider: string;
source_name: string;
name: string;
description?: string | null;
columns: TabularSourceColumn[];
schema_version: string;
fingerprint: string;
row_count?: number | null;
byte_count?: number | null;
updated_at?: string | null;
capabilities: string[];
};
export type TabularSourceCatalogue = {
available: boolean;
writable: boolean;
sources: TabularSource[];
};
export async function listDataflowNodeTypes(settings: ApiSettings): Promise<NodeTypeDefinition[]> {
const response = await apiFetch<{ nodes: NodeTypeDefinition[] }>(
settings,
"/api/v1/dataflow/node-types"
);
return response.nodes;
}
export function listDataflowSources(
settings: ApiSettings,
query = ""
): Promise<TabularSourceCatalogue> {
const suffix = query.trim() ? `?query=${encodeURIComponent(query.trim())}` : "";
return apiFetch<TabularSourceCatalogue>(settings, `/api/v1/dataflow/sources${suffix}`);
}
export function createDataflowSourceSnapshot(
settings: ApiSettings,
payload: {
name: string;
source_name: string;
description?: string | null;
} & (
{
format: "json";
rows: Record<string, unknown>[];
} | {
format: "csv";
csv_text: string;
delimiter: string;
}
)
): Promise<TabularSource> {
return apiFetch<TabularSource>(settings, "/api/v1/dataflow/sources/snapshots", {
method: "POST",
body: JSON.stringify(payload)
});
}
export async function listDataflowPipelines(settings: ApiSettings): Promise<Pipeline[]> { export async function listDataflowPipelines(settings: ApiSettings): Promise<Pipeline[]> {
const response = await apiFetch<{ pipelines: Pipeline[] }>(settings, "/api/v1/dataflow/pipelines"); const response = await apiFetch<{ pipelines: Pipeline[] }>(settings, "/api/v1/dataflow/pipelines");
return response.pipelines; return response.pipelines;

View File

@@ -15,12 +15,13 @@ import {
} from "@xyflow/react"; } from "@xyflow/react";
import type { import type {
DataflowDiagnostic, DataflowDiagnostic,
NodeTypeDefinition,
NodePreviewDiagnostic, NodePreviewDiagnostic,
PipelineGraph, PipelineGraph,
PipelineGraphNode PipelineGraphNode
} from "../../api/dataflow"; } from "../../api/dataflow";
import DataflowNode, { type DataflowFlowNode } from "./DataflowNode"; import DataflowNode, { type DataflowFlowNode } from "./DataflowNode";
import { newNode } from "./model"; import { FALLBACK_NODE_LIBRARY, newNode } from "./model";
const nodeTypes = { dataflow: DataflowNode }; const nodeTypes = { dataflow: DataflowNode };
@@ -28,6 +29,7 @@ type DataflowCanvasProps = {
graph: PipelineGraph; graph: PipelineGraph;
diagnostics: DataflowDiagnostic[]; diagnostics: DataflowDiagnostic[];
nodeDiagnostics: NodePreviewDiagnostic[]; nodeDiagnostics: NodePreviewDiagnostic[];
nodeLibrary: NodeTypeDefinition[];
selectedNodeId: string | null; selectedNodeId: string | null;
readOnly: boolean; readOnly: boolean;
onGraphChange: (graph: PipelineGraph) => void; onGraphChange: (graph: PipelineGraph) => void;
@@ -38,6 +40,7 @@ export default function DataflowCanvas({
graph, graph,
diagnostics, diagnostics,
nodeDiagnostics, nodeDiagnostics,
nodeLibrary,
selectedNodeId, selectedNodeId,
readOnly, readOnly,
onGraphChange, onGraphChange,
@@ -52,8 +55,16 @@ export default function DataflowCanvas({
() => new Map(nodeDiagnostics.map((item) => [item.node_id, item.output_rows])), () => new Map(nodeDiagnostics.map((item) => [item.node_id, item.output_rows])),
[nodeDiagnostics] [nodeDiagnostics]
); );
const definitions = useMemo(
() => new Map(nodeLibrary.map((definition) => [definition.type, definition])),
[nodeLibrary]
);
const nodes = useMemo<DataflowFlowNode[]>( const nodes = useMemo<DataflowFlowNode[]>(
() => graph.nodes.map((node) => ({ () => graph.nodes.map((node) => {
const definition = definitions.get(node.type)
?? FALLBACK_NODE_LIBRARY.find((item) => item.type === node.type)
?? FALLBACK_NODE_LIBRARY[0];
return {
id: node.id, id: node.id,
type: "dataflow", type: "dataflow",
position: node.position, position: node.position,
@@ -61,18 +72,22 @@ export default function DataflowCanvas({
data: { data: {
label: node.label, label: node.label,
transformType: node.type, transformType: node.type,
definition,
config: node.config, config: node.config,
hasError: errorNodeIds.has(node.id), hasError: errorNodeIds.has(node.id),
outputRows: rowCounts.get(node.id) outputRows: rowCounts.get(node.id)
} }
})), };
[errorNodeIds, graph.nodes, rowCounts, selectedNodeId] }),
[definitions, errorNodeIds, graph.nodes, rowCounts, selectedNodeId]
); );
const edges = useMemo<Edge[]>( const edges = useMemo<Edge[]>(
() => graph.edges.map((edge) => ({ () => graph.edges.map((edge) => ({
id: edge.id, id: edge.id,
source: edge.source, source: edge.source,
target: edge.target, target: edge.target,
sourceHandle: edge.source_port ?? "output",
targetHandle: edge.target_port ?? "input",
type: "smoothstep", type: "smoothstep",
className: "dataflow-edge" className: "dataflow-edge"
})), })),
@@ -113,10 +128,30 @@ export default function DataflowCanvas({
const source = graph.nodes.find((node) => node.id === connection.source); const source = graph.nodes.find((node) => node.id === connection.source);
const target = graph.nodes.find((node) => node.id === connection.target); const target = graph.nodes.find((node) => node.id === connection.target);
if (!source || !target || source.type === "output" || target.type.startsWith("source.")) return false; if (!source || !target || source.type === "output" || target.type.startsWith("source.")) return false;
const targetAlreadyConnected = graph.edges.some( const sourcePort = connection.sourceHandle ?? "output";
(edge) => edge.target === connection.target && edge.source !== connection.source const targetPort = connection.targetHandle ?? "input";
const sourceDefinition = definitions.get(source.type);
const targetDefinition = definitions.get(target.type);
if (
!sourceDefinition?.output_ports.some((port) => port.id === sourcePort)
|| !targetDefinition?.input_ports.some((port) => port.id === targetPort)
) {
return false;
}
const duplicate = graph.edges.some(
(edge) =>
edge.source === connection.source
&& edge.target === connection.target
&& (edge.source_port ?? "output") === sourcePort
&& (edge.target_port ?? "input") === targetPort
); );
return !targetAlreadyConnected; if (duplicate) return false;
const port = targetDefinition.input_ports.find((item) => item.id === targetPort);
const targetPortConnections = graph.edges.filter(
(edge) => edge.target === connection.target && (edge.target_port ?? "input") === targetPort
);
if (!port?.multiple && targetPortConnections.length) return false;
return !wouldCreateCycle(graph, connection.source, connection.target);
}; };
const onConnect = (connection: Connection) => { const onConnect = (connection: Connection) => {
@@ -138,7 +173,7 @@ export default function DataflowCanvas({
const type = event.dataTransfer.getData("application/x-govoplan-dataflow-node"); const type = event.dataTransfer.getData("application/x-govoplan-dataflow-node");
if (!type) return; if (!type) return;
const position = instance.screenToFlowPosition({ x: event.clientX, y: event.clientY }); const position = instance.screenToFlowPosition({ x: event.clientX, y: event.clientY });
const node = newNode(type, position); const node = newNode(type, position, nodeLibrary);
onGraphChange({ ...graph, nodes: [...graph.nodes, node] }); onGraphChange({ ...graph, nodes: [...graph.nodes, node] });
onSelectNode(node.id); onSelectNode(node.id);
}; };
@@ -199,6 +234,23 @@ export default function DataflowCanvas({
); );
} }
function wouldCreateCycle(graph: PipelineGraph, source: string, target: string): boolean {
const outgoing = new Map<string, string[]>();
graph.edges.forEach((edge) => {
outgoing.set(edge.source, [...(outgoing.get(edge.source) ?? []), edge.target]);
});
const pending = [target];
const visited = new Set<string>();
while (pending.length) {
const nodeId = pending.pop();
if (!nodeId || visited.has(nodeId)) continue;
if (nodeId === source) return true;
visited.add(nodeId);
pending.push(...(outgoing.get(nodeId) ?? []));
}
return false;
}
export function updateGraphNode( export function updateGraphNode(
graph: PipelineGraph, graph: PipelineGraph,
updatedNode: PipelineGraphNode updatedNode: PipelineGraphNode

View File

@@ -1,19 +1,11 @@
import type { ComponentType } from "react";
import {
Braces,
Columns3,
Database,
Filter,
ListEnd,
Sigma,
SortAsc,
Unplug
} from "lucide-react";
import { Handle, Position, type Node, type NodeProps } from "@xyflow/react"; import { Handle, Position, type Node, type NodeProps } from "@xyflow/react";
import type { NodeTypeDefinition } from "../../api/dataflow";
import { dataflowNodeIcon } from "./nodeIcons";
export type DataflowNodeData = { export type DataflowNodeData = {
label: string; label: string;
transformType: string; transformType: string;
definition: NodeTypeDefinition;
config: Record<string, unknown>; config: Record<string, unknown>;
hasError?: boolean; hasError?: boolean;
outputRows?: number; outputRows?: number;
@@ -21,21 +13,8 @@ export type DataflowNodeData = {
export type DataflowFlowNode = Node<DataflowNodeData, "dataflow">; export type DataflowFlowNode = Node<DataflowNodeData, "dataflow">;
const iconByType: Record<string, ComponentType<{ size?: number; strokeWidth?: number }>> = {
"source.inline": Braces,
"source.reference": Database,
filter: Filter,
select: Columns3,
aggregate: Sigma,
sort: SortAsc,
limit: ListEnd,
output: Unplug
};
export default function DataflowNode({ data, selected }: NodeProps<DataflowFlowNode>) { export default function DataflowNode({ data, selected }: NodeProps<DataflowFlowNode>) {
const Icon = iconByType[data.transformType] ?? Braces; const Icon = dataflowNodeIcon(data.definition.icon);
const isSource = data.transformType.startsWith("source.");
const isOutput = data.transformType === "output";
return ( return (
<div <div
@@ -46,42 +25,51 @@ export default function DataflowNode({ data, selected }: NodeProps<DataflowFlowN
data.hasError ? "has-error" : "" data.hasError ? "has-error" : ""
].filter(Boolean).join(" ")} ].filter(Boolean).join(" ")}
> >
{!isSource ? ( {data.definition.input_ports.map((port, index) => (
<Handle <Handle
key={port.id}
id={port.id}
type="target" type="target"
position={Position.Left} position={Position.Left}
className="dataflow-node-handle dataflow-node-handle-input" className="dataflow-node-handle dataflow-node-handle-input"
style={{ top: portPosition(index, data.definition.input_ports.length) }}
/> />
) : null} ))}
{data.definition.input_ports.length > 1
? data.definition.input_ports.map((port, index) => (
<span
key={port.id}
className="dataflow-node-port-label"
style={{ top: portPosition(index, data.definition.input_ports.length) }}
>
{port.label}
</span>
))
: null}
<span className="dataflow-node-icon" aria-hidden="true"> <span className="dataflow-node-icon" aria-hidden="true">
<Icon size={17} strokeWidth={1.8} /> <Icon size={17} strokeWidth={1.8} />
</span> </span>
<span className="dataflow-node-copy"> <span className="dataflow-node-copy">
<strong>{data.label}</strong> <strong>{data.label}</strong>
<small>{transformLabel(data.transformType)}</small> <small>{data.definition.label}</small>
</span> </span>
{typeof data.outputRows === "number" ? ( {typeof data.outputRows === "number" ? (
<span className="dataflow-node-count">{data.outputRows}</span> <span className="dataflow-node-count">{data.outputRows}</span>
) : null} ) : null}
{!isOutput ? ( {data.definition.output_ports.map((port, index) => (
<Handle <Handle
key={port.id}
id={port.id}
type="source" type="source"
position={Position.Right} position={Position.Right}
className="dataflow-node-handle dataflow-node-handle-output" className="dataflow-node-handle dataflow-node-handle-output"
style={{ top: portPosition(index, data.definition.output_ports.length) }}
/> />
) : null} ))}
</div> </div>
); );
} }
function transformLabel(type: string): string { function portPosition(index: number, count: number): string {
if (type === "source.inline") return "Inline source"; return `${((index + 1) / (count + 1)) * 100}%`;
if (type === "source.reference") return "Connector source";
if (type === "filter") return "Filter";
if (type === "select") return "Projection";
if (type === "aggregate") return "Aggregation";
if (type === "sort") return "Sort";
if (type === "limit") return "Limit";
if (type === "output") return "Output";
return type;
} }

View File

@@ -3,32 +3,27 @@ import {
useEffect, useEffect,
useMemo, useMemo,
useState, useState,
type ComponentType,
type DragEvent type DragEvent
} from "react"; } from "react";
import { import {
Braces,
CheckCircle2, CheckCircle2,
Code2, Code2,
Columns3,
Database,
Filter,
ListEnd,
Network, Network,
Play, Play,
Plus, Plus,
RefreshCw, RefreshCw,
RotateCcw, RotateCcw,
Save, Save,
Sigma,
SortAsc,
Trash2, Trash2,
TriangleAlert TriangleAlert,
Upload
} from "lucide-react"; } from "lucide-react";
import { import {
Button, Button,
ConfirmDialog, ConfirmDialog,
Dialog,
DismissibleAlert, DismissibleAlert,
FormField,
IconButton, IconButton,
LoadingFrame, LoadingFrame,
SegmentedControl, SegmentedControl,
@@ -44,24 +39,29 @@ import { ReactFlowProvider } from "@xyflow/react";
import { import {
compileDataflowSql, compileDataflowSql,
createDataflowPipeline, createDataflowPipeline,
createDataflowSourceSnapshot,
deleteDataflowPipeline, deleteDataflowPipeline,
listDataflowNodeTypes,
listDataflowPipelines, listDataflowPipelines,
listDataflowSources,
previewDataflowPipeline, previewDataflowPipeline,
renderDataflowSql, renderDataflowSql,
updateDataflowPipeline, updateDataflowPipeline,
validateDataflowPipeline, validateDataflowPipeline,
type DataflowDiagnostic, type DataflowDiagnostic,
type EditorMode, type EditorMode,
type NodeCategory,
type NodeTypeDefinition,
type NodePreviewDiagnostic, type NodePreviewDiagnostic,
type Pipeline, type Pipeline,
type PipelineGraphNode, type PipelineGraphNode,
type PipelinePreview type PipelinePreview,
type TabularSource
} from "../../api/dataflow"; } from "../../api/dataflow";
import DataflowCanvas, { updateGraphNode } from "./DataflowCanvas"; import DataflowCanvas, { updateGraphNode } from "./DataflowCanvas";
import NodeInspector from "./NodeInspector"; import NodeInspector from "./NodeInspector";
import { import {
NODE_LABELS, FALLBACK_NODE_LIBRARY,
PALETTE_NODE_TYPES,
draftFingerprint, draftFingerprint,
draftFromPipeline, draftFromPipeline,
newNode, newNode,
@@ -70,18 +70,12 @@ import {
sourceNodes, sourceNodes,
type PipelineDraft type PipelineDraft
} from "./model"; } from "./model";
import { dataflowNodeIcon } from "./nodeIcons";
type ResultTab = "preview" | "diagnostics"; type ResultTab = "preview" | "diagnostics";
type SnapshotFormat = "json" | "csv";
const paletteIcons: Record<string, ComponentType<{ size?: number; strokeWidth?: number }>> = { const NODE_CATEGORY_ORDER: NodeCategory[] = ["load", "combine", "filter", "transform", "output"];
"source.inline": Braces,
filter: Filter,
select: Columns3,
aggregate: Sigma,
sort: SortAsc,
limit: ListEnd,
output: Network
};
export default function DataflowPage({ settings, auth }: { settings: ApiSettings; auth: AuthInfo }) { export default function DataflowPage({ settings, auth }: { settings: ApiSettings; auth: AuthInfo }) {
const { requestNavigation, requestDiscard } = useUnsavedChanges(); const { requestNavigation, requestDiscard } = useUnsavedChanges();
@@ -101,9 +95,20 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
const [resultOpen, setResultOpen] = useState(false); const [resultOpen, setResultOpen] = useState(false);
const [resultTab, setResultTab] = useState<ResultTab>("preview"); const [resultTab, setResultTab] = useState<ResultTab>("preview");
const [deleteOpen, setDeleteOpen] = useState(false); const [deleteOpen, setDeleteOpen] = useState(false);
const [snapshotOpen, setSnapshotOpen] = useState(false);
const [nodeLibrary, setNodeLibrary] = useState<NodeTypeDefinition[]>(FALLBACK_NODE_LIBRARY);
const [sources, setSources] = useState<TabularSource[]>([]);
const [sourceCatalogueAvailable, setSourceCatalogueAvailable] = useState(false);
const [sourceCatalogueWritable, setSourceCatalogueWritable] = useState(false);
const canWrite = hasScope(auth, "dataflow:pipeline:write") || hasScope(auth, "dataflow:pipeline:admin"); const canWrite = hasScope(auth, "dataflow:pipeline:write") || hasScope(auth, "dataflow:pipeline:admin");
const canRun = hasScope(auth, "dataflow:pipeline:run") || hasScope(auth, "dataflow:pipeline:admin"); const canRun = hasScope(auth, "dataflow:pipeline:run") || hasScope(auth, "dataflow:pipeline:admin");
const canImportSources = canWrite
&& sourceCatalogueWritable
&& (
hasScope(auth, "connectors:source:write")
|| hasScope(auth, "connectors:source:admin")
);
const dirty = Boolean(draft) && draftFingerprint(draft) !== draftFingerprint(savedDraft); const dirty = Boolean(draft) && draftFingerprint(draft) !== draftFingerprint(savedDraft);
const selectedNode = useMemo( const selectedNode = useMemo(
() => draft?.graph.nodes.find((node) => node.id === selectedNodeId) ?? null, () => draft?.graph.nodes.find((node) => node.id === selectedNodeId) ?? null,
@@ -116,6 +121,14 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
`${pipeline.name} ${pipeline.description ?? ""} ${pipeline.status}`.toLocaleLowerCase().includes(query) `${pipeline.name} ${pipeline.description ?? ""} ${pipeline.status}`.toLocaleLowerCase().includes(query)
); );
}, [pipelines, search]); }, [pipelines, search]);
const paletteGroups = useMemo(
() => NODE_CATEGORY_ORDER.map((category) => ({
category,
label: nodeLibrary.find((item) => item.category === category)?.category_label ?? category,
nodes: nodeLibrary.filter((item) => item.category === category)
})).filter((group) => group.nodes.length),
[nodeLibrary]
);
const loadPipelines = useCallback(async (preferredId?: string | null) => { const loadPipelines = useCallback(async (preferredId?: string | null) => {
setLoading(true); setLoading(true);
@@ -145,6 +158,33 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
void loadPipelines(); void loadPipelines();
}, []); }, []);
useEffect(() => {
let cancelled = false;
void listDataflowNodeTypes(settings)
.then((items) => {
if (!cancelled && items.length) setNodeLibrary(items);
})
.catch(() => {
if (!cancelled) setNodeLibrary(FALLBACK_NODE_LIBRARY);
});
void listDataflowSources(settings)
.then((catalogue) => {
if (cancelled) return;
setSources(catalogue.sources);
setSourceCatalogueAvailable(catalogue.available);
setSourceCatalogueWritable(catalogue.writable);
})
.catch(() => {
if (cancelled) return;
setSources([]);
setSourceCatalogueAvailable(false);
setSourceCatalogueWritable(false);
});
return () => {
cancelled = true;
};
}, [settings]);
const discardDraft = useCallback(() => { const discardDraft = useCallback(() => {
if (savedDraft) { if (savedDraft) {
const reset = structuredClone(savedDraft); const reset = structuredClone(savedDraft);
@@ -383,7 +423,7 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
const node = newNode(type, { const node = newNode(type, {
x: 120 + draft.graph.nodes.length * 70, x: 120 + draft.graph.nodes.length * 70,
y: 100 + (draft.graph.nodes.length % 4) * 80 y: 100 + (draft.graph.nodes.length % 4) * 80
}); }, nodeLibrary);
updateGraph({ ...draft.graph, nodes: [...draft.graph.nodes, node] }); updateGraph({ ...draft.graph, nodes: [...draft.graph.nodes, node] });
setSelectedNodeId(node.id); setSelectedNodeId(node.id);
}; };
@@ -528,27 +568,41 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
{draft.editorMode === "graph" ? ( {draft.editorMode === "graph" ? (
<aside className="dataflow-palette" aria-label="Transform palette"> <aside className="dataflow-palette" aria-label="Transform palette">
<div className="dataflow-panel-heading"> <div className="dataflow-panel-heading">
<strong>Transforms</strong> <strong>Nodes</strong>
{canImportSources ? (
<IconButton
label="Import tabular snapshot"
icon={<Upload size={15} />}
variant="ghost"
onClick={() => setSnapshotOpen(true)}
/>
) : null}
</div> </div>
<div className="dataflow-palette-items"> <div className="dataflow-palette-items">
{PALETTE_NODE_TYPES.map((type) => { {paletteGroups.map((group) => (
const Icon = paletteIcons[type] ?? Database; <section key={group.category} className="dataflow-palette-group">
const disabled = !canWrite || uniqueNodeExists(draft, type); <h3>{group.label}</h3>
{group.nodes.map((definition) => {
const Icon = dataflowNodeIcon(definition.icon);
const disabled = !canWrite || uniqueNodeExists(draft, definition.type);
return ( return (
<button <button
key={type} key={definition.type}
type="button" type="button"
title={definition.description}
draggable={!disabled} draggable={!disabled}
disabled={disabled} disabled={disabled}
onDragStart={(event) => startPaletteDrag(event, type)} onDragStart={(event) => startPaletteDrag(event, definition.type)}
onClick={() => addNode(type)} onClick={() => addNode(definition.type)}
> >
<Icon size={16} /> <Icon size={16} />
<span>{NODE_LABELS[type]}</span> <span>{definition.label}</span>
<Plus size={14} className="dataflow-palette-add" /> <Plus size={14} className="dataflow-palette-add" />
</button> </button>
); );
})} })}
</section>
))}
</div> </div>
</aside> </aside>
) : null} ) : null}
@@ -559,6 +613,7 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
graph={draft.graph} graph={draft.graph}
diagnostics={diagnostics} diagnostics={diagnostics}
nodeDiagnostics={nodeDiagnostics} nodeDiagnostics={nodeDiagnostics}
nodeLibrary={nodeLibrary}
selectedNodeId={selectedNodeId} selectedNodeId={selectedNodeId}
readOnly={!canWrite} readOnly={!canWrite}
onGraphChange={updateGraph} onGraphChange={updateGraph}
@@ -589,6 +644,9 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
</section> </section>
<NodeInspector <NodeInspector
node={selectedNode} node={selectedNode}
nodeLibrary={nodeLibrary}
sources={sources}
sourceCatalogueAvailable={sourceCatalogueAvailable}
readOnly={!canWrite} readOnly={!canWrite}
onChange={(node: PipelineGraphNode) => updateGraph(updateGraphNode(draft.graph, node))} onChange={(node: PipelineGraphNode) => updateGraph(updateGraphNode(draft.graph, node))}
onDelete={(nodeId) => { onDelete={(nodeId) => {
@@ -634,10 +692,218 @@ export default function DataflowPage({ settings, auth }: { settings: ApiSettings
onCancel={() => setDeleteOpen(false)} onCancel={() => setDeleteOpen(false)}
onConfirm={() => void removePipeline()} onConfirm={() => void removePipeline()}
/> />
<SourceSnapshotDialog
open={snapshotOpen}
settings={settings}
onClose={() => setSnapshotOpen(false)}
onCreated={(source) => {
setSources((current) => [
source,
...current.filter((item) => item.ref !== source.ref)
]);
setSourceCatalogueAvailable(true);
setSnapshotOpen(false);
if (!draft) return;
const node = newNode(
"source.reference",
{
x: 100 + draft.graph.nodes.length * 60,
y: 90 + (draft.graph.nodes.length % 4) * 90
},
nodeLibrary
);
node.label = source.name;
node.config = {
...node.config,
source_ref: source.ref,
source_name: source.source_name,
expected_fingerprint: source.fingerprint,
source_columns: source.columns
};
updateGraph({ ...draft.graph, nodes: [...draft.graph.nodes, node] });
setSelectedNodeId(node.id);
}}
/>
</main> </main>
); );
} }
function SourceSnapshotDialog({
open,
settings,
onClose,
onCreated
}: {
open: boolean;
settings: ApiSettings;
onClose: () => void;
onCreated: (source: TabularSource) => void;
}) {
const [name, setName] = useState("");
const [sourceName, setSourceName] = useState("");
const [description, setDescription] = useState("");
const [format, setFormat] = useState<SnapshotFormat>("json");
const [rowsText, setRowsText] = useState("[]");
const [csvText, setCsvText] = useState("");
const [delimiter, setDelimiter] = useState(",");
const [fileInputKey, setFileInputKey] = useState(0);
const [busy, setBusy] = useState(false);
const [error, setError] = useState("");
useEffect(() => {
if (!open) return;
setName("");
setSourceName("");
setDescription("");
setFormat("json");
setRowsText("[]");
setCsvText("");
setDelimiter(",");
setFileInputKey((current) => current + 1);
setError("");
}, [open]);
const create = async () => {
setError("");
let rows: Record<string, unknown>[] = [];
if (format === "json") {
try {
const parsed: unknown = JSON.parse(rowsText);
if (!Array.isArray(parsed) || parsed.some((row) => !isRecord(row))) {
throw new Error("Rows must be a JSON array of objects.");
}
rows = parsed;
} catch (parseError) {
setError(parseError instanceof Error ? parseError.message : "Rows could not be parsed.");
return;
}
} else if (!csvText.trim()) {
setError("Choose a CSV file or paste CSV data.");
return;
}
setBusy(true);
try {
const source = await createDataflowSourceSnapshot(settings, {
name: name.trim(),
source_name: sourceName.trim(),
description: description.trim() || null,
...(format === "json"
? { format, rows }
: { format, csv_text: csvText, delimiter })
});
onCreated(source);
} catch (createError) {
setError(apiErrorMessage(createError));
} finally {
setBusy(false);
}
};
const loadCsvFile = async (file: File | undefined) => {
if (!file) return;
setError("");
try {
setCsvText(await file.text());
const baseName = file.name.replace(/\.[^.]+$/, "");
setName((current) => current || baseName);
setSourceName((current) => current || sourceNameFromFile(baseName));
} catch {
setError("The selected CSV file could not be read.");
}
};
return (
<Dialog
open={open}
title="Import tabular snapshot"
className="dataflow-source-dialog"
onClose={() => {
if (!busy) onClose();
}}
footer={(
<>
<Button onClick={onClose} disabled={busy}>Cancel</Button>
<Button
variant="primary"
onClick={() => void create()}
disabled={busy || !name.trim() || !sourceName.trim()}
>
{busy ? "Importing..." : "Import"}
</Button>
</>
)}
>
<div className="dataflow-source-dialog-fields">
{error ? (
<DismissibleAlert tone="danger" resetKey={error}>{error}</DismissibleAlert>
) : null}
<div className="dataflow-source-dialog-format">
<SegmentedControl<SnapshotFormat>
ariaLabel="Snapshot format"
options={[
{ id: "json", label: "JSON" },
{ id: "csv", label: "CSV" }
]}
value={format}
onChange={setFormat}
/>
</div>
<FormField label="Name">
<input value={name} onChange={(event) => setName(event.target.value)} />
</FormField>
<FormField label="Logical source name">
<input
value={sourceName}
onChange={(event) => setSourceName(event.target.value)}
pattern="[A-Za-z_][A-Za-z0-9_]*"
placeholder="monthly_input"
/>
</FormField>
<FormField label="Description">
<input value={description} onChange={(event) => setDescription(event.target.value)} />
</FormField>
{format === "csv" ? (
<>
<FormField label="CSV file">
<input
key={fileInputKey}
type="file"
accept=".csv,.tsv,text/csv,text/tab-separated-values,text/plain"
onChange={(event) => void loadCsvFile(event.target.files?.[0])}
/>
</FormField>
<FormField label="Delimiter">
<select value={delimiter} onChange={(event) => setDelimiter(event.target.value)}>
<option value=",">Comma</option>
<option value=";">Semicolon</option>
<option value={"\t"}>Tab</option>
<option value="|">Pipe</option>
</select>
</FormField>
<FormField label="CSV data">
<textarea
className="dataflow-json-editor"
value={csvText}
onChange={(event) => setCsvText(event.target.value)}
spellCheck={false}
/>
</FormField>
</>
) : (
<FormField label="Rows">
<textarea
className="dataflow-json-editor"
value={rowsText}
onChange={(event) => setRowsText(event.target.value)}
spellCheck={false}
/>
</FormField>
)}
</div>
</Dialog>
);
}
function ResultPanel({ function ResultPanel({
tab, tab,
onTabChange, onTabChange,
@@ -702,9 +968,17 @@ function PreviewTable({ preview }: { preview: PipelinePreview | null }) {
))} ))}
</tbody> </tbody>
</table> </table>
{preview.truncated ? ( <div className="dataflow-preview-summary">
<div className="dataflow-preview-truncated">Showing {preview.rows.length} of {preview.total_rows} rows</div> <span>
) : null} Showing {preview.rows.length} of {preview.total_rows} output rows
</span>
<span>{preview.input_row_count} input rows</span>
<span>
{preview.source_fingerprints.length} source
{preview.source_fingerprints.length === 1 ? "" : "s"}
</span>
<span>{preview.executor_version}</span>
</div>
</div> </div>
); );
} }
@@ -747,7 +1021,6 @@ function DiagnosticsPanel({
function uniqueNodeExists(draft: PipelineDraft, type: string): boolean { function uniqueNodeExists(draft: PipelineDraft, type: string): boolean {
if (type === "output") return draft.graph.nodes.some((node) => node.type === "output"); if (type === "output") return draft.graph.nodes.some((node) => node.type === "output");
if (type.startsWith("source.")) return draft.graph.nodes.some((node) => node.type.startsWith("source."));
return false; return false;
} }
@@ -762,6 +1035,22 @@ function formatCell(value: unknown): string {
return String(value); return String(value);
} }
function sourceNameFromFile(value: string): string {
const normalized = value
.normalize("NFKD")
.replace(/[\u0300-\u036f]/g, "")
.replace(/[^A-Za-z0-9_]+/g, "_")
.replace(/^_+|_+$/g, "")
.toLocaleLowerCase()
.slice(0, 120);
if (!normalized) return "imported_source";
return (/^[A-Za-z_]/.test(normalized) ? normalized : `source_${normalized}`).slice(0, 120);
}
function isRecord(value: unknown): value is Record<string, unknown> {
return typeof value === "object" && value !== null && !Array.isArray(value);
}
function apiErrorMessage(error: unknown): string { function apiErrorMessage(error: unknown): string {
if (!isApiError(error)) return error instanceof Error ? error.message : "The request failed."; if (!isApiError(error)) return error instanceof Error ? error.message : "The request failed.";
try { try {

View File

@@ -1,16 +1,31 @@
import { useEffect, useState } from "react"; import { useEffect, useState } from "react";
import { Trash2 } from "lucide-react"; import { Trash2 } from "lucide-react";
import { Button, DismissibleAlert, FormField } from "@govoplan/core-webui"; import { Button, DismissibleAlert, FormField } from "@govoplan/core-webui";
import type { PipelineGraphNode } from "../../api/dataflow"; import type {
NodeTypeDefinition,
PipelineGraphNode,
TabularSource
} from "../../api/dataflow";
type NodeInspectorProps = { type NodeInspectorProps = {
node: PipelineGraphNode | null; node: PipelineGraphNode | null;
nodeLibrary: NodeTypeDefinition[];
sources: TabularSource[];
sourceCatalogueAvailable: boolean;
readOnly: boolean; readOnly: boolean;
onChange: (node: PipelineGraphNode) => void; onChange: (node: PipelineGraphNode) => void;
onDelete: (nodeId: string) => void; onDelete: (nodeId: string) => void;
}; };
export default function NodeInspector({ node, readOnly, onChange, onDelete }: NodeInspectorProps) { export default function NodeInspector({
node,
nodeLibrary,
sources,
sourceCatalogueAvailable,
readOnly,
onChange,
onDelete
}: NodeInspectorProps) {
const [rowsText, setRowsText] = useState(""); const [rowsText, setRowsText] = useState("");
const [aggregateText, setAggregateText] = useState(""); const [aggregateText, setAggregateText] = useState("");
const [sortText, setSortText] = useState(""); const [sortText, setSortText] = useState("");
@@ -34,6 +49,7 @@ export default function NodeInspector({ node, readOnly, onChange, onDelete }: No
); );
} }
const definition = nodeLibrary.find((item) => item.type === node.type);
const updateConfig = (patch: Record<string, unknown>) => { const updateConfig = (patch: Record<string, unknown>) => {
onChange({ ...node, config: { ...node.config, ...patch } }); onChange({ ...node, config: { ...node.config, ...patch } });
}; };
@@ -76,7 +92,7 @@ export default function NodeInspector({ node, readOnly, onChange, onDelete }: No
<div className="dataflow-panel-heading"> <div className="dataflow-panel-heading">
<span> <span>
<strong>Inspector</strong> <strong>Inspector</strong>
<small>{node.type}</small> <small>{definition?.label ?? node.type}</small>
</span> </span>
<Button <Button
variant="ghost" variant="ghost"
@@ -111,6 +127,50 @@ export default function NodeInspector({ node, readOnly, onChange, onDelete }: No
/> />
</FormField> </FormField>
) : null} ) : null}
{node.type === "source.reference" ? (
<>
<FormField label="Source">
<select
value={textValue(node.config.source_ref)}
onChange={(event) => {
const source = sources.find((item) => item.ref === event.target.value);
if (!source) {
updateConfig({
source_ref: "",
expected_fingerprint: ""
});
return;
}
updateConfig({
source_ref: source.ref,
source_name: source.source_name,
expected_fingerprint: source.fingerprint,
source_columns: source.columns
});
}}
disabled={readOnly || !sourceCatalogueAvailable}
>
<option value="">
{sourceCatalogueAvailable ? "Choose a source" : "Connectors unavailable"}
</option>
{sources.map((source) => (
<option key={source.ref} value={source.ref}>
{source.name} ({source.row_count ?? "?"} rows)
</option>
))}
</select>
</FormField>
{textValue(node.config.expected_fingerprint) ? (
<FormField label="Pinned fingerprint">
<input
value={textValue(node.config.expected_fingerprint)}
readOnly
title={textValue(node.config.expected_fingerprint)}
/>
</FormField>
) : null}
</>
) : null}
{node.type === "source.inline" ? ( {node.type === "source.inline" ? (
<FormField label="Rows"> <FormField label="Rows">
<textarea <textarea
@@ -160,6 +220,65 @@ export default function NodeInspector({ node, readOnly, onChange, onDelete }: No
) : null} ) : null}
</> </>
) : null} ) : null}
{node.type === "distinct" ? (
<FormField label="Key columns">
<input
value={stringList(node.config.columns).join(", ")}
onChange={(event) => updateConfig({ columns: commaList(event.target.value) })}
placeholder="All columns"
disabled={readOnly}
/>
</FormField>
) : null}
{node.type === "combine.union" ? (
<FormField label="Duplicates">
<select
value={textValue(node.config.mode) || "all"}
onChange={(event) => updateConfig({ mode: event.target.value })}
disabled={readOnly}
>
<option value="all">Keep all rows</option>
<option value="distinct">Remove duplicates</option>
</select>
</FormField>
) : null}
{node.type === "combine.join" ? (
<>
<FormField label="Join type">
<select
value={textValue(node.config.join_type) || "inner"}
onChange={(event) => updateConfig({ join_type: event.target.value })}
disabled={readOnly}
>
<option value="inner">Matching rows</option>
<option value="left">All left rows</option>
<option value="right">All right rows</option>
<option value="full">All rows</option>
</select>
</FormField>
<FormField label="Left keys">
<input
value={stringList(node.config.left_keys).join(", ")}
onChange={(event) => updateConfig({ left_keys: commaList(event.target.value) })}
disabled={readOnly}
/>
</FormField>
<FormField label="Right keys">
<input
value={stringList(node.config.right_keys).join(", ")}
onChange={(event) => updateConfig({ right_keys: commaList(event.target.value) })}
disabled={readOnly}
/>
</FormField>
<FormField label="Right-column prefix">
<input
value={textValue(node.config.right_prefix)}
onChange={(event) => updateConfig({ right_prefix: event.target.value })}
disabled={readOnly}
/>
</FormField>
</>
) : null}
{node.type === "select" ? ( {node.type === "select" ? (
<FormField label="Columns"> <FormField label="Columns">
<input <input
@@ -190,6 +309,51 @@ export default function NodeInspector({ node, readOnly, onChange, onDelete }: No
</FormField> </FormField>
</> </>
) : null} ) : null}
{node.type === "derive" ? (
<>
<FormField label="Output column">
<input
value={textValue(node.config.target_column)}
onChange={(event) => updateConfig({ target_column: event.target.value })}
disabled={readOnly}
/>
</FormField>
<FormField label="Operation">
<select
value={textValue(node.config.operation) || "copy"}
onChange={(event) => updateConfig({ operation: event.target.value })}
disabled={readOnly}
>
<option value="copy">Copy</option>
<option value="upper">Uppercase</option>
<option value="lower">Lowercase</option>
<option value="trim">Trim whitespace</option>
<option value="concat">Concatenate</option>
<option value="coalesce">First non-empty</option>
<option value="add">Add</option>
<option value="subtract">Subtract</option>
<option value="multiply">Multiply</option>
<option value="divide">Divide</option>
</select>
</FormField>
<FormField label="Source columns">
<input
value={stringList(node.config.source_columns).join(", ")}
onChange={(event) => updateConfig({ source_columns: commaList(event.target.value) })}
disabled={readOnly}
/>
</FormField>
{textValue(node.config.operation) === "concat" ? (
<FormField label="Separator">
<input
value={textValue(node.config.separator)}
onChange={(event) => updateConfig({ separator: event.target.value })}
disabled={readOnly}
/>
</FormField>
) : null}
</>
) : null}
{node.type === "sort" ? ( {node.type === "sort" ? (
<FormField label="Sort fields"> <FormField label="Sort fields">
<textarea <textarea

View File

@@ -4,7 +4,8 @@ import type {
PipelineGraph, PipelineGraph,
PipelineGraphNode, PipelineGraphNode,
PipelinePayload, PipelinePayload,
PipelineStatus PipelineStatus,
NodeTypeDefinition
} from "../../api/dataflow"; } from "../../api/dataflow";
export type PipelineDraft = { export type PipelineDraft = {
@@ -18,26 +19,125 @@ export type PipelineDraft = {
editorMode: EditorMode; editorMode: EditorMode;
}; };
export const NODE_LABELS: Record<string, string> = { const input = [{ id: "input", label: "Input", required: true, multiple: false, minimum_connections: 1 }];
"source.inline": "Inline source", const output = [{ id: "output", label: "Output", required: true, multiple: false, minimum_connections: 1 }];
"source.reference": "Connector source",
filter: "Filter rows",
select: "Select columns",
aggregate: "Aggregate",
sort: "Sort rows",
limit: "Limit rows",
output: "Output"
};
export const PALETTE_NODE_TYPES = [ export const FALLBACK_NODE_LIBRARY: NodeTypeDefinition[] = [
"source.inline", nodeType("source.inline", "load", "Load", "Inline data", "Enter a small JSON table.", "braces", [], output, {
source_name: "inline_source",
rows: []
}),
nodeType(
"source.reference",
"load",
"Load",
"Connector source",
"Load a governed tabular source.",
"database",
[],
output,
{ source_ref: "", source_name: "connector_source", expected_fingerprint: "" }
),
nodeType(
"combine.union",
"combine",
"Combine",
"Append rows",
"Append two or more inputs.",
"combine",
[{ id: "input", label: "Inputs", required: true, multiple: true, minimum_connections: 2 }],
output,
{ mode: "all" }
),
nodeType(
"combine.join",
"combine",
"Combine",
"Join tables",
"Match two inputs using key columns.",
"git-merge",
[
{ id: "left", label: "Left", required: true, multiple: false, minimum_connections: 1 },
{ id: "right", label: "Right", required: true, multiple: false, minimum_connections: 1 }
],
output,
{ join_type: "inner", left_keys: [""], right_keys: [""], right_prefix: "right_" }
),
nodeType("filter", "filter", "Filter", "Filter rows", "Keep matching rows.", "filter", input, output, {
column: "",
operator: "eq",
value: ""
}),
nodeType(
"distinct",
"filter", "filter",
"Filter",
"Remove duplicates",
"Keep one row per key.",
"list-filter",
input,
output,
{ columns: [] }
),
nodeType(
"select", "select",
"transform",
"Transform",
"Select columns",
"Choose and rename fields.",
"columns-3",
input,
output,
{ fields: [{ column: "", alias: "" }] }
),
nodeType(
"derive",
"transform",
"Transform",
"Derive column",
"Create a constrained calculated field.",
"variable",
input,
output,
{ target_column: "", operation: "copy", source_columns: [""], separator: " " }
),
nodeType(
"aggregate", "aggregate",
"transform",
"Transform",
"Aggregate",
"Group and summarize rows.",
"sigma",
input,
output,
{ group_by: [], aggregates: [{ function: "count", column: "*", alias: "row_count" }] }
),
nodeType(
"sort", "sort",
"limit", "transform",
"output" "Transform",
] as const; "Sort rows",
"Order rows by fields.",
"arrow-up-down",
input,
output,
{ fields: [{ column: "", direction: "asc" }] }
),
nodeType("limit", "transform", "Transform", "Limit rows", "Keep the first rows.", "list-end", input, output, {
count: 100
}),
nodeType(
"output",
"output",
"Output",
"Preview output",
"Expose the terminal table.",
"panel-top-open",
input,
[],
{}
)
];
export function draftFromPipeline(pipeline: Pipeline): PipelineDraft { export function draftFromPipeline(pipeline: Pipeline): PipelineDraft {
return { return {
@@ -152,26 +252,45 @@ export function draftFingerprint(draft: PipelineDraft | null): string {
}); });
} }
export function newNode(type: string, position: { x: number; y: number }): PipelineGraphNode { export function newNode(
type: string,
position: { x: number; y: number },
library: NodeTypeDefinition[] = FALLBACK_NODE_LIBRARY
): PipelineGraphNode {
const id = `${type.replace(".", "-")}-${crypto.randomUUID()}`; const id = `${type.replace(".", "-")}-${crypto.randomUUID()}`;
const config: Record<string, unknown> = defaultNodeConfig(type); const definition = library.find((item) => item.type === type);
const config = structuredClone(definition?.default_config ?? {});
return { return {
id, id,
type, type,
label: NODE_LABELS[type] ?? type, label: definition?.label ?? type,
position, position,
config config
}; };
} }
function defaultNodeConfig(type: string): Record<string, unknown> { function nodeType(
if (type === "source.inline") return { source_name: "new_source", rows: [] }; type: string,
if (type === "filter") return { column: "", operator: "eq", value: "" }; category: NodeTypeDefinition["category"],
if (type === "select") return { fields: [{ column: "", alias: "" }] }; categoryLabel: string,
if (type === "aggregate") { label: string,
return { group_by: [], aggregates: [{ function: "count", column: "*", alias: "row_count" }] }; description: string,
} icon: string,
if (type === "sort") return { fields: [{ column: "", direction: "asc" }] }; inputPorts: NodeTypeDefinition["input_ports"],
if (type === "limit") return { count: 100 }; outputPorts: NodeTypeDefinition["output_ports"],
return {}; defaultConfig: Record<string, unknown>
): NodeTypeDefinition {
return {
type,
category,
category_label: categoryLabel,
label,
description,
icon,
input_ports: inputPorts,
output_ports: outputPorts,
config_fields: [],
default_config: defaultConfig,
sql_support: ["combine.union", "combine.join", "distinct", "derive"].includes(type) ? "partial" : "full"
};
} }

View File

@@ -0,0 +1,34 @@
import {
ArrowUpDown,
Braces,
Columns3,
Combine,
Database,
Filter,
GitMerge,
ListEnd,
ListFilter,
PanelTopOpen,
Sigma,
Variable,
type LucideIcon
} from "lucide-react";
const icons: Record<string, LucideIcon> = {
"arrow-up-down": ArrowUpDown,
braces: Braces,
"columns-3": Columns3,
combine: Combine,
database: Database,
filter: Filter,
"git-merge": GitMerge,
"list-end": ListEnd,
"list-filter": ListFilter,
"panel-top-open": PanelTopOpen,
sigma: Sigma,
variable: Variable
};
export function dataflowNodeIcon(icon: string): LucideIcon {
return icons[icon] ?? Braces;
}

View File

@@ -240,7 +240,7 @@
.dataflow-editor { .dataflow-editor {
flex: 1 1 auto; flex: 1 1 auto;
display: grid; display: grid;
grid-template-columns: 170px minmax(0, 1fr) minmax(260px, 310px); grid-template-columns: 210px minmax(0, 1fr) minmax(260px, 310px);
min-height: 0; min-height: 0;
overflow: hidden; overflow: hidden;
} }
@@ -297,10 +297,32 @@
.dataflow-palette-items { .dataflow-palette-items {
display: grid; display: grid;
gap: 4px; gap: 4px;
height: calc(100% - 44px);
padding: 8px; padding: 8px;
overflow: auto; overflow: auto;
} }
.dataflow-palette-group {
display: grid;
gap: 2px;
}
.dataflow-palette-group + .dataflow-palette-group {
margin-top: 6px;
padding-top: 8px;
border-top: var(--border-line);
}
.dataflow-palette-group h3 {
margin: 0;
color: var(--muted);
font-size: 10px;
font-weight: 700;
letter-spacing: 0;
padding: 3px 8px;
text-transform: uppercase;
}
.dataflow-palette-items button { .dataflow-palette-items button {
display: grid; display: grid;
grid-template-columns: 18px minmax(0, 1fr) 14px; grid-template-columns: 18px minmax(0, 1fr) 14px;
@@ -425,6 +447,17 @@
border-left-color: #b7791f; border-left-color: #b7791f;
} }
.dataflow-node-combine-union,
.dataflow-node-combine-join {
min-height: 64px;
border-left-color: #2f7d6d;
}
.dataflow-node-distinct {
border-left-color: #b7791f;
}
.dataflow-node-derive,
.dataflow-node-aggregate { .dataflow-node-aggregate {
border-left-color: #76569b; border-left-color: #76569b;
} }
@@ -507,6 +540,24 @@
background: var(--accent); background: var(--accent);
} }
.dataflow-node-port-label {
position: absolute;
left: 8px;
max-width: 42px;
overflow: hidden;
color: var(--muted);
font-size: 8px;
line-height: 1;
pointer-events: none;
text-overflow: ellipsis;
transform: translateY(-50%);
white-space: nowrap;
}
.dataflow-node-combine-join .dataflow-node-icon {
margin-left: 31px;
}
.dataflow-inspector-fields { .dataflow-inspector-fields {
display: grid; display: grid;
gap: 12px; gap: 12px;
@@ -546,6 +597,27 @@
color: var(--danger-text); color: var(--danger-text);
} }
.dataflow-source-dialog {
width: min(720px, calc(100vw - 32px));
}
.dataflow-source-dialog-fields {
display: grid;
grid-template-columns: 1fr 1fr;
gap: 12px;
}
.dataflow-source-dialog-fields .alert,
.dataflow-source-dialog-format,
.dataflow-source-dialog-fields .form-field:last-child {
grid-column: 1 / -1;
}
.dataflow-source-dialog-fields .dataflow-json-editor {
width: 100%;
min-height: 260px;
}
.dataflow-sql-workbench { .dataflow-sql-workbench {
display: flex; display: flex;
height: 100%; height: 100%;
@@ -651,9 +723,12 @@
background: var(--primary-soft); background: var(--primary-soft);
} }
.dataflow-preview-truncated { .dataflow-preview-summary {
position: sticky; position: sticky;
bottom: 0; bottom: 0;
display: flex;
flex-wrap: wrap;
gap: 6px 16px;
padding: 6px 10px; padding: 6px 10px;
border-top: var(--border-line); border-top: var(--border-line);
background: var(--panel-header); background: var(--panel-header);
@@ -754,7 +829,7 @@
} }
.dataflow-editor { .dataflow-editor {
grid-template-columns: 150px minmax(0, 1fr) 270px; grid-template-columns: 180px minmax(0, 1fr) 270px;
} }
.dataflow-editor.is-sql { .dataflow-editor.is-sql {
@@ -839,6 +914,23 @@
overflow-x: auto; overflow-x: auto;
} }
.dataflow-palette-group {
display: flex;
flex: 0 0 auto;
}
.dataflow-palette-group + .dataflow-palette-group {
margin-top: 0;
padding-top: 0;
padding-left: 6px;
border-top: 0;
border-left: var(--border-line);
}
.dataflow-palette-group h3 {
display: none;
}
.dataflow-palette-items button { .dataflow-palette-items button {
min-width: 128px; min-width: 128px;
} }