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
+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",
]