feat: impact analysis

This commit is contained in:
Stefan
2026-08-25 16:40:32 +02:00
parent adf7cdf44a
commit 2d33e14dee
20 changed files with 1915 additions and 54 deletions
+7
View File
@@ -23,3 +23,10 @@ tokio-util = { version = "0.7", features = ["rt"] }
zen-engine = { path = "../../core/engine", features = ["arbitrary_precision"] }
zen-expression = { path = "../../core/expression", features = ["arbitrary_precision"] }
zen-tmpl = { path = "../../core/template" }
arrow-array = { version = "56", optional = true, features = ["ffi"] }
arrow-schema = { version = "56", optional = true }
[features]
default = []
data = ["zen-engine/data"]
arrow = ["data", "dep:arrow-array", "dep:arrow-schema"]
+158
View File
@@ -0,0 +1,158 @@
use std::sync::Arc;
use anyhow::anyhow;
use arrow_array::cast::AsArray;
use arrow_array::types::{Float32Type, Float64Type, Int16Type, Int32Type, Int64Type, Int8Type};
use arrow_array::{Array, RecordBatch};
use arrow_schema::DataType;
use rust_decimal::Decimal;
use zen_engine::Variable;
use zen_expression::variable::VariableMap;
use crate::data::insert_path;
pub fn inputs_from_batch(batch: &RecordBatch) -> Result<Vec<Variable>, anyhow::Error> {
let rows = batch.num_rows();
let schema = batch.schema();
let mut columns: Vec<(Vec<String>, Vec<Option<Variable>>)> = Vec::new();
for (index, field) in schema.fields().iter().enumerate() {
let name = field.name();
let segments = name.split('.').map(str::to_string).collect();
columns.push((segments, column_values(batch.column(index), name)?));
}
Ok((0..rows)
.map(|row| {
let mut map = VariableMap::default();
for (segments, values) in &columns {
if let Some(value) = &values[row] {
insert_path(&mut map, segments, value.clone());
}
}
Variable::from_object(map)
})
.collect())
}
fn column_values(
column: &Arc<dyn Array>,
name: &str,
) -> Result<Vec<Option<Variable>>, anyhow::Error> {
let rows = column.len();
let values = match column.data_type() {
DataType::Null => vec![None; rows],
DataType::Boolean => {
let array = column.as_boolean();
(0..rows)
.map(|row| (!array.is_null(row)).then(|| Variable::Bool(array.value(row))))
.collect()
}
DataType::Int8 => int_values(column.as_primitive::<Int8Type>(), |v| v as i64),
DataType::Int16 => int_values(column.as_primitive::<Int16Type>(), |v| v as i64),
DataType::Int32 => int_values(column.as_primitive::<Int32Type>(), |v| v as i64),
DataType::Int64 => int_values(column.as_primitive::<Int64Type>(), |v| v),
DataType::Float32 => float_values(column.as_primitive::<Float32Type>(), |v| v as f64),
DataType::Float64 => float_values(column.as_primitive::<Float64Type>(), |v| v),
DataType::Utf8 => {
let array = column.as_string::<i32>();
(0..rows)
.map(|row| {
(!array.is_null(row)).then(|| Variable::String(array.value(row).into()))
})
.collect()
}
DataType::LargeUtf8 => {
let array = column.as_string::<i64>();
(0..rows)
.map(|row| {
(!array.is_null(row)).then(|| Variable::String(array.value(row).into()))
})
.collect()
}
DataType::Dictionary(_, value_type) if **value_type == DataType::Utf8 => {
let array = column.as_any_dictionary();
let dictionary = array
.values()
.as_string_opt::<i32>()
.ok_or_else(|| anyhow!("column '{name}': unsupported dictionary values"))?;
let interned: Vec<Variable> = (0..dictionary.len())
.map(|index| Variable::String(dictionary.value(index).into()))
.collect();
let keys = array.normalized_keys();
(0..rows)
.map(|row| (!array.is_null(row)).then(|| interned[keys[row]].clone()))
.collect()
}
DataType::Struct(fields) => {
let array = column.as_struct();
let mut children: Vec<(String, Vec<Option<Variable>>)> =
Vec::with_capacity(fields.len());
for (field, child) in fields.iter().zip(array.columns()) {
children.push((field.name().clone(), column_values(child, field.name())?));
}
(0..rows)
.map(|row| {
if array.is_null(row) {
return None;
}
let mut map = VariableMap::with_capacity(children.len());
for (key, values) in &children {
if let Some(value) = &values[row] {
map.insert(key.as_str().into(), value.clone());
}
}
Some(Variable::from_object(map))
})
.collect()
}
DataType::List(_) => {
let array = column.as_list::<i32>();
let values = column_values(array.values(), name)?;
let offsets = array.offsets();
(0..rows)
.map(|row| {
if array.is_null(row) {
return None;
}
let start = offsets[row] as usize;
let end = offsets[row + 1] as usize;
let items: Vec<Variable> = values[start..end]
.iter()
.map(|value| value.clone().unwrap_or(Variable::Null))
.collect();
Some(Variable::from_array(items))
})
.collect()
}
other => return Err(anyhow!("column '{name}': unsupported arrow type {other}")),
};
Ok(values)
}
fn int_values<T: arrow_array::ArrowPrimitiveType>(
array: &arrow_array::PrimitiveArray<T>,
to_i64: impl Fn(T::Native) -> i64,
) -> Vec<Option<Variable>> {
(0..array.len())
.map(|row| {
(!array.is_null(row)).then(|| Variable::Number(Decimal::from(to_i64(array.value(row)))))
})
.collect()
}
fn float_values<T: arrow_array::ArrowPrimitiveType>(
array: &arrow_array::PrimitiveArray<T>,
to_f64: impl Fn(T::Native) -> f64,
) -> Vec<Option<Variable>> {
(0..array.len())
.map(|row| {
if array.is_null(row) {
return None;
}
Decimal::try_from(to_f64(array.value(row)))
.ok()
.map(Variable::Number)
})
.collect()
}
+185
View File
@@ -0,0 +1,185 @@
use std::sync::Arc;
use anyhow::anyhow;
use pyo3::types::PyBytes;
use pyo3::{pyclass, pymethods, Py, PyResult, Python};
use zen_engine::data::impact::ImpactAnalysis;
use zen_engine::EvaluationOptions;
use crate::engine::PyZenEngine;
use crate::mt::{block_on, worker_pool};
use zen_engine::Variable;
use zen_expression::variable::VariableMap;
/// Dotted top-level keys compose into nested objects — `customer.firstName`
/// becomes `{customer: {firstName: …}}`, matching how zen expressions resolve
/// member access. Spark's `df.toJSON()` emits flat dotted keys for dotted
/// column names; without this they would silently never match a rule.
pub(crate) fn normalize_dotted(variable: Variable) -> Variable {
let Variable::Object(object) = &variable else {
return variable;
};
if !object.borrow().keys().any(|key| key.contains('.')) {
return variable;
}
let mut out = VariableMap::new();
for (key, value) in object.borrow().iter() {
let segments: Vec<String> = key.as_str().split('.').map(str::to_string).collect();
insert_path(&mut out, &segments, value.clone());
}
Variable::from_object(out)
}
pub(crate) fn insert_path(map: &mut VariableMap, segments: &[String], value: Variable) {
let (first, rest) = match segments.split_first() {
Some(parts) => parts,
None => return,
};
if rest.is_empty() {
map.insert(first.as_str().into(), value);
return;
}
let key = first.as_str().into();
if !matches!(map.get(&key), Some(Variable::Object(_))) {
map.insert(key.clone(), Variable::from_object(VariableMap::new()));
}
if let Some(Variable::Object(child)) = map.get(&key) {
let child = child.clone();
insert_path(&mut child.borrow_mut(), rest, value);
}
}
#[pyclass]
#[pyo3(name = "ZenImpactAnalysis")]
pub struct PyZenImpactAnalysis {
inner: Arc<ImpactAnalysis>,
}
#[pymethods]
impl PyZenImpactAnalysis {
#[new]
pub fn new(candidate: &PyZenEngine, baseline: &PyZenEngine) -> Self {
Self {
inner: Arc::new(ImpactAnalysis::new(
candidate.engine.clone(),
baseline.engine.clone(),
)),
}
}
/// One entry point for both feeds: a list of JSON documents (bytes or
/// str), or any object implementing the Arrow PyCapsule protocol
/// (`__arrow_c_stream__` — a pyarrow RecordBatch, Table, polars frame…).
/// Inputs are parsed once and shared by both arms; dotted top-level keys
/// and dotted column names compose into nested objects. Returns the
/// serialized batch — `{rows, summary}`, with `before`/`after` present
/// only for changed or failed records and summary counts that merge
/// additively across batches. The GIL is released for the batch.
#[pyo3(signature = (candidate_key, baseline_key, data, max_depth=None))]
pub fn run_batch(
&self,
py: Python,
candidate_key: String,
baseline_key: String,
data: pyo3::Bound<'_, pyo3::PyAny>,
max_depth: Option<u8>,
) -> PyResult<Py<PyBytes>> {
use pyo3::prelude::PyAnyMethods;
enum Feed {
Lines(Vec<Vec<u8>>),
#[cfg(feature = "arrow")]
Batches(Vec<arrow_array::RecordBatch>),
}
let feed = if let Ok(list) = data.downcast::<pyo3::types::PyList>() {
use pyo3::prelude::PyListMethods;
let mut lines = Vec::with_capacity(list.len());
for item in list.iter() {
let payload = item
.extract::<Vec<u8>>()
.or_else(|_| item.extract::<String>().map(String::into_bytes))
.map_err(|_| {
pyo3::exceptions::PyTypeError::new_err(
"list items must be JSON documents as bytes or str",
)
})?;
lines.push(payload);
}
Feed::Lines(lines)
} else if data.hasattr("__arrow_c_stream__")? {
#[cfg(not(feature = "arrow"))]
{
return Err(pyo3::exceptions::PyTypeError::new_err(
"Arrow input requires a build with the 'arrow' feature",
));
}
#[cfg(feature = "arrow")]
{
use pyo3::types::PyCapsuleMethods;
let capsule_any = data.call_method0("__arrow_c_stream__")?;
let capsule = capsule_any.downcast::<pyo3::types::PyCapsule>()?;
let pointer =
capsule.pointer() as *mut arrow_array::ffi_stream::FFI_ArrowArrayStream;
let reader =
unsafe { arrow_array::ffi_stream::ArrowArrayStreamReader::from_raw(pointer) }
.map_err(|e| anyhow!("arrow stream: {e}"))?;
let mut batches = Vec::new();
for batch in reader {
batches.push(batch.map_err(|e| anyhow!("arrow batch: {e}"))?);
}
Feed::Batches(batches)
}
} else {
return Err(pyo3::exceptions::PyTypeError::new_err(
"expected a list of JSON documents or an Arrow stream object",
));
};
let analysis = self.inner.clone();
let options = EvaluationOptions {
trace: false,
max_depth: max_depth.unwrap_or(10),
};
let body = py.allow_threads(move || {
block_on(
worker_pool().spawn_pinned(move || async move {
let variables = match feed {
Feed::Lines(lines) => {
let mut variables = Vec::with_capacity(lines.len());
for payload in &lines {
let parsed =
serde_json::from_slice::<zen_engine::Variable>(payload)
.map_err(|e| anyhow!("input is not JSON: {e}"))?;
variables.push(normalize_dotted(parsed));
}
variables
}
#[cfg(feature = "arrow")]
Feed::Batches(batches) => {
let mut variables = Vec::new();
for batch in &batches {
variables.extend(crate::columnar::inputs_from_batch(batch)?);
}
variables
}
};
let comparison = analysis
.compare(&candidate_key, &baseline_key)
.await
.map_err(|e| anyhow!(e.to_string()))?;
let batch = comparison.run_batch(variables, options).await;
serde_json::to_vec(&batch).map_err(|e| anyhow!(e))
}),
)
.map_err(|_| anyhow!("evaluation worker panicked"))?
})?;
Ok(PyBytes::new(py, &body).into())
}
}
+1 -1
View File
@@ -22,7 +22,7 @@ use zen_engine::{DecisionEngine, EvaluationOptions};
#[pyclass]
#[pyo3(name = "ZenEngine")]
pub struct PyZenEngine {
engine: Arc<DecisionEngine>,
pub(crate) engine: Arc<DecisionEngine>,
}
#[derive(Serialize, Deserialize)]
+6
View File
@@ -9,6 +9,10 @@ use pyo3::prelude::PyModuleMethods;
use pyo3::types::PyModule;
use pyo3::{pymodule, wrap_pyfunction, Bound, PyResult, Python};
#[cfg(feature = "arrow")]
mod columnar;
#[cfg(feature = "data")]
mod data;
mod content;
mod convert;
mod custom_node;
@@ -27,6 +31,8 @@ fn zen(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<PyZenDecision>()?;
m.add_class::<PyExpression>()?;
m.add_class::<PyZenDecisionContent>()?;
#[cfg(feature = "data")]
m.add_class::<data::PyZenImpactAnalysis>()?;
m.add_function(wrap_pyfunction!(evaluate_expression, m)?)?;
m.add_function(wrap_pyfunction!(evaluate_unary_expression, m)?)?;
m.add_function(wrap_pyfunction!(render_template, m)?)?;