Files
govoplan-dataflow/src/govoplan_dataflow/backend/batches.py
T

223 lines
5.8 KiB
Python

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