diff --git a/bindings/c/src/engine.rs b/bindings/c/src/engine.rs index 42c49fdd..5db2ff04 100644 --- a/bindings/c/src/engine.rs +++ b/bindings/c/src/engine.rs @@ -156,9 +156,13 @@ pub extern "C" fn zen_engine_evaluate( let zen_engine = unsafe { &*(engine as *mut ZenEngine) }; + let Ok(var_context) = zen_engine::Variable::try_from_value(val_context) else { + return ZenResult::error(ZenError::JsonDeserializationFailed); + }; + let maybe_result = tokio_runtime().block_on(zen_engine.evaluate_with_opts( str_key, - val_context.into(), + var_context, options.into(), )); let result = match maybe_result { @@ -232,8 +236,10 @@ pub extern "C" fn zen_engine_evaluate_batch( Ok(value) => { let engine = decision_engine.clone(); BatchTask::Pending(pool.spawn_pinned(move || async move { + let value = zen_engine::Variable::try_from_value(value) + .map_err(|e| json!(format!("invalid context: {e}")))?; engine - .evaluate_with_opts(key, value.into(), eval_options) + .evaluate_with_opts(key, value, eval_options) .await .map(|response| serde_json::to_value(&response).unwrap_or(Value::Null)) .map_err(|e| { diff --git a/bindings/nodejs/src/decision.rs b/bindings/nodejs/src/decision.rs index 7e16a198..3b4df2f9 100644 --- a/bindings/nodejs/src/decision.rs +++ b/bindings/nodejs/src/decision.rs @@ -34,9 +34,10 @@ impl ZenDecision { let options = opts.unwrap_or_default(); async move { - decision - .evaluate_serialized(context.into(), options.into()) - .await + let context = zen_engine::Variable::try_from_value(context).map_err( + |e| serde_json::json!({ "type": "ContextError", "source": e.to_string() }), + )?; + decision.evaluate_serialized(context, options.into()).await } }) .await diff --git a/bindings/nodejs/src/engine.rs b/bindings/nodejs/src/engine.rs index 48550bfd..7848937b 100644 --- a/bindings/nodejs/src/engine.rs +++ b/bindings/nodejs/src/engine.rs @@ -264,8 +264,11 @@ impl ZenEngine { }; async move { + let context = zen_engine::Variable::try_from_value(context).map_err( + |e| serde_json::json!({ "type": "ContextError", "source": e.to_string() }), + )?; graph - .evaluate_with_opts(key, context.into(), options) + .evaluate_with_opts(key, context, options) .await .map(|response| NodeEvalResponse::build(response, mode)) .map_err(|e| { @@ -357,8 +360,11 @@ impl ZenEngine { max_depth, }; + let context = zen_engine::Variable::try_from_value(context).map_err( + |e| serde_json::json!({ "type": "ContextError", "source": e.to_string() }), + )?; engine - .evaluate_with_opts(key, context.into(), eval_opts) + .evaluate_with_opts(key, context, eval_opts) .await .map(|response| NodeEvalResponse::build(response, mode)) .map_err(|e| { diff --git a/bindings/nodejs/src/expression.rs b/bindings/nodejs/src/expression.rs index 719ae93c..0767a9b2 100644 --- a/bindings/nodejs/src/expression.rs +++ b/bindings/nodejs/src/expression.rs @@ -4,10 +4,11 @@ use serde_json::Value; #[napi] pub fn evaluate_expression_sync(expression: String, context: Option) -> napi::Result { - let ctx = context.unwrap_or(Value::Null); + let ctx = zen_expression::Variable::try_from_value(context.unwrap_or(Value::Null)) + .map_err(|e| anyhow!(e))?; Ok( - zen_expression::evaluate_expression(expression.as_str(), ctx.into()) + zen_expression::evaluate_expression(expression.as_str(), ctx) .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))? .to_value(), ) @@ -16,8 +17,9 @@ pub fn evaluate_expression_sync(expression: String, context: Option) -> n #[allow(dead_code)] #[napi] pub fn evaluate_unary_expression_sync(expression: String, context: Value) -> napi::Result { + let context = zen_expression::Variable::try_from_value(context).map_err(|e| anyhow!(e))?; Ok( - zen_expression::evaluate_unary_expression(expression.as_str(), context.into()) + zen_expression::evaluate_unary_expression(expression.as_str(), context) .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))?, ) } @@ -25,7 +27,8 @@ pub fn evaluate_unary_expression_sync(expression: String, context: Value) -> nap #[allow(dead_code)] #[napi] pub fn render_template_sync(template: String, context: Value) -> napi::Result { - Ok(zen_tmpl::render(template.as_str(), context.into()) + let context = zen_expression::Variable::try_from_value(context).map_err(|e| anyhow!(e))?; + Ok(zen_tmpl::render(template.as_str(), context) .map_err(|e| anyhow!(serde_json::to_string(&e).unwrap_or_else(|_| e.to_string())))? .to_value()) } diff --git a/bindings/python/src/decision.rs b/bindings/python/src/decision.rs index f4a9ca53..2433c259 100644 --- a/bindings/python/src/decision.rs +++ b/bindings/python/src/decision.rs @@ -58,8 +58,10 @@ impl PyZenDecision { let result = tokio::future_into_py_with_locals(py, get_current_locals(py)?, async move { let response = worker_pool() .spawn_pinned(move || async move { + let context = + zen_engine::Variable::try_from_value(context).map_err(|e| anyhow!(e))?; decision - .evaluate_with_opts(context.into(), options.into()) + .evaluate_with_opts(context, options.into()) .await .map(crate::convert::PortableResponse::build) .map_err(|e| { diff --git a/bindings/python/src/engine.rs b/bindings/python/src/engine.rs index 473d0b03..081eb2d9 100644 --- a/bindings/python/src/engine.rs +++ b/bindings/python/src/engine.rs @@ -203,8 +203,10 @@ impl PyZenEngine { .map(|request| { let engine = self.engine.clone(); worker_pool().spawn_pinned(move || async move { + let context = zen_engine::Variable::try_from_value(request.context) + .map_err(|e| Value::String(format!("invalid context: {e}")))?; engine - .evaluate_with_opts(request.key, request.context.into(), options) + .evaluate_with_opts(request.key, context, options) .await .map(crate::convert::PortableResponse::build) .map_err(|e| { @@ -263,8 +265,10 @@ impl PyZenEngine { let result = tokio::future_into_py_with_locals(py, get_current_locals(py)?, async move { let response = worker_pool() .spawn_pinned(move || async move { + let context = + zen_engine::Variable::try_from_value(context).map_err(|e| anyhow!(e))?; engine - .evaluate_with_opts(key, context.into(), options.into()) + .evaluate_with_opts(key, context, options.into()) .await .map(crate::convert::PortableResponse::build) .map_err(|e| { diff --git a/bindings/uniffi/src/decision.rs b/bindings/uniffi/src/decision.rs index 8b6bae9e..75073003 100644 --- a/bindings/uniffi/src/decision.rs +++ b/bindings/uniffi/src/decision.rs @@ -36,8 +36,10 @@ impl ZenDecision { let response = task::spawn_blocking(move || { // The blocking code that uses non-Send types Handle::current().block_on(async move { + let context = zen_engine::Variable::try_from_value(context) + .map_err(|e| ZenError::ValidationError(e.to_string()))?; decision - .evaluate_with_opts(context.into(), options.into()) + .evaluate_with_opts(context, options.into()) .await .map(|response| ZenEngineResponse::try_from(response)) .map_err(|err| { diff --git a/bindings/uniffi/src/engine.rs b/bindings/uniffi/src/engine.rs index c6673c73..f6f8cb69 100644 --- a/bindings/uniffi/src/engine.rs +++ b/bindings/uniffi/src/engine.rs @@ -83,8 +83,10 @@ impl ZenEngine { let response = task::spawn_blocking(move || { // The blocking code that uses non-Send types Handle::current().block_on(async move { + let context = zen_engine::Variable::try_from_value(context) + .map_err(|e| ZenError::ValidationError(e.to_string()))?; engine - .evaluate_with_opts(key, context.into(), options.into()) + .evaluate_with_opts(key, context, options.into()) .await .map(|response| ZenEngineResponse::try_from(response)) .map_err(|err| { @@ -115,8 +117,10 @@ impl ZenEngine { task::spawn_blocking(move || { Handle::current().block_on(async move { let context: Value = request.context.try_into()?; + let context = zen_engine::Variable::try_from_value(context) + .map_err(|e| ZenError::ValidationError(e.to_string()))?; let response = engine - .evaluate_with_opts(request.key, context.into(), options) + .evaluate_with_opts(request.key, context, options) .await .map_err(|err| { ZenError::EvaluationError( diff --git a/core/expression/src/functions/internal.rs b/core/expression/src/functions/internal.rs index 97388b1a..47354cd9 100644 --- a/core/expression/src/functions/internal.rs +++ b/core/expression/src/functions/internal.rs @@ -661,7 +661,10 @@ pub(crate) mod imp { pub fn avg(args: Arguments) -> anyhow::Result { let a = __internal_number_array(&args, 0)?; - let sum = a.iter().fold(Decimal::ZERO, |acc, x| acc + x); + let sum = a + .iter() + .try_fold(Decimal::ZERO, |acc, x| acc.checked_add(*x)) + .context("Number overflow")?; Ok(V::Number(Decimal::from( sum.checked_div(Decimal::from(a.len())) @@ -671,7 +674,10 @@ pub(crate) mod imp { pub fn sum(args: Arguments) -> anyhow::Result { let a = __internal_number_array(&args, 0)?; - let sum = a.iter().fold(Decimal::ZERO, |acc, v| acc + v); + let sum = a + .iter() + .try_fold(Decimal::ZERO, |acc, v| acc.checked_add(*v)) + .context("Number overflow")?; Ok(V::Number(Decimal::from(sum))) } @@ -688,7 +694,10 @@ pub(crate) mod imp { let center_left = a.get(center - 1).context("Index out of bounds")?; let center_right = a.get(center).context("Index out of bounds")?; - let median = ((*center_left) + (*center_right)) / dec!(2); + let median = center_left + .checked_add(*center_right) + .context("Number overflow")? + / dec!(2); Ok(V::Number(median)) } } diff --git a/core/expression/src/vm/vm.rs b/core/expression/src/vm/vm.rs index d04e44bd..2e857d5a 100644 --- a/core/expression/src/vm/vm.rs +++ b/core/expression/src/vm/vm.rs @@ -467,7 +467,13 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { let a = self.pop()?; match (a, b) { - (Number(a), Number(b)) => self.push(Number(a + b)), + (Number(a), Number(b)) => { + let result = a.checked_add(b).ok_or_else(|| OpcodeErr { + opcode: "Add".into(), + message: "Number overflow".into(), + })?; + self.push(Number(result)); + } (String(a), String(b)) => { let mut c = StdString::with_capacity(a.len() + b.len()); @@ -489,7 +495,13 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { let a = self.pop()?; match (a, b) { - (Number(a), Number(b)) => self.push(Number(a - b)), + (Number(a), Number(b)) => { + let result = a.checked_sub(b).ok_or_else(|| OpcodeErr { + opcode: "Subtract".into(), + message: "Number overflow".into(), + })?; + self.push(Number(result)); + } _ => { return Err(OpcodeErr { opcode: "Subtract".into(), @@ -503,7 +515,13 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { let a = self.pop()?; match (a, b) { - (Number(a), Number(b)) => self.push(Number(a * b)), + (Number(a), Number(b)) => { + let result = a.checked_mul(b).ok_or_else(|| OpcodeErr { + opcode: "Multiply".into(), + message: "Number overflow".into(), + })?; + self.push(Number(result)); + } _ => { return Err(OpcodeErr { opcode: "Multiply".into(), @@ -535,7 +553,13 @@ impl<'arena, 'parent_ref, 'bytecode_ref> VMInner<'parent_ref, 'bytecode_ref> { let a = self.pop()?; match (a, b) { - (Number(a), Number(b)) => self.push(Number(a % b)), + (Number(a), Number(b)) => { + let result = match a.checked_rem(b) { + Some(r) => Number(r), + None => Null, + }; + self.push(result); + } _ => { return Err(OpcodeErr { opcode: "Modulo".into(), diff --git a/core/expression/tests/isolate.rs b/core/expression/tests/isolate.rs index 71a47419..93c7548c 100644 --- a/core/expression/tests/isolate.rs +++ b/core/expression/tests/isolate.rs @@ -966,3 +966,33 @@ mod test { } } } + +#[test] +fn arithmetic_overflow_errors_instead_of_panicking() { + let mut isolate = Isolate::new(); + let max = "79228162514264337593543950335"; + + assert!(isolate.run_standard(&format!("{max} + 1")).is_err()); + assert!(isolate.run_standard(&format!("-{max} - 1")).is_err()); + assert!(isolate.run_standard(&format!("{max} * 2")).is_err()); + assert!(isolate + .run_standard(&format!("sum([{max}, {max}])")) + .is_err()); + assert!(isolate + .run_standard(&format!("avg([{max}, {max}])")) + .is_err()); + assert!(isolate + .run_standard(&format!("median([{max}, {max}])")) + .is_err()); +} + +#[test] +fn division_and_modulo_by_zero_return_null() { + let mut isolate = Isolate::new(); + + let division = isolate.run_standard("1 / 0").unwrap(); + assert_eq!(division, Variable::Null); + + let modulo = isolate.run_standard("1 % 0").unwrap(); + assert_eq!(modulo, Variable::Null); +} diff --git a/core/types/src/variable/conv.rs b/core/types/src/variable/conv.rs index 0ee276b1..fd0f3846 100644 --- a/core/types/src/variable/conv.rs +++ b/core/types/src/variable/conv.rs @@ -1,47 +1,75 @@ use crate::rccell::RcCell; use crate::variable::Variable; use rust_decimal::Decimal; -#[cfg(not(feature = "arbitrary_precision"))] use rust_decimal::prelude::FromPrimitive; use serde_json::{Number, Value}; #[cfg(not(feature = "arbitrary_precision"))] use std::str::FromStr; +use thiserror::Error; -impl From for Variable { - fn from(value: Value) -> Self { - match value { - Value::Null => Variable::Null, - Value::Bool(b) => Variable::Bool(b), - Value::Number(n) => { - #[cfg(feature = "arbitrary_precision")] - { - Variable::Number( - Decimal::from_str_exact(n.as_str()) - .or_else(|_| Decimal::from_scientific(n.as_str())) - .expect("Allowed number"), - ) - } +#[derive(Debug, Error)] +pub enum VariableConversionError { + #[error("number out of range: {0}")] + NumberOutOfRange(String), +} - #[cfg(not(feature = "arbitrary_precision"))] - { - if let Some(n) = n.as_u64() { - return Variable::Number(n.into()); - } +impl Variable { + fn decimal_from_number(n: &Number) -> Option { + #[cfg(feature = "arbitrary_precision")] + { + Decimal::from_str_exact(n.as_str()) + .or_else(|_| Decimal::from_scientific(n.as_str())) + .ok() + .or_else(|| n.as_f64().and_then(Decimal::from_f64)) + } - if let Some(n) = n.as_i64() { - return Variable::Number(n.into()); - } + #[cfg(not(feature = "arbitrary_precision"))] + { + if let Some(u) = n.as_u64() { + return Some(u.into()); + } + if let Some(i) = n.as_i64() { + return Some(i.into()); + } + n.as_f64().and_then(Decimal::from_f64) + } + } +} - if let Some(n) = n.as_f64() { - return Variable::Number(Decimal::from_f64(n).expect("Allowed number")); - } +impl Variable { + pub fn try_from_value(value: Value) -> Result { + match value { + Value::Null => Ok(Variable::Null), + Value::Bool(b) => Ok(Variable::Bool(b)), + Value::Number(n) => Self::decimal_from_number(&n) + .map(Variable::Number) + .ok_or_else(|| VariableConversionError::NumberOutOfRange(n.to_string())), + Value::String(s) => Ok(Variable::String((s.as_str()).into())), + Value::Array(arr) => Ok(Variable::from_array( + arr.into_iter() + .map(Variable::try_from_value) + .collect::>()?, + )), + Value::Object(obj) => Ok(Variable::from_object( + obj.into_iter() + .map(|(k, v)| { + Ok(( + crate::symbol::Symbol::from(k.as_str()), + Variable::try_from_value(v)?, + )) + }) + .collect::>()?, + )), + } + } +} - unreachable!( - "serde_json::Number is always u64, i64, or f64 without arbitrary_precision" - ) - } - } - Value::String(s) => Variable::String((s.as_str()).into()), +impl From for Variable { + fn from(value: Value) -> Self { + match value { + Value::Number(n) => Variable::decimal_from_number(&n) + .map(Variable::Number) + .unwrap_or(Variable::Null), Value::Array(arr) => { Variable::from_array(arr.into_iter().map(Variable::from).collect()) } @@ -50,6 +78,7 @@ impl From for Variable { .map(|(k, v)| (crate::symbol::Symbol::from(k.as_str()), Variable::from(v))) .collect(), ), + other => Variable::from(&other), } } } @@ -59,35 +88,9 @@ impl From<&Value> for Variable { match value { Value::Null => Variable::Null, Value::Bool(b) => Variable::Bool(*b), - Value::Number(n) => { - #[cfg(feature = "arbitrary_precision")] - { - Variable::Number( - Decimal::from_str_exact(n.as_str()) - .or_else(|_| Decimal::from_scientific(n.as_str())) - .expect("Allowed number"), - ) - } - - #[cfg(not(feature = "arbitrary_precision"))] - { - if let Some(u) = n.as_u64() { - return Variable::Number(u.into()); - } - - if let Some(i) = n.as_i64() { - return Variable::Number(i.into()); - } - - if let Some(f) = n.as_f64() { - return Variable::Number(Decimal::from_f64(f).expect("Allowed number")); - } - - unreachable!( - "serde_json::Number is always u64, i64, or f64 without arbitrary_precision" - ); - } - } + Value::Number(n) => Variable::decimal_from_number(n) + .map(Variable::Number) + .unwrap_or(Variable::Null), Value::String(s) => Variable::String((s.as_str()).into()), Value::Array(arr) => Variable::from_array(arr.iter().map(Variable::from).collect()), Value::Object(obj) => Variable::from_object( @@ -99,6 +102,53 @@ impl From<&Value> for Variable { } } +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn out_of_range_number_converts_to_null() { + assert_eq!(Variable::from(json!(1e40)), Variable::Null); + assert_eq!(Variable::from(json!(-1e40)), Variable::Null); + assert_eq!( + Variable::from(json!({ "nested": [1e40] })), + Variable::from(json!({ "nested": [null] })) + ); + } + + #[test] + fn underflowing_number_rounds_to_zero() { + assert_eq!( + Variable::from(json!(1e-32)), + Variable::Number(Decimal::ZERO) + ); + } + + #[test] + fn try_from_value_rejects_out_of_range_numbers() { + assert!(matches!( + Variable::try_from_value(json!(1e40)), + Err(VariableConversionError::NumberOutOfRange(_)) + )); + assert!(matches!( + Variable::try_from_value(json!({ "a": { "b": [1, 2, 1e40] } })), + Err(VariableConversionError::NumberOutOfRange(_)) + )); + } + + #[test] + fn try_from_value_accepts_regular_payloads() { + let converted = + Variable::try_from_value(json!({ "a": [1, 2.5, -3], "b": "x", "c": null, "d": true })) + .unwrap(); + assert_eq!( + converted, + Variable::from(json!({ "a": [1, 2.5, -3], "b": "x", "c": null, "d": true })) + ); + } +} + impl From for Value { fn from(value: Variable) -> Self { match value { diff --git a/core/types/src/variable/de.rs b/core/types/src/variable/de.rs index 8d248274..71b83b5a 100644 --- a/core/types/src/variable/de.rs +++ b/core/types/src/variable/de.rs @@ -101,7 +101,9 @@ impl<'de> Visitor<'de> for VariableVisitor { return Ok(Variable::Number( Decimal::from_str_exact(str) .or_else(|_| Decimal::from_scientific(str)) - .map_err(|_| Error::custom("invalid number"))?, + .ok() + .or_else(|| str.parse::().ok().and_then(Decimal::from_f64)) + .ok_or_else(|| Error::custom(format!("number out of range: {str}")))?, )); } diff --git a/core/types/src/variable/mod.rs b/core/types/src/variable/mod.rs index e897fac7..e03f6112 100644 --- a/core/types/src/variable/mod.rs +++ b/core/types/src/variable/mod.rs @@ -11,6 +11,7 @@ use std::ops::Deref; use std::rc::Rc; use crate::rcvalue::RcValue; +pub use crate::variable::conv::VariableConversionError; pub use crate::variable::ref_deser::RefDeserializeError; use crate::variable::ref_deser::RefDeserializer; pub use de::VariableDeserializer;