use std::ffi::{CStr, CString, c_char, c_void};
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::ptr;
use onnx_runtime_session::{InferenceSession, SessionBuilder, SessionError, Tensor};
#[repr(C)]
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum OrtErrorCode {
Ok = 0,
Fail = 1,
InvalidArgument = 2,
NoSuchFile = 3,
NoModel = 4,
EngineMismatch = 5,
InvalidProtobuf = 6,
ModelLoaded = 7,
NotImplemented = 8,
InvalidGraph = 10,
EpFail = 11,
}
const VERSION: &[u8] = concat!(env!("CARGO_PKG_VERSION"), "\0").as_bytes();
#[unsafe(no_mangle)]
pub extern "C" fn nxrt_get_version_string() -> *const c_char {
VERSION.as_ptr() as *const c_char
}
pub struct OrtStatus {
code: OrtErrorCode,
message: CString,
}
impl OrtStatus {
fn boxed(code: OrtErrorCode, message: impl Into<Vec<u8>>) -> *mut OrtStatus {
let message = CString::new(message).unwrap_or_else(|_| {
CString::new("status message contained an interior NUL").expect("literal is NUL-free")
});
Box::into_raw(Box::new(OrtStatus { code, message }))
}
}
type FfiResult = Result<(), (OrtErrorCode, String)>;
fn guard<F: FnOnce() -> FfiResult>(f: F) -> *mut OrtStatus {
match catch_unwind(AssertUnwindSafe(f)) {
Ok(Ok(())) => ptr::null_mut(),
Ok(Err((code, msg))) => OrtStatus::boxed(code, msg),
Err(_) => OrtStatus::boxed(
OrtErrorCode::Fail,
"panic caught at the FFI boundary (this is a runtime bug)",
),
}
}
fn map_session_error(err: &SessionError) -> OrtErrorCode {
use onnx_runtime_session::SessionError as E;
match err {
E::InputNotFound { .. }
| E::DtypeMismatch { .. }
| E::ShapeMismatch { .. }
| E::UnknownOption { .. }
| E::InvalidOption { .. }
| E::DynamicShape { .. }
| E::SymbolConflict { .. }
| E::RankMismatch { .. }
| E::RuntimeBroadcastIncompatible { .. } => OrtErrorCode::InvalidArgument,
E::NoModelSource => OrtErrorCode::NoModel,
E::UnsupportedOp { .. } => OrtErrorCode::NotImplemented,
E::Ep(_) | E::ExecutionProviderUnavailable(_) => OrtErrorCode::EpFail,
E::DanglingEpContext { .. }
| E::Ir(_)
| E::Graph(_)
| E::Optimize(_)
| E::ShapeInfer(_) => OrtErrorCode::InvalidGraph,
E::Load(load) => map_loader_error(load),
E::NotInitialized
| E::Internal(_)
| E::UnresolvedShape { .. }
| E::ShapeOverflow { .. }
| E::OutputShapeCountMismatch { .. }
| E::SequenceOp { .. }
| E::ControlFlow { .. }
| E::HeterogeneousPlacementRequired { .. } => OrtErrorCode::Fail,
}
}
fn map_loader_error(err: &onnx_runtime_loader::LoaderError) -> OrtErrorCode {
use onnx_runtime_loader::LoaderError as L;
match err {
L::Io { source, .. } if source.kind() == std::io::ErrorKind::NotFound => {
OrtErrorCode::NoSuchFile
}
L::ExternalDataNotFound { .. } => OrtErrorCode::NoSuchFile,
L::ProtobufParse(_) => OrtErrorCode::InvalidProtobuf,
L::ExternalDataPath { .. } | L::EpContextPath { .. } | L::Ir(_) | L::GraphBuild(_) => {
OrtErrorCode::InvalidGraph
}
_ => OrtErrorCode::Fail,
}
}
fn session_err(err: SessionError) -> (OrtErrorCode, String) {
(map_session_error(&err), err.to_string())
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_get_error_code(status: *const OrtStatus) -> OrtErrorCode {
if status.is_null() {
return OrtErrorCode::Ok;
}
unsafe { (*status).code }
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_get_error_message(status: *const OrtStatus) -> *const c_char {
if status.is_null() {
return c"".as_ptr();
}
unsafe { (*status).message.as_ptr() }
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_release_status(status: *mut OrtStatus) {
if status.is_null() {
return;
}
drop(unsafe { Box::from_raw(status) });
}
pub struct OrtSession {
inner: InferenceSession,
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_create_session(
model_path: *const c_char,
out: *mut *mut OrtSession,
) -> *mut OrtStatus {
guard(|| {
if out.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out pointer is null".into()));
}
unsafe { *out = ptr::null_mut() };
if model_path.is_null() {
return Err((OrtErrorCode::InvalidArgument, "model_path is null".into()));
}
let path = unsafe { CStr::from_ptr(model_path) }
.to_str()
.map_err(|_| {
(
OrtErrorCode::InvalidArgument,
"model_path is not valid UTF-8".into(),
)
})?;
let inner = InferenceSession::load(path).map_err(session_err)?;
let handle = Box::into_raw(Box::new(OrtSession { inner }));
unsafe { *out = handle };
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_release_session(session: *mut OrtSession) {
if session.is_null() {
return;
}
drop(unsafe { Box::from_raw(session) });
}
#[derive(Default)]
pub struct OrtSessionOptions {
entries: Vec<(String, String)>,
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_create_session_options(
out: *mut *mut OrtSessionOptions,
) -> *mut OrtStatus {
guard(|| {
if out.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out pointer is null".into()));
}
unsafe { *out = Box::into_raw(Box::new(OrtSessionOptions::default())) };
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_release_session_options(options: *mut OrtSessionOptions) {
if options.is_null() {
return;
}
drop(unsafe { Box::from_raw(options) });
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_add_session_config_entry(
options: *mut OrtSessionOptions,
key: *const c_char,
value: *const c_char,
) -> *mut OrtStatus {
guard(|| {
if options.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
"options handle is null".into(),
));
}
if key.is_null() || value.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
"key or value pointer is null".into(),
));
}
let options = unsafe { &mut *options };
let key = unsafe { CStr::from_ptr(key) }.to_str().map_err(|_| {
(
OrtErrorCode::InvalidArgument,
"key is not valid UTF-8".into(),
)
})?;
let value = unsafe { CStr::from_ptr(value) }.to_str().map_err(|_| {
(
OrtErrorCode::InvalidArgument,
"value is not valid UTF-8".into(),
)
})?;
options.entries.push((key.to_string(), value.to_string()));
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_create_session_with_options(
model_path: *const c_char,
options: *const OrtSessionOptions,
out: *mut *mut OrtSession,
) -> *mut OrtStatus {
guard(|| {
if out.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out pointer is null".into()));
}
unsafe { *out = ptr::null_mut() };
if model_path.is_null() {
return Err((OrtErrorCode::InvalidArgument, "model_path is null".into()));
}
let path = unsafe { CStr::from_ptr(model_path) }
.to_str()
.map_err(|_| {
(
OrtErrorCode::InvalidArgument,
"model_path is not valid UTF-8".into(),
)
})?;
let mut builder = SessionBuilder::new().model(path);
if !options.is_null() {
for (key, value) in &unsafe { &*options }.entries {
builder = builder.option(key, value);
}
}
let inner = builder.build().map_err(session_err)?;
let handle = Box::into_raw(Box::new(OrtSession { inner }));
unsafe { *out = handle };
Ok(())
})
}
pub struct OrtValue {
inner: Tensor,
}
fn shape_numel(shape: &[i64]) -> Result<usize, (OrtErrorCode, String)> {
let mut numel: usize = 1;
for (i, &d) in shape.iter().enumerate() {
if d < 0 {
return Err((
OrtErrorCode::InvalidArgument,
format!("shape dim {i} is negative ({d})"),
));
}
numel = numel.checked_mul(d as usize).ok_or((
OrtErrorCode::InvalidArgument,
"shape element count overflows usize".into(),
))?;
}
Ok(numel)
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_create_tensor(
data: *const c_void,
data_len: usize,
shape: *const i64,
rank: usize,
data_type: i32,
out: *mut *mut OrtValue,
) -> *mut OrtStatus {
guard(|| {
if out.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out pointer is null".into()));
}
unsafe { *out = ptr::null_mut() };
let dtype = onnx_runtime_ir::DataType::from_onnx(data_type).ok_or((
OrtErrorCode::InvalidArgument,
format!("unsupported ONNX data_type {data_type}"),
))?;
if shape.is_null() && rank != 0 {
return Err((
OrtErrorCode::InvalidArgument,
"shape is null but rank is non-zero".into(),
));
}
let dims: Vec<i64> = if rank == 0 {
Vec::new()
} else {
unsafe { std::slice::from_raw_parts(shape, rank) }.to_vec()
};
let numel = shape_numel(&dims)?;
let expected = dtype.storage_bytes(numel);
if data_len != expected {
return Err((
OrtErrorCode::InvalidArgument,
format!(
"data_len {data_len} does not match {expected} bytes for dtype {dtype:?} shape {dims:?}"
),
));
}
if data.is_null() && data_len != 0 {
return Err((
OrtErrorCode::InvalidArgument,
"data is null but data_len is non-zero".into(),
));
}
let bytes: &[u8] = if data_len == 0 {
&[]
} else {
unsafe { std::slice::from_raw_parts(data as *const u8, data_len) }
};
let shape_usize: Vec<usize> = dims.iter().map(|&d| d as usize).collect();
let tensor = Tensor::from_raw(dtype, shape_usize, bytes).map_err(session_err)?;
let handle = Box::into_raw(Box::new(OrtValue { inner: tensor }));
unsafe { *out = handle };
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_release_value(value: *mut OrtValue) {
if value.is_null() {
return;
}
drop(unsafe { Box::from_raw(value) });
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_get_tensor_dtype(
value: *const OrtValue,
out: *mut i32,
) -> *mut OrtStatus {
guard(|| {
let value = unsafe { value.as_ref() }
.ok_or((OrtErrorCode::InvalidArgument, "value handle is null".into()))?;
if out.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out pointer is null".into()));
}
unsafe { *out = value.inner.dtype.to_onnx() };
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_get_tensor_rank(
value: *const OrtValue,
out: *mut usize,
) -> *mut OrtStatus {
guard(|| {
let value = unsafe { value.as_ref() }
.ok_or((OrtErrorCode::InvalidArgument, "value handle is null".into()))?;
if out.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out pointer is null".into()));
}
unsafe { *out = value.inner.shape.len() };
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_get_tensor_shape(
value: *const OrtValue,
out_dims: *mut i64,
rank: usize,
) -> *mut OrtStatus {
guard(|| {
let value = unsafe { value.as_ref() }
.ok_or((OrtErrorCode::InvalidArgument, "value handle is null".into()))?;
let shape = &value.inner.shape;
if rank != shape.len() {
return Err((
OrtErrorCode::InvalidArgument,
format!("rank {rank} does not match tensor rank {}", shape.len()),
));
}
if rank == 0 {
return Ok(());
}
if out_dims.is_null() {
return Err((OrtErrorCode::InvalidArgument, "out_dims is null".into()));
}
let dst = unsafe { std::slice::from_raw_parts_mut(out_dims, rank) };
for (slot, &d) in dst.iter_mut().zip(shape.iter()) {
*slot = d as i64;
}
Ok(())
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn nxrt_get_tensor_data(
value: *const OrtValue,
out_data: *mut *const c_void,
out_len: *mut usize,
) -> *mut OrtStatus {
guard(|| {
let value = unsafe { value.as_ref() }
.ok_or((OrtErrorCode::InvalidArgument, "value handle is null".into()))?;
if out_data.is_null() || out_len.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
"out_data or out_len is null".into(),
));
}
let bytes = value.inner.as_bytes();
unsafe {
*out_data = bytes.as_ptr() as *const c_void;
*out_len = bytes.len();
}
Ok(())
})
}
#[unsafe(no_mangle)]
#[allow(clippy::too_many_arguments)]
pub unsafe extern "C" fn nxrt_run(
session: *mut OrtSession,
input_names: *const *const c_char,
input_values: *const *const OrtValue,
n_inputs: usize,
output_names: *const *const c_char,
n_outputs: usize,
out_values: *mut *mut OrtValue,
) -> *mut OrtStatus {
guard(|| {
let session = unsafe { session.as_mut() }.ok_or((
OrtErrorCode::InvalidArgument,
"session handle is null".into(),
))?;
if n_outputs != 0 {
if out_values.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
"out_values is null but n_outputs is non-zero".into(),
));
}
if output_names.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
"output_names is null but n_outputs is non-zero".into(),
));
}
let slots = unsafe { std::slice::from_raw_parts_mut(out_values, n_outputs) };
for slot in slots.iter_mut() {
*slot = ptr::null_mut();
}
}
if n_inputs != 0 && (input_names.is_null() || input_values.is_null()) {
return Err((
OrtErrorCode::InvalidArgument,
"input_names or input_values is null but n_inputs is non-zero".into(),
));
}
let mut inputs: Vec<(&str, &Tensor)> = Vec::with_capacity(n_inputs);
for i in 0..n_inputs {
let name_ptr = unsafe { *input_names.add(i) };
let val_ptr = unsafe { *input_values.add(i) };
if name_ptr.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
format!("input name #{i} is null"),
));
}
let name = unsafe { CStr::from_ptr(name_ptr) }.to_str().map_err(|_| {
(
OrtErrorCode::InvalidArgument,
format!("input name #{i} is not valid UTF-8"),
)
})?;
let value = unsafe { val_ptr.as_ref() }.ok_or((
OrtErrorCode::InvalidArgument,
format!("input value #{i} is null"),
))?;
inputs.push((name, &value.inner));
}
let declared: Vec<&str> = session
.inner
.outputs()
.iter()
.map(|m| m.name.as_str())
.collect();
let mut want: Vec<usize> = Vec::with_capacity(n_outputs);
for i in 0..n_outputs {
let name_ptr = unsafe { *output_names.add(i) };
if name_ptr.is_null() {
return Err((
OrtErrorCode::InvalidArgument,
format!("output name #{i} is null"),
));
}
let name = unsafe { CStr::from_ptr(name_ptr) }.to_str().map_err(|_| {
(
OrtErrorCode::InvalidArgument,
format!("output name #{i} is not valid UTF-8"),
)
})?;
let pos = declared.iter().position(|d| *d == name).ok_or((
OrtErrorCode::InvalidArgument,
format!("requested output {name:?} is not a model output"),
))?;
want.push(pos);
}
let outputs = session.inner.run(&inputs).map_err(session_err)?;
let mut outputs: Vec<Option<Tensor>> = outputs.into_iter().map(Some).collect();
let mut produced: Vec<*mut OrtValue> = Vec::with_capacity(n_outputs);
for &pos in &want {
let tensor = match outputs.get_mut(pos).and_then(Option::take) {
Some(t) => t,
None => {
for p in produced {
drop(unsafe { Box::from_raw(p) });
}
return Err((
OrtErrorCode::Fail,
format!("output index {pos} requested more than once"),
));
}
};
produced.push(Box::into_raw(Box::new(OrtValue { inner: tensor })));
}
for (i, handle) in produced.into_iter().enumerate() {
unsafe { *out_values.add(i) = handle };
}
Ok(())
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn version_is_nul_terminated() {
assert_eq!(*VERSION.last().unwrap(), 0);
let ptr = nxrt_get_version_string();
assert!(!ptr.is_null());
}
#[test]
fn status_codes_match_ort() {
assert_eq!(OrtErrorCode::Ok as i32, 0);
assert_eq!(OrtErrorCode::InvalidGraph as i32, 10);
assert_eq!(OrtErrorCode::EpFail as i32, 11);
}
#[test]
fn null_status_accessors_are_ok() {
assert_eq!(
unsafe { nxrt_get_error_code(ptr::null()) },
OrtErrorCode::Ok
);
let msg = unsafe { nxrt_get_error_message(ptr::null()) };
assert!(!msg.is_null());
unsafe { nxrt_release_status(ptr::null_mut()) };
}
#[test]
fn session_error_mapping_covers_structural_and_shape_variants() {
use onnx_runtime_session::SessionError as E;
assert_eq!(
map_session_error(&E::DanglingEpContext {
source_key: Some("QNN".into()),
partition_name: Some("encoder".into()),
}),
OrtErrorCode::InvalidGraph
);
assert_eq!(
map_session_error(&E::SymbolConflict {
symbol: "N".into(),
first: 2,
second: 3,
}),
OrtErrorCode::InvalidArgument
);
assert_eq!(
map_session_error(&E::RankMismatch {
name: "x".into(),
expected: 2,
got: 3,
}),
OrtErrorCode::InvalidArgument
);
assert_eq!(
map_session_error(&E::RuntimeBroadcastIncompatible {
node: "add".into(),
domain: String::new(),
op_type: "Add".into(),
input_shapes: vec![vec![2, 3], vec![2, 4]],
}),
OrtErrorCode::InvalidArgument
);
assert_eq!(
map_session_error(&E::UnresolvedShape {
value: "y".into(),
op: "Reshape".into(),
}),
OrtErrorCode::Fail
);
assert_eq!(
map_session_error(&E::ShapeOverflow {
value: "y".into(),
dims: vec![usize::MAX, 2],
}),
OrtErrorCode::Fail
);
assert_eq!(
map_session_error(&E::OutputShapeCountMismatch {
op: "NonZero".into(),
expected: 1,
got: 2,
}),
OrtErrorCode::Fail
);
}
}