Add typed Arrow execution backend boundary

This commit is contained in:
2026-07-29 16:47:14 +02:00
parent 69509d5cc2
commit 8a11d7571b
17 changed files with 2264 additions and 49 deletions
@@ -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",
]
@@ -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",
]
@@ -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",
]
@@ -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"]
@@ -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",
]
+222
View File
@@ -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",
]
+491
View File
@@ -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",
]
+99
View File
@@ -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",
]
@@ -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(
+1
View File
@@ -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):
+177 -43
View File
@@ -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,
)
@@ -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