mirror of
https://github.com/gorules/zen.git
synced 2026-10-07 00:02:19 +00:00
feat: impact analysis
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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)]
|
||||
|
||||
@@ -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)?)?;
|
||||
|
||||
Reference in New Issue
Block a user