diff --git a/rust/src/lib.rs b/rust/src/lib.rs index 427808a..fbe0583 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -179,6 +179,7 @@ unsafe fn CreateEpFactories_impl( ); } + mlx::install_error_handler(); let factory = MlxEpFactory::new(registration_name, ort_api, ep_api); *factories.add(0) = factory.as_ptr(); *num_factories = 1; diff --git a/rust/src/mlx.rs b/rust/src/mlx.rs index 5d10fc3..0c5d72b 100644 --- a/rust/src/mlx.rs +++ b/rust/src/mlx.rs @@ -499,3 +499,34 @@ mod float64_primitive_tests { } } } + +unsafe extern "C" fn log_mlx_error( + msg: *const std::os::raw::c_char, + _data: *mut std::os::raw::c_void, +) { + let msg = unsafe { std::ffi::CStr::from_ptr(msg) }.to_string_lossy(); + log::error!("MLX error: {msg}"); +} + +pub fn install_error_handler() { + unsafe { mlx::mlx_set_error_handler(Some(log_mlx_error), std::ptr::null_mut(), None) }; +} + +#[cfg(test)] +mod error_handler_tests { + use super::*; + use crate::sys::mlx as sys; + + #[test] + fn op_failure_returns_error_code_instead_of_exiting() { + install_error_handler(); + let (a_data, b_data) = ([0f32; 2], [0f32; 3]); + let a = Array::from_data(a_data.as_ptr().cast(), &[2], sys::mlx_dtype__MLX_FLOAT32); + let b = Array::from_data(b_data.as_ptr().cast(), &[3], sys::mlx_dtype__MLX_FLOAT32); + let stream = Stream::new_default_cpu(); + let mut raw = unsafe { sys::mlx_array_new() }; + let rc = unsafe { sys::mlx_add(&mut raw, a.as_raw(), b.as_raw(), stream.as_raw()) }; + drop(Array::from_raw(raw)); + assert_ne!(rc, 0, "broadcasting [2] with [3] must fail"); + } +}