Add typed Arrow execution backend boundary
This commit is contained in:
@@ -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",
|
||||
]
|
||||
Reference in New Issue
Block a user