From 8a11d7571bdb218b59048ff76b77b5b5c736c828 Mon Sep 17 00:00:00 2001 From: Albrecht Degering Date: Wed, 29 Jul 2026 16:47:14 +0200 Subject: [PATCH] Add typed Arrow execution backend boundary --- README.md | 24 +- pyproject.toml | 6 + .../backend/backends/__init__.py | 25 + .../backend/backends/base.py | 196 +++++++ .../backend/backends/duckdb.py | 412 +++++++++++++++ .../backend/backends/reference.py | 156 ++++++ .../backend/backends/registry.py | 157 ++++++ src/govoplan_dataflow/backend/batches.py | 222 ++++++++ src/govoplan_dataflow/backend/ir.py | 491 ++++++++++++++++++ src/govoplan_dataflow/backend/planner.py | 99 ++++ .../backend/schema_validation.py | 11 +- src/govoplan_dataflow/backend/schemas.py | 1 + src/govoplan_dataflow/backend/service.py | 220 ++++++-- src/govoplan_dataflow/backend/sql_compiler.py | 2 +- tests/test_service.py | 32 ++ tests/test_typed_execution.py | 258 +++++++++ webui/src/api/dataflow.ts | 1 + 17 files changed, 2264 insertions(+), 49 deletions(-) create mode 100644 src/govoplan_dataflow/backend/backends/__init__.py create mode 100644 src/govoplan_dataflow/backend/backends/base.py create mode 100644 src/govoplan_dataflow/backend/backends/duckdb.py create mode 100644 src/govoplan_dataflow/backend/backends/reference.py create mode 100644 src/govoplan_dataflow/backend/backends/registry.py create mode 100644 src/govoplan_dataflow/backend/batches.py create mode 100644 src/govoplan_dataflow/backend/ir.py create mode 100644 src/govoplan_dataflow/backend/planner.py create mode 100644 tests/test_typed_execution.py diff --git a/README.md b/README.md index a68768a..3a1e1b9 100644 --- a/README.md +++ b/README.md @@ -59,10 +59,26 @@ expression filters and calculated columns. It supports literals, columns, arithmetic, comparisons, boolean logic, `CASE`, safe casts, and a bounded string/numeric function catalogue without Python evaluation or effectful SQL. -Node definitions, validators, preview executors, and SQL renderers are -registered independently in the operator registry. Adding a node no longer -requires another branch in the preview or graph-to-SQL dispatch loop. Operators -that cannot be represented by constrained SQL declare that explicitly. +Node definitions, validators, schema propagators, preview executors, and SQL +renderers are registered independently in the operator registry. Adding a node +no longer requires another branch in the preview or graph-to-SQL dispatch loop. +Operators that cannot be represented by constrained SQL declare that +explicitly. + +Every executable graph also has a versioned typed IR. It preserves graph +identity and layout while giving ports, schemas, expressions, parameters, +lineage, diagnostics, semantic hashes, and physical results stable contracts. +The backend boundary exchanges bounded typed columnar batches that can be +serialized as Arrow IPC. The deterministic Python executor remains the default +and fallback. + +Install `.[analytics]` to enable the analytical backend. It executes only SQL +generated from validated graphs in a separate short-lived DuckDB process. +Inputs and outputs cross the process boundary as Arrow IPC; external access, +extension installation/loading, persistent secrets, configuration changes, +temporary spill files, and multiple DuckDB threads are disabled. Wall-clock, +row, byte, memory, file-descriptor, and backend-concurrency limits are applied. +Callers can explicitly request `reference`, `duckdb`, or `auto` for a preview. 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 diff --git a/pyproject.toml b/pyproject.toml index 8e68e8f..5cc460d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,6 +15,12 @@ dependencies = [ "sqlglot>=30.14,<31", ] +[project.optional-dependencies] +analytics = [ + "duckdb>=1.5.5,<2", + "pyarrow>=24,<26", +] + [tool.setuptools.packages.find] where = ["src"] diff --git a/src/govoplan_dataflow/backend/backends/__init__.py b/src/govoplan_dataflow/backend/backends/__init__.py new file mode 100644 index 0000000..80d12df --- /dev/null +++ b/src/govoplan_dataflow/backend/backends/__init__.py @@ -0,0 +1,25 @@ +from govoplan_dataflow.backend.backends.base import ( + BackendExecutionError, + BackendExecutionRequest, + BackendExecutionResult, + BackendSource, + ExecutionBackend, + ExecutionBudget, +) +from govoplan_dataflow.backend.backends.registry import ( + EXECUTION_BACKENDS, + ExecutionBackendRegistry, + execute_typed_graph, +) + +__all__ = [ + "BackendExecutionError", + "BackendExecutionRequest", + "BackendExecutionResult", + "BackendSource", + "EXECUTION_BACKENDS", + "ExecutionBackend", + "ExecutionBackendRegistry", + "ExecutionBudget", + "execute_typed_graph", +] diff --git a/src/govoplan_dataflow/backend/backends/base.py b/src/govoplan_dataflow/backend/backends/base.py new file mode 100644 index 0000000..fb2ae81 --- /dev/null +++ b/src/govoplan_dataflow/backend/backends/base.py @@ -0,0 +1,196 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Mapping, Protocol, runtime_checkable + +from govoplan_dataflow.backend.batches import TypedBatch +from govoplan_dataflow.backend.ir import ( + IrExecutionResult, + IrField, + IrLineage, + IrSchema, + TypedGraphIr, + schema_from_fields, +) +from govoplan_dataflow.backend.planner import ExecutionPlan +from govoplan_dataflow.backend.schemas import ( + DataflowDiagnostic, + NodePreviewDiagnostic, + NodePreviewResult, +) + + +@dataclass(frozen=True, slots=True) +class ExecutionBudget: + max_output_rows: int = 500 + max_batch_bytes: int = 1_000_000 + max_wall_seconds: float = 2.0 + max_memory_bytes: int = 256 * 1024 * 1024 + max_concurrency: int = 1 + + def __post_init__(self) -> None: + if self.max_output_rows < 1: + raise ValueError("Execution output row limit must be positive.") + if self.max_batch_bytes < 1: + raise ValueError("Execution byte limit must be positive.") + if self.max_wall_seconds <= 0: + raise ValueError("Execution time limit must be positive.") + if self.max_memory_bytes < 64 * 1024 * 1024: + raise ValueError("Execution memory limit must be at least 64 MiB.") + if self.max_concurrency < 1: + raise ValueError("Execution concurrency limit must be positive.") + + +@dataclass(frozen=True, slots=True) +class BackendSource: + node_id: str + batch: TypedBatch + source_ref: str + provider: str + fingerprint: str + total_rows: int + truncated: bool = False + source_name: str | None = None + kind: str = "datasource" + + def lineage_payload(self) -> dict[str, Any]: + payload: dict[str, Any] = { + "node_id": self.node_id, + "source_name": self.source_name, + "kind": self.kind, + "fingerprint": self.fingerprint, + "row_count": self.total_rows, + } + if self.kind != "inline": + payload.update( + { + "source_ref": self.source_ref, + "provider": self.provider, + "preview_rows": self.batch.row_count, + "truncated": self.truncated, + } + ) + return payload + + +@dataclass(frozen=True, slots=True) +class BackendExecutionRequest: + plan: ExecutionPlan + sources: Mapping[str, BackendSource] = field(default_factory=dict) + budget: ExecutionBudget = field(default_factory=ExecutionBudget) + preview_node_id: str | None = None + + @property + def graph(self) -> TypedGraphIr: + return self.plan.graph + + +@dataclass(frozen=True, slots=True) +class BackendExecutionResult: + contract: IrExecutionResult + batch: TypedBatch + node_diagnostics: tuple[NodePreviewDiagnostic, ...] = () + node_preview: NodePreviewResult | None = None + metadata: Mapping[str, Any] = field(default_factory=dict) + + @property + def rows(self) -> list[dict[str, Any]]: + return self.batch.to_rows() + + +class BackendExecutionError(RuntimeError): + def __init__( + self, + message: str, + *, + code: str = "backend.execution", + node_id: str | None = None, + diagnostics: tuple[DataflowDiagnostic, ...] = (), + ) -> None: + super().__init__(message) + self.code = code + self.node_id = node_id + self.diagnostics = diagnostics + + +@runtime_checkable +class ExecutionBackend(Protocol): + name: str + version: str + + def available(self) -> bool: + ... + + def supports(self, request: BackendExecutionRequest) -> bool: + ... + + def execute( + self, + request: BackendExecutionRequest, + ) -> BackendExecutionResult: + ... + + +def request_lineage(request: BackendExecutionRequest) -> IrLineage: + fingerprints = tuple( + source.lineage_payload() + for source in request.sources.values() + ) + return request.graph.lineage.model_copy( + update={"source_fingerprints": fingerprints} + ) + + +def canonical_result_schema( + request: BackendExecutionRequest, + observed: IrSchema, +) -> IrSchema: + output = next( + ( + node.output_schema + for node in request.graph.nodes + if node.type == "output" + ), + None, + ) + if output is None: + return observed + declared = { + field.name: field + for field in output.fields + } + fields = tuple( + _resolved_result_field(field, declared.get(field.name)) + for field in observed.fields + ) + return schema_from_fields(fields, open=output.open) + + +def _resolved_result_field( + observed: IrField, + declared: IrField | None, +) -> IrField: + if declared is None: + return observed + return declared.model_copy( + update={ + "type": ( + observed.type + if declared.type in {"unknown", "null"} + else declared.type + ), + "nullable": declared.nullable, + } + ) + + +__all__ = [ + "BackendExecutionError", + "BackendExecutionRequest", + "BackendExecutionResult", + "BackendSource", + "ExecutionBackend", + "ExecutionBudget", + "canonical_result_schema", + "request_lineage", +] diff --git a/src/govoplan_dataflow/backend/backends/duckdb.py b/src/govoplan_dataflow/backend/backends/duckdb.py new file mode 100644 index 0000000..5967029 --- /dev/null +++ b/src/govoplan_dataflow/backend/backends/duckdb.py @@ -0,0 +1,412 @@ +from __future__ import annotations + +import hashlib +import importlib.util +import json +import multiprocessing +import os +import signal +from typing import Any + +import sqlglot +from sqlglot import exp + +from govoplan_dataflow.backend.backends.base import ( + BackendExecutionError, + BackendExecutionRequest, + BackendExecutionResult, + BackendSource, + canonical_result_schema, +) +from govoplan_dataflow.backend.batches import ( + ArrowDependencyError, + TypedBatch, +) +from govoplan_dataflow.backend.ir import ( + IrExecutionResult, + IrLineage, + ir_to_graph, +) +from govoplan_dataflow.backend.schemas import GraphNode + + +DUCKDB_BACKEND_VERSION = "duckdb-isolated-v1" + + +class DuckDbExecutionBackend: + name = "duckdb" + version = DUCKDB_BACKEND_VERSION + + def available(self) -> bool: + return _analytics_available() + + def supports(self, request: BackendExecutionRequest) -> bool: + return ( + self.available() + and _previews_final_output(request) + and request.plan.generated_sql is not None + ) + + def execute( + self, + request: BackendExecutionRequest, + ) -> BackendExecutionResult: + if not self.available(): + raise BackendExecutionError( + "DuckDB execution requires the govoplan-dataflow analytics extra.", + code="backend.unavailable", + ) + if not _previews_final_output(request): + raise BackendExecutionError( + "DuckDB execution currently produces the final output only.", + code="backend.preview_node", + node_id=request.preview_node_id, + ) + sql_text = request.plan.generated_sql + if sql_text is None: + raise BackendExecutionError( + "This graph cannot be executed by the DuckDB backend.", + code="backend.unsupported_graph", + diagnostics=request.plan.sql_diagnostics, + ) + graph = ir_to_graph(request.graph) + _validate_generated_select(sql_text) + sources = _source_batches(request, graph.nodes) + worker_payload = { + "sql": sql_text, + "sources": { + name: source.batch.to_arrow_ipc() + for name, source in sources.items() + }, + "max_rows": request.budget.max_output_rows, + "max_bytes": request.budget.max_batch_bytes, + "max_memory_bytes": request.budget.max_memory_bytes, + } + worker_result = _run_worker( + worker_payload, + timeout=request.budget.max_wall_seconds, + ) + try: + observed_batch = TypedBatch.from_arrow_ipc(worker_result["ipc"]) + batch = TypedBatch.from_rows( + observed_batch.to_rows(), + schema=canonical_result_schema( + request, + observed_batch.schema, + ), + ) + except (ArrowDependencyError, KeyError, TypeError, ValueError) as exc: + raise BackendExecutionError( + "DuckDB returned an invalid Arrow result.", + code="backend.result_invalid", + ) from exc + batch.ensure_within( + max_rows=request.budget.max_output_rows, + max_bytes=request.budget.max_batch_bytes, + ) + lineage = _source_lineage(request, sources) + contract = IrExecutionResult( + backend=self.name, + backend_version=f"{self.version}+{worker_result['duckdb_version']}", + graph_semantic_hash=request.graph.semantic_hash, + result_schema=batch.schema, + row_count=int(worker_result["total_rows"]), + byte_count=batch.byte_count, + truncated=bool(worker_result["truncated"]), + diagnostics=request.plan.diagnostics, + lineage=lineage, + ) + return BackendExecutionResult( + contract=contract, + batch=batch, + metadata={ + "source_fingerprints": tuple( + source.lineage_payload() + for source in sources.values() + ), + "isolated_process": True, + "external_access": False, + }, + ) + + +def _analytics_available() -> bool: + return ( + importlib.util.find_spec("duckdb") is not None + and importlib.util.find_spec("pyarrow") is not None + ) + + +def _previews_final_output(request: BackendExecutionRequest) -> bool: + if request.preview_node_id is None: + return True + return any( + node.id == request.preview_node_id and node.type == "output" + for node in request.graph.nodes + ) + + +def _validate_generated_select(sql_text: str) -> None: + statements = sqlglot.parse(sql_text, read="duckdb") + if len(statements) != 1 or not isinstance( + statements[0], + (exp.Select, exp.Union), + ): + raise BackendExecutionError( + "DuckDB accepts one generated SELECT statement only.", + code="backend.sql_rejected", + ) + prohibited = ( + exp.Command, + exp.Copy, + exp.Create, + exp.Delete, + exp.Drop, + exp.Insert, + exp.Merge, + exp.Update, + ) + if any(isinstance(item, prohibited) for item in statements[0].walk()): + raise BackendExecutionError( + "DuckDB rejected an effectful statement.", + code="backend.sql_rejected", + ) + + +def _source_batches( + request: BackendExecutionRequest, + nodes: list[GraphNode], +) -> dict[str, BackendSource]: + sources: dict[str, BackendSource] = {} + for node in nodes: + if not node.type.startswith("source."): + continue + source_name = str(node.config.get("source_name") or "").strip() + source = ( + _inline_source(node) + if node.type == "source.inline" + else _request_source(request, node) + ) + try: + source.batch.ensure_within( + max_rows=10_000, + max_bytes=request.budget.max_batch_bytes, + ) + except ValueError as exc: + raise BackendExecutionError( + str(exc), + code="backend.source_budget", + node_id=node.id, + ) from exc + sources[source_name] = source + return sources + + +def _inline_source(node: GraphNode) -> BackendSource: + rows = [dict(row) for row in node.config.get("rows") or []] + batch = TypedBatch.from_rows(rows) + fingerprint = hashlib.sha256( + json.dumps( + rows, + sort_keys=True, + separators=(",", ":"), + default=str, + ).encode("utf-8") + ).hexdigest() + return BackendSource( + node_id=node.id, + batch=batch, + source_ref=f"inline:{node.id}", + provider="inline", + fingerprint=fingerprint, + total_rows=batch.row_count, + source_name=str(node.config.get("source_name") or ""), + kind="inline", + ) + + +def _request_source( + request: BackendExecutionRequest, + node: GraphNode, +) -> BackendSource: + keys = ( + node.id, + str(node.config.get("source_ref") or ""), + str(node.config.get("source_name") or ""), + ) + source = next( + ( + request.sources[key] + for key in keys + if key and key in request.sources + ), + None, + ) + if source is None: + raise BackendExecutionError( + f"No typed source batch was supplied for node {node.id!r}.", + code="backend.source_missing", + node_id=node.id, + ) + return source + + +def _source_lineage( + request: BackendExecutionRequest, + sources: dict[str, BackendSource], +) -> IrLineage: + return request.graph.lineage.model_copy( + update={ + "source_fingerprints": tuple( + source.lineage_payload() + for source in sources.values() + ) + } + ) + + +def _run_worker(payload: dict[str, Any], *, timeout: float) -> dict[str, Any]: + context = multiprocessing.get_context("spawn") + receive, send = context.Pipe(duplex=False) + process = context.Process( + target=_duckdb_worker, + args=(send, payload), + name="govoplan-dataflow-duckdb", + daemon=True, + ) + process.start() + send.close() + try: + if not receive.poll(timeout): + _terminate_worker(process) + raise BackendExecutionError( + "DuckDB execution exceeded its wall-clock budget.", + code="backend.timeout", + ) + result = receive.recv() + except EOFError as exc: + process.join(timeout=0.2) + raise BackendExecutionError( + "DuckDB worker exited without a result.", + code="backend.worker_exit", + ) from exc + finally: + receive.close() + process.join(timeout=0.2) + if process.is_alive(): + _terminate_worker(process) + if not isinstance(result, dict) or not result.get("ok"): + message = ( + str(result.get("error")) + if isinstance(result, dict) + else "DuckDB worker returned an invalid response." + ) + raise BackendExecutionError( + message, + code="backend.duckdb", + ) + return result + + +def _terminate_worker(process: multiprocessing.Process) -> None: + process.terminate() + process.join(timeout=0.5) + if process.is_alive() and process.pid: + os.kill(process.pid, signal.SIGKILL) + process.join(timeout=0.5) + + +def _duckdb_worker(send: Any, payload: dict[str, Any]) -> None: + try: + _set_worker_limits(float(payload["max_memory_bytes"])) + result = _execute_duckdb_payload(payload) + send.send({"ok": True, **result}) + except BaseException as exc: + send.send( + { + "ok": False, + "error": f"{type(exc).__name__}: {exc}", + } + ) + finally: + send.close() + + +def _set_worker_limits(max_memory_bytes: float) -> None: + try: + import resource + except ImportError: + return + resource.setrlimit(resource.RLIMIT_CORE, (0, 0)) + resource.setrlimit(resource.RLIMIT_FSIZE, (0, 0)) + resource.setrlimit(resource.RLIMIT_NOFILE, (64, 64)) + cpu_seconds = 5 + resource.setrlimit(resource.RLIMIT_CPU, (cpu_seconds, cpu_seconds)) + os.environ["GOVOPLAN_DATAFLOW_MEMORY_BUDGET"] = str(int(max_memory_bytes)) + + +def _execute_duckdb_payload(payload: dict[str, Any]) -> dict[str, Any]: + import secrets + + import duckdb + import pyarrow as pa + + memory_bytes = int(payload["max_memory_bytes"]) + connection = duckdb.connect( + database=":memory:", + config={ + "allow_community_extensions": "false", + "allow_unsigned_extensions": "false", + "autoinstall_known_extensions": "false", + "autoload_known_extensions": "false", + "allow_persistent_secrets": "false", + "memory_limit": f"{memory_bytes}B", + "max_temp_directory_size": "0B", + "threads": "1", + }, + ) + try: + for name, ipc_payload in payload["sources"].items(): + with pa.ipc.open_stream(ipc_payload) as reader: + connection.register(name, reader.read_all()) + connection.execute("SET enable_external_access = false") + connection.execute("SET lock_configuration = true") + total_column = ( + f"__govoplan_internal_total_rows_{secrets.token_hex(16)}" + ) + query = ( + f'SELECT *, COUNT(*) OVER () AS "{total_column}" ' + f'FROM ({payload["sql"]}) AS "_govoplan_result" ' + f'LIMIT {int(payload["max_rows"]) + 1}' + ) + table = connection.execute(query).to_arrow_table() + total_rows = ( + int(table.column(total_column)[0].as_py()) + if table.num_rows + else 0 + ) + table = table.drop([total_column]) + truncated = table.num_rows > int(payload["max_rows"]) + if truncated: + table = table.slice(0, int(payload["max_rows"])) + if table.nbytes > int(payload["max_bytes"]): + raise MemoryError( + "DuckDB output exceeds the configured byte budget." + ) + sink = pa.BufferOutputStream() + with pa.ipc.new_stream(sink, table.schema) as writer: + writer.write_table(table) + return { + "ipc": sink.getvalue().to_pybytes(), + "total_rows": total_rows, + "truncated": truncated, + "duckdb_version": duckdb.__version__, + } + finally: + connection.close() + + +__all__ = [ + "DUCKDB_BACKEND_VERSION", + "DuckDbExecutionBackend", +] diff --git a/src/govoplan_dataflow/backend/backends/reference.py b/src/govoplan_dataflow/backend/backends/reference.py new file mode 100644 index 0000000..b59650c --- /dev/null +++ b/src/govoplan_dataflow/backend/backends/reference.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +from govoplan_dataflow.backend.backends.base import ( + BackendExecutionError, + BackendExecutionRequest, + BackendExecutionResult, + BackendSource, + canonical_result_schema, +) +from govoplan_dataflow.backend.batches import TypedBatch +from govoplan_dataflow.backend.executor import ( + EXECUTOR_VERSION, + PipelineExecutionError, + ResolvedSource, + execute_preview, +) +from govoplan_dataflow.backend.ir import IrExecutionResult, ir_to_graph +from govoplan_dataflow.backend.schemas import GraphNode + + +class ReferenceExecutionBackend: + name = "reference" + version = EXECUTOR_VERSION + + def available(self) -> bool: + return True + + def supports(self, request: BackendExecutionRequest) -> bool: + return request.graph.ir_version == 1 + + def execute( + self, + request: BackendExecutionRequest, + ) -> BackendExecutionResult: + _validate_source_batches(request) + try: + result = execute_preview( + ir_to_graph(request.graph), + row_limit=request.budget.max_output_rows, + source_resolver=lambda node, limit: _resolve_source( + request, + node, + limit, + ), + preview_node_id=request.preview_node_id, + ) + except PipelineExecutionError as exc: + raise BackendExecutionError( + str(exc), + code="backend.reference", + node_id=exc.node_id, + diagnostics=tuple(exc.diagnostics), + ) from exc + observed_batch = TypedBatch.from_rows(result.rows) + batch = TypedBatch.from_rows( + result.rows, + schema=canonical_result_schema( + request, + observed_batch.schema, + ), + ) + try: + batch.ensure_within( + max_rows=request.budget.max_output_rows, + max_bytes=request.budget.max_batch_bytes, + ) + except ValueError as exc: + raise BackendExecutionError( + str(exc), + code="backend.budget", + ) from exc + lineage = request.graph.lineage.model_copy( + update={ + "source_fingerprints": tuple(result.source_fingerprints) + } + ) + contract = IrExecutionResult( + backend=self.name, + backend_version=self.version, + graph_semantic_hash=request.graph.semantic_hash, + result_schema=batch.schema, + row_count=result.total_rows, + byte_count=batch.byte_count, + truncated=result.truncated, + diagnostics=tuple(result.diagnostics), + lineage=lineage, + ) + return BackendExecutionResult( + contract=contract, + batch=batch, + node_diagnostics=tuple(result.node_diagnostics), + node_preview=result.node_preview, + metadata={ + "input_row_count": result.input_row_count, + "source_fingerprints": tuple(result.source_fingerprints), + }, + ) + + +def _validate_source_batches(request: BackendExecutionRequest) -> None: + for source in request.sources.values(): + try: + source.batch.ensure_within( + max_rows=10_000, + max_bytes=request.budget.max_batch_bytes, + ) + except ValueError as exc: + raise BackendExecutionError( + str(exc), + code="backend.source_budget", + node_id=source.node_id, + ) from exc + + +def _resolve_source( + request: BackendExecutionRequest, + node: GraphNode, + limit: int, +) -> ResolvedSource: + source = _source_for_node(request, node) + if source is None: + raise PipelineExecutionError( + f"No typed source batch was supplied for node {node.id!r}.", + node_id=node.id, + ) + rows = source.batch.to_rows() + return ResolvedSource( + rows=tuple(rows[:limit]), + source_ref=source.source_ref, + provider=source.provider, + fingerprint=source.fingerprint, + total_rows=source.total_rows, + truncated=source.truncated or len(rows) > limit, + ) + + +def _source_for_node( + request: BackendExecutionRequest, + node: GraphNode, +) -> BackendSource | None: + candidates = ( + node.id, + str(node.config.get("source_ref") or ""), + str(node.config.get("source_name") or ""), + ) + return next( + ( + request.sources[key] + for key in candidates + if key and key in request.sources + ), + None, + ) + + +__all__ = ["ReferenceExecutionBackend"] diff --git a/src/govoplan_dataflow/backend/backends/registry.py b/src/govoplan_dataflow/backend/backends/registry.py new file mode 100644 index 0000000..e66e22d --- /dev/null +++ b/src/govoplan_dataflow/backend/backends/registry.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +import threading +from typing import Mapping + +from govoplan_dataflow.backend.backends.base import ( + BackendExecutionError, + BackendExecutionRequest, + BackendExecutionResult, + BackendSource, + ExecutionBackend, + ExecutionBudget, +) +from govoplan_dataflow.backend.backends.duckdb import DuckDbExecutionBackend +from govoplan_dataflow.backend.backends.reference import ReferenceExecutionBackend +from govoplan_dataflow.backend.ir import TypedGraphIr, graph_to_ir +from govoplan_dataflow.backend.planner import PlanningError, plan_execution +from govoplan_dataflow.backend.schemas import PipelineGraph + + +class ExecutionBackendRegistry: + def __init__(self) -> None: + self._backends: dict[str, ExecutionBackend] = {} + self._slots: dict[str, threading.BoundedSemaphore] = {} + + def register( + self, + backend: ExecutionBackend, + *, + max_concurrency: int, + ) -> None: + if backend.name in self._backends: + raise ValueError( + f"Duplicate Dataflow execution backend {backend.name!r}." + ) + if max_concurrency < 1: + raise ValueError("Backend concurrency must be positive.") + self._backends[backend.name] = backend + self._slots[backend.name] = threading.BoundedSemaphore( + max_concurrency + ) + + def names(self) -> tuple[str, ...]: + return tuple(self._backends) + + def available(self) -> tuple[str, ...]: + return tuple( + name + for name, backend in self._backends.items() + if backend.available() + ) + + def resolve( + self, + name: str, + request: BackendExecutionRequest, + ) -> ExecutionBackend: + if name == "auto": + return self._automatic_backend(request) + backend = self._backends.get(name) + if backend is None: + raise BackendExecutionError( + f"Unknown Dataflow execution backend {name!r}.", + code="backend.unknown", + ) + if not backend.available(): + raise BackendExecutionError( + f"Dataflow execution backend {name!r} is unavailable.", + code="backend.unavailable", + ) + if not backend.supports(request): + raise BackendExecutionError( + f"Dataflow execution backend {name!r} does not support this graph.", + code="backend.unsupported_graph", + ) + return backend + + def execute( + self, + name: str, + request: BackendExecutionRequest, + ) -> BackendExecutionResult: + backend = self.resolve(name, request) + slot = self._slots[backend.name] + acquired = slot.acquire(timeout=request.budget.max_wall_seconds) + if not acquired: + raise BackendExecutionError( + f"Dataflow execution backend {backend.name!r} is at capacity.", + code="backend.capacity", + ) + try: + return backend.execute(request) + finally: + slot.release() + + def _automatic_backend( + self, + request: BackendExecutionRequest, + ) -> ExecutionBackend: + for name in ("duckdb", "reference"): + backend = self._backends.get(name) + if ( + backend is not None + and backend.available() + and backend.supports(request) + ): + return backend + raise BackendExecutionError( + "No Dataflow execution backend supports this graph.", + code="backend.unavailable", + ) + + +EXECUTION_BACKENDS = ExecutionBackendRegistry() +EXECUTION_BACKENDS.register( + ReferenceExecutionBackend(), + max_concurrency=4, +) +EXECUTION_BACKENDS.register( + DuckDbExecutionBackend(), + max_concurrency=1, +) + + +def execute_typed_graph( + graph: PipelineGraph | TypedGraphIr, + *, + backend: str = "auto", + sources: Mapping[str, BackendSource] | None = None, + budget: ExecutionBudget | None = None, + preview_node_id: str | None = None, + registry: ExecutionBackendRegistry = EXECUTION_BACKENDS, +) -> BackendExecutionResult: + typed_graph = graph if isinstance(graph, TypedGraphIr) else graph_to_ir(graph) + resolved_budget = budget or ExecutionBudget() + try: + plan = plan_execution(typed_graph) + except PlanningError as exc: + raise BackendExecutionError( + str(exc), + code="backend.planning", + diagnostics=exc.diagnostics, + ) from exc + request = BackendExecutionRequest( + plan=plan, + sources=sources or {}, + budget=resolved_budget, + preview_node_id=preview_node_id, + ) + return registry.execute(backend, request) + + +__all__ = [ + "EXECUTION_BACKENDS", + "ExecutionBackendRegistry", + "execute_typed_graph", +] diff --git a/src/govoplan_dataflow/backend/batches.py b/src/govoplan_dataflow/backend/batches.py new file mode 100644 index 0000000..9c4ff48 --- /dev/null +++ b/src/govoplan_dataflow/backend/batches.py @@ -0,0 +1,222 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from typing import Any, Iterable, Mapping + +from govoplan_dataflow.backend.ir import ( + IrField, + IrSchema, + schema_from_fields, + schema_from_rows, +) + + +class ArrowDependencyError(RuntimeError): + pass + + +@dataclass(frozen=True, slots=True) +class TypedBatch: + schema: IrSchema + columns: Mapping[str, tuple[Any, ...]] + row_count: int + byte_count: int + + @classmethod + def from_rows( + cls, + rows: Iterable[Mapping[str, Any]], + *, + schema: IrSchema | None = None, + ) -> "TypedBatch": + normalized = [dict(row) for row in rows] + resolved_schema = schema or schema_from_rows(normalized) + names = _column_names(normalized, resolved_schema) + columns = { + name: tuple(row.get(name) for row in normalized) + for name in names + } + return cls( + schema=resolved_schema, + columns=columns, + row_count=len(normalized), + byte_count=_estimated_bytes(normalized), + ) + + @classmethod + def from_arrow_ipc(cls, payload: bytes) -> "TypedBatch": + pa = _pyarrow() + with pa.ipc.open_stream(payload) as reader: + table = reader.read_all() + return cls.from_arrow_table(table) + + @classmethod + def from_arrow_table(cls, table: object) -> "TypedBatch": + pa = _pyarrow() + if not isinstance(table, pa.Table): + raise TypeError("Expected a pyarrow.Table.") + schema = schema_from_fields( + tuple( + IrField( + name=field.name, + type=_arrow_data_type(field.type), + nullable=field.nullable, + ) + for field in table.schema + ), + ) + columns = { + name: tuple(table.column(name).to_pylist()) + for name in table.column_names + } + return cls( + schema=schema, + columns=columns, + row_count=table.num_rows, + byte_count=table.nbytes, + ) + + def to_rows(self) -> list[dict[str, Any]]: + names = list(self.columns) + return [ + { + name: self.columns[name][index] + for name in names + } + for index in range(self.row_count) + ] + + def to_arrow_table(self) -> object: + pa = _pyarrow() + if self.columns: + return pa.table( + { + name: list(values) + for name, values in self.columns.items() + } + ) + fields = [ + pa.field( + field.name, + _pyarrow_type(field.type), + nullable=field.nullable, + ) + for field in self.schema.fields + ] + schema = pa.schema(fields) + return pa.Table.from_arrays( + [pa.array([], type=field.type) for field in fields], + schema=schema, + ) + + def to_arrow_ipc(self) -> bytes: + pa = _pyarrow() + table = self.to_arrow_table() + sink = pa.BufferOutputStream() + with pa.ipc.new_stream(sink, table.schema) as writer: + writer.write_table(table) + return sink.getvalue().to_pybytes() + + def ensure_within( + self, + *, + max_rows: int, + max_bytes: int, + ) -> None: + if self.row_count > max_rows: + raise ValueError( + f"Typed batch has {self.row_count} rows; limit is {max_rows}." + ) + if self.byte_count > max_bytes: + raise ValueError( + f"Typed batch uses {self.byte_count} bytes; limit is {max_bytes}." + ) + + +def arrow_available() -> bool: + try: + _pyarrow() + except ArrowDependencyError: + return False + return True + + +def _column_names( + rows: list[dict[str, Any]], + schema: IrSchema, +) -> list[str]: + names = [field.name for field in schema.fields] + seen = set(names) + for row in rows: + for name in row: + if name not in seen: + names.append(name) + seen.add(name) + return names + + +def _estimated_bytes(rows: list[dict[str, Any]]) -> int: + encoded = json.dumps( + rows, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=True, + default=str, + ) + return len(encoded.encode("utf-8")) + + +def _pyarrow() -> Any: + try: + import pyarrow as pa + except ImportError as exc: + raise ArrowDependencyError( + "Arrow execution requires the govoplan-dataflow analytics extra." + ) from exc + return pa + + +def _arrow_data_type(data_type: object) -> str: + pa = _pyarrow() + if pa.types.is_boolean(data_type): + return "boolean" + if pa.types.is_integer(data_type): + return "integer" + if pa.types.is_floating(data_type) or pa.types.is_decimal(data_type): + return "number" + if pa.types.is_date(data_type): + return "date" + if pa.types.is_timestamp(data_type): + return "datetime" + if pa.types.is_binary(data_type): + return "binary" + if pa.types.is_string(data_type) or pa.types.is_large_string(data_type): + return "string" + if pa.types.is_null(data_type): + return "null" + return "unknown" + + +def _pyarrow_type(data_type: str) -> object: + pa = _pyarrow() + mapping = { + "boolean": pa.bool_(), + "integer": pa.int64(), + "number": pa.float64(), + "date": pa.date32(), + "datetime": pa.timestamp("us", tz="UTC"), + "binary": pa.binary(), + "string": pa.string(), + "null": pa.null(), + "json": pa.string(), + "unknown": pa.null(), + } + return mapping.get(data_type, pa.null()) + + +__all__ = [ + "ArrowDependencyError", + "TypedBatch", + "arrow_available", +] diff --git a/src/govoplan_dataflow/backend/ir.py b/src/govoplan_dataflow/backend/ir.py new file mode 100644 index 0000000..fcbc7d4 --- /dev/null +++ b/src/govoplan_dataflow/backend/ir.py @@ -0,0 +1,491 @@ +from __future__ import annotations + +import hashlib +import json +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field + +from govoplan_dataflow.backend.expressions import ( + ExpressionError, + infer_expression_type, + parse_expression, +) +from govoplan_dataflow.backend.graph import topological_order +from govoplan_dataflow.backend.node_library import node_definition +from govoplan_dataflow.backend.schema_validation import ( + SchemaState, + propagate_graph_schemas, +) +from govoplan_dataflow.backend.schemas import ( + DataflowDiagnostic, + GraphEdge, + GraphNode, + GraphPosition, + PipelineGraph, +) + + +DATAFLOW_IR_VERSION = 1 +OPERATOR_CONTRACT_VERSION = 1 +DataType = Literal[ + "unknown", + "null", + "string", + "integer", + "number", + "boolean", + "date", + "datetime", + "binary", + "json", +] + + +class IrModel(BaseModel): + model_config = ConfigDict(frozen=True) + + +class IrField(IrModel): + name: str = Field(min_length=1, max_length=300) + type: DataType = "unknown" + nullable: bool = True + metadata: dict[str, Any] = Field(default_factory=dict) + + +class IrSchema(IrModel): + fields: tuple[IrField, ...] = () + open: bool = False + semantic_hash: str = Field(default="", max_length=64) + + +class IrPort(IrModel): + id: str = Field(min_length=1, max_length=80) + direction: Literal["input", "output"] + required: bool = True + multiple: bool = False + minimum_connections: int = Field(default=1, ge=0) + data_schema: IrSchema | None = None + + +class IrExpression(IrModel): + source: str = Field(min_length=1, max_length=10_000) + dialect: Literal["duckdb"] = "duckdb" + columns: tuple[str, ...] = () + result_type: DataType = "unknown" + semantic_hash: str = Field(min_length=64, max_length=64) + + +class IrParameter(IrModel): + id: str = Field(min_length=1, max_length=120) + type: DataType = "unknown" + required: bool = False + default: Any = None + description: str | None = Field(default=None, max_length=1_000) + + +class IrNode(IrModel): + id: str = Field(min_length=1, max_length=100) + type: str = Field(min_length=1, max_length=100) + operator_version: int = Field(default=OPERATOR_CONTRACT_VERSION, ge=1) + label: str = Field(min_length=1, max_length=300) + position: GraphPosition + config: dict[str, Any] = Field(default_factory=dict) + input_ports: tuple[IrPort, ...] = () + output_ports: tuple[IrPort, ...] = () + expressions: tuple[IrExpression, ...] = () + output_schema: IrSchema = Field(default_factory=IrSchema) + + +class IrEdge(IrModel): + id: str = Field(min_length=1, max_length=255) + source: str = Field(min_length=1, max_length=100) + target: str = Field(min_length=1, max_length=100) + source_port: str = Field(default="output", min_length=1, max_length=80) + target_port: str = Field(default="input", min_length=1, max_length=80) + + +class IrLineage(IrModel): + generated_from: Literal["graph", "sql", "template"] = "graph" + source_fingerprints: tuple[dict[str, Any], ...] = () + parent_semantic_hashes: tuple[str, ...] = () + loss_diagnostics: tuple[DataflowDiagnostic, ...] = () + + +class TypedGraphIr(IrModel): + ir_version: Literal[1] = DATAFLOW_IR_VERSION + graph_schema_version: int = Field(default=1, ge=1) + nodes: tuple[IrNode, ...] = Field(default=(), max_length=100) + edges: tuple[IrEdge, ...] = Field(default=(), max_length=200) + parameters: tuple[IrParameter, ...] = Field(default=(), max_length=100) + lineage: IrLineage = Field(default_factory=IrLineage) + semantic_hash: str = Field(min_length=64, max_length=64) + + +class IrExecutionResult(IrModel): + ir_version: Literal[1] = DATAFLOW_IR_VERSION + backend: str = Field(min_length=1, max_length=100) + backend_version: str = Field(min_length=1, max_length=100) + graph_semantic_hash: str = Field(min_length=64, max_length=64) + result_schema: IrSchema + row_count: int = Field(ge=0) + byte_count: int = Field(ge=0) + truncated: bool = False + diagnostics: tuple[DataflowDiagnostic, ...] = () + lineage: IrLineage = Field(default_factory=IrLineage) + + +def graph_to_ir( + graph: PipelineGraph, + *, + parameters: tuple[IrParameter, ...] = (), + lineage: IrLineage | None = None, +) -> TypedGraphIr: + ordered, cyclic = topological_order(graph) + schemas, schema_diagnostics = ( + ({}, []) + if cyclic + else propagate_graph_schemas(graph, ordered=ordered) + ) + resolved_lineage = _lineage_with_diagnostics( + lineage or IrLineage(), + schema_diagnostics, + ) + nodes = tuple( + _node_ir(node, schemas.get(node.id, _unknown_schema())) + for node in graph.nodes + ) + edges = tuple(_edge_ir(edge) for edge in graph.edges) + semantic_hash = _graph_semantic_hash( + graph_schema_version=graph.schema_version, + nodes=nodes, + edges=edges, + parameters=parameters, + ) + return TypedGraphIr( + graph_schema_version=graph.schema_version, + nodes=nodes, + edges=edges, + parameters=parameters, + lineage=resolved_lineage, + semantic_hash=semantic_hash, + ) + + +def ir_to_graph(ir: TypedGraphIr) -> PipelineGraph: + return PipelineGraph( + schema_version=ir.graph_schema_version, + nodes=[ + GraphNode( + id=node.id, + type=node.type, + label=node.label, + position=node.position.model_copy(deep=True), + config=dict(node.config), + ) + for node in ir.nodes + ], + edges=[ + GraphEdge( + id=edge.id, + source=edge.source, + target=edge.target, + source_port=edge.source_port, + target_port=edge.target_port, + ) + for edge in ir.edges + ], + ) + + +def schema_from_rows(rows: list[dict[str, Any]]) -> IrSchema: + names: list[str] = [] + seen: set[str] = set() + for row in rows: + for name in row: + normalized = str(name) + if normalized not in seen: + names.append(normalized) + seen.add(normalized) + columns = { + name: [row.get(name) for row in rows] + for name in names + } + fields = tuple( + IrField( + name=name, + type=_values_type(values), + nullable=any(value is None for value in values), + ) + for name, values in columns.items() + ) + return _schema(fields) + + +def schema_from_fields( + fields: tuple[IrField, ...], + *, + open: bool = False, +) -> IrSchema: + return _schema(fields, open=open) + + +def _node_ir(node: GraphNode, output_state: SchemaState) -> IrNode: + definition = node_definition(node.type) + input_ports = ( + tuple( + IrPort( + id=port.id, + direction="input", + required=port.required, + multiple=port.multiple, + minimum_connections=port.minimum_connections, + ) + for port in definition.input_ports + ) + if definition + else () + ) + output_schema = _state_schema(output_state) + output_ports = ( + tuple( + IrPort( + id=port.id, + direction="output", + required=port.required, + multiple=port.multiple, + minimum_connections=port.minimum_connections, + data_schema=output_schema, + ) + for port in definition.output_ports + ) + if definition + else () + ) + return IrNode( + id=node.id, + type=node.type, + label=node.label, + position=node.position.model_copy(deep=True), + config=dict(node.config), + input_ports=input_ports, + output_ports=output_ports, + expressions=_node_expressions(node, output_state), + output_schema=output_schema, + ) + + +def _node_expressions( + node: GraphNode, + output_state: SchemaState, +) -> tuple[IrExpression, ...]: + source = node.config.get("expression") + if node.type not in {"expression", "filter.expression"} or not source: + return () + try: + parsed = parse_expression(str(source)) + except ExpressionError: + return () + inferred = infer_expression_type(parsed, output_state.types) + return ( + IrExpression( + source=parsed.source, + columns=parsed.columns, + result_type=_data_type(inferred), + semantic_hash=_hash_payload( + { + "dialect": "duckdb", + "source": parsed.sql(), + "result_type": inferred, + } + ), + ), + ) + + +def _edge_ir(edge: GraphEdge) -> IrEdge: + return IrEdge( + id=edge.id, + source=edge.source, + target=edge.target, + source_port=edge.source_port, + target_port=edge.target_port, + ) + + +def _state_schema(state: SchemaState) -> IrSchema: + return _schema( + tuple( + IrField( + name=name, + type=_data_type(state.type_of(name)), + nullable=True, + ) + for name in sorted(state.columns) + ), + open=state.open, + ) + + +def _schema( + fields: tuple[IrField, ...], + *, + open: bool = False, +) -> IrSchema: + semantic_hash = _hash_payload( + { + "fields": [ + { + "name": item.name, + "type": item.type, + "nullable": item.nullable, + } + for item in fields + ], + "open": open, + } + ) + return IrSchema( + fields=fields, + open=open, + semantic_hash=semantic_hash, + ) + + +def _graph_semantic_hash( + *, + graph_schema_version: int, + nodes: tuple[IrNode, ...], + edges: tuple[IrEdge, ...], + parameters: tuple[IrParameter, ...], +) -> str: + return _hash_payload( + { + "ir_version": DATAFLOW_IR_VERSION, + "graph_schema_version": graph_schema_version, + "nodes": [ + { + "id": node.id, + "type": node.type, + "operator_version": node.operator_version, + "config": node.config, + "expressions": [ + expression.model_dump(mode="json") + for expression in node.expressions + ], + } + for node in nodes + ], + "edges": [ + edge.model_dump(mode="json") + for edge in edges + ], + "parameters": [ + parameter.model_dump(mode="json") + for parameter in parameters + ], + } + ) + + +def _lineage_with_diagnostics( + lineage: IrLineage, + diagnostics: list[DataflowDiagnostic], +) -> IrLineage: + if not diagnostics: + return lineage + return lineage.model_copy( + update={ + "loss_diagnostics": ( + *lineage.loss_diagnostics, + *diagnostics, + ) + } + ) + + +def _values_type(values: list[Any]) -> DataType: + concrete = { + _value_type(value) + for value in values + if value is not None + } + if not concrete: + return "null" + if concrete <= {"integer", "number"}: + return "number" if "number" in concrete else "integer" + return concrete.pop() if len(concrete) == 1 else "unknown" + + +def _value_type(value: Any) -> DataType: + from datetime import date, datetime + from decimal import Decimal + + if isinstance(value, bool): + return "boolean" + if isinstance(value, int): + return "integer" + if isinstance(value, (float, Decimal)): + return "number" + if isinstance(value, datetime): + return "datetime" + if isinstance(value, date): + return "date" + if isinstance(value, bytes): + return "binary" + if isinstance(value, str): + return "string" + if isinstance(value, (dict, list, tuple)): + return "json" + return "unknown" + + +def _data_type(value: str) -> DataType: + known = { + "unknown", + "null", + "string", + "integer", + "number", + "boolean", + "date", + "datetime", + "binary", + "json", + } + return value if value in known else "unknown" # type: ignore[return-value] + + +def _unknown_schema() -> SchemaState: + return SchemaState(frozenset(), open=True) + + +def _hash_payload(payload: object) -> str: + encoded = json.dumps( + payload, + sort_keys=True, + separators=(",", ":"), + ensure_ascii=True, + default=str, + ) + return hashlib.sha256(encoded.encode("utf-8")).hexdigest() + + +__all__ = [ + "DATAFLOW_IR_VERSION", + "DataType", + "IrEdge", + "IrExecutionResult", + "IrExpression", + "IrField", + "IrLineage", + "IrNode", + "IrParameter", + "IrPort", + "IrSchema", + "TypedGraphIr", + "graph_to_ir", + "ir_to_graph", + "schema_from_fields", + "schema_from_rows", +] diff --git a/src/govoplan_dataflow/backend/planner.py b/src/govoplan_dataflow/backend/planner.py new file mode 100644 index 0000000..c2ce3fe --- /dev/null +++ b/src/govoplan_dataflow/backend/planner.py @@ -0,0 +1,99 @@ +from __future__ import annotations + +import hashlib +from dataclasses import dataclass + +from govoplan_dataflow.backend.graph import topological_order, validate_graph +from govoplan_dataflow.backend.ir import TypedGraphIr, ir_to_graph +from govoplan_dataflow.backend.schemas import DataflowDiagnostic +from govoplan_dataflow.backend.sql_compiler import ( + SqlCompilationError, + render_sql, +) + + +@dataclass(frozen=True, slots=True) +class ExecutionPlan: + graph: TypedGraphIr + ordered_node_ids: tuple[str, ...] + diagnostics: tuple[DataflowDiagnostic, ...] + generated_sql: str | None + sql_diagnostics: tuple[DataflowDiagnostic, ...] + semantic_hash: str + + +class PlanningError(ValueError): + def __init__(self, diagnostics: tuple[DataflowDiagnostic, ...]) -> None: + super().__init__( + diagnostics[0].message + if diagnostics + else "Dataflow planning failed." + ) + self.diagnostics = diagnostics + + +def plan_execution(graph: TypedGraphIr) -> ExecutionPlan: + pipeline_graph = ir_to_graph(graph) + diagnostics = tuple(validate_graph(pipeline_graph)) + errors = tuple( + item + for item in diagnostics + if item.severity == "error" + ) + if errors: + raise PlanningError(errors) + ordered, cyclic = topological_order(pipeline_graph) + if cyclic: + cycle = DataflowDiagnostic( + severity="error", + code="graph.cycle", + message="A cyclic graph cannot be planned.", + ) + raise PlanningError((cycle,)) + sql_text, sql_diagnostics = _generated_sql(pipeline_graph) + semantic_hash = _plan_hash( + graph.semantic_hash, + ordered, + sql_text=sql_text, + ) + return ExecutionPlan( + graph=graph, + ordered_node_ids=tuple(ordered), + diagnostics=diagnostics, + generated_sql=sql_text, + sql_diagnostics=sql_diagnostics, + semantic_hash=semantic_hash, + ) + + +def _generated_sql( + graph, +) -> tuple[str | None, tuple[DataflowDiagnostic, ...]]: + try: + sql_text, diagnostics = render_sql(graph) + except SqlCompilationError as exc: + return None, tuple(exc.diagnostics) + return sql_text, tuple(diagnostics) + + +def _plan_hash( + graph_hash: str, + ordered: list[str], + *, + sql_text: str | None, +) -> str: + payload = "\n".join( + ( + graph_hash, + *ordered, + (sql_text or "").strip(), + ) + ) + return hashlib.sha256(payload.encode("utf-8")).hexdigest() + + +__all__ = [ + "ExecutionPlan", + "PlanningError", + "plan_execution", +] diff --git a/src/govoplan_dataflow/backend/schema_validation.py b/src/govoplan_dataflow/backend/schema_validation.py index 632cbcf..9c45855 100644 --- a/src/govoplan_dataflow/backend/schema_validation.py +++ b/src/govoplan_dataflow/backend/schema_validation.py @@ -55,6 +55,15 @@ def validate_graph_schemas( *, ordered: list[str], ) -> list[DataflowDiagnostic]: + _, diagnostics = propagate_graph_schemas(graph, ordered=ordered) + return diagnostics + + +def propagate_graph_schemas( + graph: PipelineGraph, + *, + ordered: list[str], +) -> tuple[dict[str, SchemaState], list[DataflowDiagnostic]]: node_by_id = {node.id: node for node in graph.nodes} inputs = _graph_inputs_by_port(graph) schemas: dict[str, SchemaState] = {} @@ -81,7 +90,7 @@ def validate_graph_schemas( ) schemas[node.id] = result.state diagnostics.extend(result.diagnostics) - return diagnostics + return schemas, diagnostics def _graph_inputs_by_port( diff --git a/src/govoplan_dataflow/backend/schemas.py b/src/govoplan_dataflow/backend/schemas.py index e72522c..01d9780 100644 --- a/src/govoplan_dataflow/backend/schemas.py +++ b/src/govoplan_dataflow/backend/schemas.py @@ -193,6 +193,7 @@ class PipelinePreviewRequest(BaseModel): pattern=r"^[A-Za-z0-9_.:-]+$", ) row_limit: int = Field(default=100, ge=1, le=500) + execution_backend: Literal["reference", "duckdb", "auto"] = "reference" class PreviewColumn(BaseModel): diff --git a/src/govoplan_dataflow/backend/service.py b/src/govoplan_dataflow/backend/service.py index 7ba3a16..d1d68a0 100644 --- a/src/govoplan_dataflow/backend/service.py +++ b/src/govoplan_dataflow/backend/service.py @@ -25,6 +25,13 @@ from govoplan_core.core.datasources import ( datasource_publication, ) from govoplan_core.db.base import utcnow +from govoplan_dataflow.backend.backends import ( + BackendExecutionError, + BackendSource, + ExecutionBudget, + execute_typed_graph, +) +from govoplan_dataflow.backend.batches import TypedBatch from govoplan_dataflow.backend.db.models import ( DataflowPipeline, DataflowPipelineRevision, @@ -32,6 +39,7 @@ from govoplan_dataflow.backend.db.models import ( ) from govoplan_dataflow.backend.executor import ( EXECUTOR_VERSION, + MAX_SOURCE_ROWS, PipelineExecutionError, PipelineExecutionResult, ResolvedSource, @@ -50,6 +58,7 @@ from govoplan_dataflow.backend.graph import ( from govoplan_dataflow.backend.schemas import ( DataflowDiagnostic, GraphNode, + NodePreviewResult, PipelineCreateRequest, PipelineDeriveRequest, PipelineDraftRequest, @@ -62,6 +71,7 @@ from govoplan_dataflow.backend.schemas import ( PipelineSqlResponse, PipelineUpdateRequest, PipelineValidationResponse, + PreviewColumn, ) from govoplan_dataflow.backend.sql_compiler import ( SqlCompilationError, @@ -588,49 +598,13 @@ def preview_pipeline( started_at = utcnow() run: DataflowRun | None = None try: - provider = datasource_catalogue(registry) - - def resolve_source(node: GraphNode, limit: int) -> ResolvedSource: - if provider is None: - raise PipelineExecutionError( - "Datasource-backed preview requires the Datasources catalogue capability.", - node_id=node.id, - ) - if principal is None: - raise PipelineExecutionError( - "Datasource-backed preview requires a tenant API principal.", - node_id=node.id, - ) - try: - resolved = provider.read_datasource( - session, - principal, - request=DatasourceReadRequest( - datasource_ref=str(node.config["source_ref"]), - consistency=str( - node.config.get("consistency") or "current" - ), # type: ignore[arg-type] - limit=limit, - expected_fingerprint=_clean_optional( - node.config.get("expected_fingerprint") - ), - ), - ) - except DatasourceError 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.datasource.ref, - provider=resolved.datasource.provider or "datasources", - fingerprint=resolved.datasource.fingerprint, - total_rows=resolved.total_rows, - truncated=resolved.truncated, - ) - - result = execute_preview( + result, executor_version = _execute_pipeline_preview( graph, + session=session, + principal=principal, + registry=registry, + backend=payload.execution_backend, row_limit=payload.row_limit, - source_resolver=resolve_source, preview_node_id=payload.preview_node_id, ) status = "succeeded" @@ -645,6 +619,11 @@ def preview_pipeline( source_fingerprints = result.source_fingerprints input_row_count = result.input_row_count except PipelineExecutionError as exc: + executor_version = ( + payload.execution_backend + if payload.execution_backend != "reference" + else EXECUTOR_VERSION + ) status = "failed" error = str(exc) diagnostics = [ @@ -672,7 +651,7 @@ def preview_pipeline( pipeline_revision_id=revision.id, run_type="preview", status=status, - executor_version=EXECUTOR_VERSION, + executor_version=executor_version, definition_hash=graph_hash, source_fingerprints=source_fingerprints, result_schema=[item.model_dump(mode="json") for item in columns], @@ -702,7 +681,162 @@ def preview_pipeline( source_fingerprints=source_fingerprints, input_row_count=input_row_count, definition_hash=graph_hash, - executor_version=EXECUTOR_VERSION, + executor_version=executor_version, + ) + + +def _execute_pipeline_preview( + graph: PipelineGraph, + *, + session: Session, + principal: ApiPrincipal | None, + registry: object | None, + backend: str, + row_limit: int, + preview_node_id: str | None, +) -> tuple[PipelineExecutionResult, str]: + source_resolver = _preview_source_resolver( + session=session, + principal=principal, + registry=registry, + ) + if backend == "reference": + return ( + execute_preview( + graph, + row_limit=row_limit, + source_resolver=source_resolver, + preview_node_id=preview_node_id, + ), + EXECUTOR_VERSION, + ) + sources = _typed_backend_sources( + graph, + source_resolver=source_resolver, + ) + try: + result = execute_typed_graph( + graph, + backend=backend, + sources=sources, + budget=ExecutionBudget(max_output_rows=row_limit), + preview_node_id=preview_node_id, + ) + except BackendExecutionError as exc: + raise PipelineExecutionError( + str(exc), + node_id=exc.node_id, + diagnostics=tuple(exc.diagnostics), + ) from exc + columns = [ + PreviewColumn( + name=field.name, + type=field.type, + nullable=field.nullable, + ) + for field in result.batch.schema.fields + ] + source_fingerprints = list( + result.metadata.get( + "source_fingerprints", + result.contract.lineage.source_fingerprints, + ) + ) + node_preview = _typed_node_preview( + result, + preview_node_id=preview_node_id, + columns=columns, + ) + return ( + PipelineExecutionResult( + rows=result.rows, + total_rows=result.contract.row_count, + truncated=result.contract.truncated, + columns=columns, + diagnostics=list(result.contract.diagnostics), + node_diagnostics=list(result.node_diagnostics), + node_preview=node_preview, + source_fingerprints=source_fingerprints, + input_row_count=int( + result.metadata.get( + "input_row_count", + sum( + int( + item.get( + "row_count", + item.get("total_rows", 0), + ) + ) + for item in source_fingerprints + ), + ) + ), + ), + result.contract.backend_version, + ) + + +def _preview_source_resolver( + *, + session: Session, + principal: ApiPrincipal | None, + registry: object | None, +): + if principal is not None: + return _datasource_source_resolver( + session=session, + principal=principal, + registry=registry, + ) + + def unavailable(node: GraphNode, _limit: int) -> ResolvedSource: + raise PipelineExecutionError( + "Datasource-backed preview requires a tenant API principal.", + node_id=node.id, + ) + + return unavailable + + +def _typed_backend_sources( + graph: PipelineGraph, + *, + source_resolver, +) -> dict[str, BackendSource]: + sources: dict[str, BackendSource] = {} + for node in graph.nodes: + if node.type != "source.reference": + continue + resolved = source_resolver(node, MAX_SOURCE_ROWS) + sources[node.id] = BackendSource( + node_id=node.id, + batch=TypedBatch.from_rows(resolved.rows), + source_ref=resolved.source_ref, + provider=resolved.provider, + fingerprint=resolved.fingerprint, + total_rows=resolved.total_rows, + truncated=resolved.truncated, + source_name=str(node.config.get("source_name") or ""), + ) + return sources + + +def _typed_node_preview( + result, + *, + preview_node_id: str | None, + columns: list[PreviewColumn], +) -> NodePreviewResult | None: + if result.node_preview is not None: + return result.node_preview + if preview_node_id is None: + return None + return NodePreviewResult( + node_id=preview_node_id, + columns=columns, + rows=result.rows, + total_rows=result.contract.row_count, + truncated=result.contract.truncated, ) diff --git a/src/govoplan_dataflow/backend/sql_compiler.py b/src/govoplan_dataflow/backend/sql_compiler.py index ada79e3..a9a350c 100644 --- a/src/govoplan_dataflow/backend/sql_compiler.py +++ b/src/govoplan_dataflow/backend/sql_compiler.py @@ -1409,7 +1409,7 @@ def _render_select( else str(field_config.get("alias") or column) ) expression: exp.Expression = column_expression(column) - if alias != column: + if alias != expression.alias_or_name: expression = expression.as_(alias) state.select_expressions.append(expression) state.selected = True diff --git a/tests/test_service.py b/tests/test_service.py index bb77fc8..49ea3c2 100644 --- a/tests/test_service.py +++ b/tests/test_service.py @@ -18,6 +18,7 @@ from govoplan_core.core.datasources import ( DatasourcePublicationResult, ) from govoplan_core.db.base import Base +from govoplan_dataflow.backend.backends.duckdb import DuckDbExecutionBackend from govoplan_dataflow.backend.db.models import ( DataflowPipeline, DataflowPipelineRevision, @@ -320,6 +321,37 @@ class DataflowServiceTests(unittest.TestCase): self.session.scalar(select(func.count()).select_from(DataflowRun)), ) + @unittest.skipUnless( + DuckDbExecutionBackend().available(), + "analytics extra is not installed", + ) + def test_saved_preview_can_use_isolated_duckdb_backend(self) -> None: + pipeline = self._create() + + response = preview_pipeline( + self.session, + tenant_id="tenant-1", + actor_id="user-1", + payload=PipelinePreviewRequest( + pipeline_id=pipeline.id, + execution_backend="duckdb", + ), + principal=principal(), + ) + self.session.commit() + + run = self.session.get(DataflowRun, response.run_id) + self.assertEqual("succeeded", response.status) + self.assertEqual( + [{"id": 2, "amount": 15}, {"id": 3, "amount": 25}], + response.rows, + ) + self.assertEqual(2, response.total_rows) + self.assertFalse(response.truncated) + self.assertEqual(3, response.input_row_count) + self.assertTrue(response.executor_version.startswith("duckdb-isolated-v1")) + self.assertEqual(response.executor_version, run.executor_version) + 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"}] diff --git a/tests/test_typed_execution.py b/tests/test_typed_execution.py new file mode 100644 index 0000000..4cb646e --- /dev/null +++ b/tests/test_typed_execution.py @@ -0,0 +1,258 @@ +from __future__ import annotations + +import unittest + +from govoplan_dataflow.backend.backends import ( + EXECUTION_BACKENDS, + ExecutionBudget, + execute_typed_graph, +) +from govoplan_dataflow.backend.backends.duckdb import DuckDbExecutionBackend +from govoplan_dataflow.backend.batches import TypedBatch +from govoplan_dataflow.backend.ir import ( + DATAFLOW_IR_VERSION, + graph_to_ir, + ir_to_graph, +) +from govoplan_dataflow.backend.planner import plan_execution +from govoplan_dataflow.backend.schemas import ( + GraphEdge, + GraphNode, + GraphPosition, + PipelineGraph, +) +from govoplan_dataflow.backend.sql_compiler import compile_sql + + +def inline_source( + node_id: str = "source", + source_name: str = "records", + rows: list[dict] | None = None, +) -> GraphNode: + return GraphNode( + id=node_id, + type="source.inline", + label=source_name, + position=GraphPosition(x=40, y=120), + config={ + "source_name": source_name, + "rows": rows + or [ + {"department": "a", "amount": 10}, + {"department": "a", "amount": 5}, + {"department": "b", "amount": 20}, + ], + }, + ) + + +def simple_graph() -> PipelineGraph: + source = inline_source() + derive = GraphNode( + id="derive", + type="expression", + label="Double amount", + position=GraphPosition(x=280, y=120), + config={ + "target_column": "double_amount", + "expression": "amount * 2", + "result_type": "integer", + }, + ) + output = GraphNode( + id="output", + type="output", + label="Output", + position=GraphPosition(x=520, y=120), + config={}, + ) + return PipelineGraph( + nodes=[source, derive, output], + edges=[ + GraphEdge(id="edge-1", source=source.id, target=derive.id), + GraphEdge(id="edge-2", source=derive.id, target=output.id), + ], + ) + + +class TypedIrTests(unittest.TestCase): + def test_ir_roundtrip_preserves_identity_layout_and_expressions(self) -> None: + graph = simple_graph() + + ir = graph_to_ir(graph) + roundtrip = ir_to_graph(ir) + + self.assertEqual(DATAFLOW_IR_VERSION, ir.ir_version) + self.assertEqual(64, len(ir.semantic_hash)) + self.assertEqual( + graph.model_dump(mode="json"), + roundtrip.model_dump(mode="json"), + ) + expression = ir.nodes[1].expressions[0] + self.assertEqual(("amount",), expression.columns) + self.assertEqual("integer", expression.result_type) + self.assertEqual( + {"amount", "department", "double_amount"}, + {field.name for field in ir.nodes[-1].output_schema.fields}, + ) + + def test_semantic_hash_ignores_layout_but_changes_with_behavior(self) -> None: + graph = simple_graph() + moved = graph.model_copy(deep=True) + moved.nodes[1].position = GraphPosition(x=900, y=700) + changed = graph.model_copy(deep=True) + changed.nodes[1].config["expression"] = "amount * 3" + + self.assertEqual( + graph_to_ir(graph).semantic_hash, + graph_to_ir(moved).semantic_hash, + ) + self.assertNotEqual( + graph_to_ir(graph).semantic_hash, + graph_to_ir(changed).semantic_hash, + ) + + def test_typed_batch_is_columnar_and_roundtrips_rows(self) -> None: + rows = [ + {"name": "Ada", "active": True}, + {"name": "Grace", "active": None}, + ] + + batch = TypedBatch.from_rows(rows) + + self.assertEqual(2, batch.row_count) + self.assertEqual(("Ada", "Grace"), batch.columns["name"]) + self.assertEqual(rows, batch.to_rows()) + self.assertGreater(batch.byte_count, 0) + + def test_physical_plan_pins_order_sql_and_semantic_hash(self) -> None: + plan = plan_execution(graph_to_ir(simple_graph())) + + self.assertEqual( + ("source", "derive", "output"), + plan.ordered_node_ids, + ) + self.assertIn("amount * 2", plan.generated_sql) + self.assertEqual(64, len(plan.semantic_hash)) + self.assertEqual((), plan.sql_diagnostics) + + +class ExecutionBackendTests(unittest.TestCase): + def test_reference_backend_executes_typed_ir_with_budgets(self) -> None: + result = execute_typed_graph( + simple_graph(), + backend="reference", + budget=ExecutionBudget(max_output_rows=2), + ) + + self.assertEqual("reference", result.contract.backend) + self.assertEqual(3, result.contract.row_count) + self.assertTrue(result.contract.truncated) + self.assertEqual( + [ + {"department": "a", "amount": 10, "double_amount": 20}, + {"department": "a", "amount": 5, "double_amount": 10}, + ], + result.rows, + ) + + def test_backend_registry_keeps_reference_available(self) -> None: + self.assertIn("reference", EXECUTION_BACKENDS.names()) + self.assertIn("reference", EXECUTION_BACKENDS.available()) + + @unittest.skipUnless( + DuckDbExecutionBackend().available(), + "analytics extra is not installed", + ) + def test_duckdb_matches_reference_for_aggregate_join_and_union(self) -> None: + graphs = [ + compile_sql( + """ + SELECT department, COUNT(*) AS records, SUM(amount) AS total + FROM records + GROUP BY department + ORDER BY department + """, + source_nodes=[inline_source()], + )[0], + compile_sql( + """ + SELECT records.department, labels.label + FROM records + LEFT JOIN labels ON records.department = labels.department + ORDER BY records.department + """, + source_nodes=[ + inline_source(), + inline_source( + "labels-source", + "labels", + [ + {"department": "a", "label": "Alpha"}, + {"department": "b", "label": "Beta"}, + ], + ), + ], + )[0], + compile_sql( + """ + SELECT * FROM first + UNION ALL BY NAME + SELECT * FROM second + ORDER BY amount + """, + source_nodes=[ + inline_source( + "first-source", + "first", + [{"amount": 2}, {"amount": None}], + ), + inline_source( + "second-source", + "second", + [{"amount": 1}], + ), + ], + )[0], + ] + for graph in graphs: + with self.subTest(nodes=[node.type for node in graph.nodes]): + reference = execute_typed_graph( + graph, + backend="reference", + budget=ExecutionBudget(max_output_rows=100), + ) + analytical = execute_typed_graph( + graph, + backend="duckdb", + budget=ExecutionBudget( + max_output_rows=100, + max_wall_seconds=10, + ), + ) + + self.assertEqual(reference.rows, analytical.rows) + self.assertEqual( + reference.contract.row_count, + analytical.contract.row_count, + ) + self.assertEqual( + [field.name for field in reference.batch.schema.fields], + [field.name for field in analytical.batch.schema.fields], + ) + self.assertEqual( + reference.batch.schema.semantic_hash, + analytical.batch.schema.semantic_hash, + ) + self.assertEqual( + reference.contract.diagnostics, + analytical.contract.diagnostics, + ) + self.assertEqual( + reference.contract.lineage.source_fingerprints, + analytical.contract.lineage.source_fingerprints, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/webui/src/api/dataflow.ts b/webui/src/api/dataflow.ts index a713dc2..b9ff9c2 100644 --- a/webui/src/api/dataflow.ts +++ b/webui/src/api/dataflow.ts @@ -478,6 +478,7 @@ export function previewDataflowPipeline( source_nodes?: PipelineGraphNode[]; preview_node_id?: string; row_limit?: number; + execution_backend?: "reference" | "duckdb" | "auto"; } ): Promise { return apiFetch(settings, "/api/v1/dataflow/preview", {