use core::ffi::{CStr, c_char, c_int};
use core::marker::PhantomData;
use std::mem::MaybeUninit;
use executorch_sys as sys;
use crate::util::{ArrayRef, IntoRust, c_new, str2chars, try_c_new};
use crate::{Error, Result};
#[repr(transparent)]
pub struct BackendOption(sys::ET_BackendOption);
impl BackendOption {
pub fn new_bool(key: &str, value: bool) -> Result<Self> {
let key = ArrayRef::from_chars(str2chars(key));
unsafe { try_c_new(|out| sys::executorch_BackendOption_new_bool(key.0, value, out)) }
.map(Self)
}
pub fn new_int(key: &str, value: i64) -> Result<Self> {
let value: c_int = value.try_into().map_err(|_| Error::InvalidArgument)?;
let key = ArrayRef::from_chars(str2chars(key));
unsafe { try_c_new(|out| sys::executorch_BackendOption_new_int(key.0, value, out)) }
.map(Self)
}
pub fn new_str(key: &str, value: &str) -> Result<Self> {
let key = ArrayRef::from_chars(str2chars(key));
let value = ArrayRef::from_chars(str2chars(value));
unsafe { try_c_new(|out| sys::executorch_BackendOption_new_str(key.0, value.0, out)) }
.map(Self)
}
pub fn key(&self) -> &str {
let key = unsafe { CStr::from_ptr(sys::executorch_BackendOption_key(&self.0)) };
key.to_str().unwrap_or("")
}
pub fn is_bool(&self) -> bool {
unsafe { sys::executorch_BackendOption_is_bool(&self.0) }
}
pub fn is_int(&self) -> bool {
unsafe { sys::executorch_BackendOption_is_int(&self.0) }
}
pub fn is_str(&self) -> bool {
unsafe { sys::executorch_BackendOption_is_str(&self.0) }
}
pub fn as_bool(&self) -> Option<bool> {
unsafe { try_c_new(|out| sys::executorch_BackendOption_as_bool(&self.0, out)).ok() }
}
pub fn as_int(&self) -> Option<i64> {
unsafe { try_c_new(|out| sys::executorch_BackendOption_as_int(&self.0, out)).ok() }
.map(|v| v as i64)
}
pub fn as_str(&self) -> Option<&str> {
let ptr =
unsafe { try_c_new(|out| sys::executorch_BackendOption_as_str(&self.0, out)).ok()? };
Some(unsafe { CStr::from_ptr(ptr) }.to_str().unwrap_or(""))
}
}
#[repr(transparent)]
pub struct LoadBackendOptionsMap<'a>(sys::ET_LoadBackendOptionsMap, PhantomData<&'a ()>);
impl<'a> LoadBackendOptionsMap<'a> {
pub fn new() -> Self {
let inner = unsafe { c_new(|out| sys::executorch_LoadBackendOptionsMap_new(out)) };
Self(inner, PhantomData)
}
pub fn set_options(&mut self, backend_id: &str, options: &'a [BackendOption]) -> Result<()> {
let backend_id = ArrayRef::from_chars(str2chars(backend_id));
let ptr = options.as_ptr().cast::<sys::ET_BackendOption>();
unsafe {
sys::executorch_LoadBackendOptionsMap_set_options(
&mut self.0,
backend_id.0,
ptr,
options.len(),
)
}
.rs()
}
pub fn len(&self) -> usize {
unsafe { sys::executorch_LoadBackendOptionsMap_size(&self.0) }
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn get(&self, index: usize) -> Result<(&str, &[BackendOption])> {
let (backend_id, options, n_options) = unsafe {
let mut backend_id = MaybeUninit::<*const c_char>::uninit();
let mut options = MaybeUninit::<*const sys::ET_BackendOption>::uninit();
let mut n_options = MaybeUninit::<usize>::uninit();
sys::executorch_LoadBackendOptionsMap_entry_at(
&self.0,
index,
backend_id.as_mut_ptr(),
options.as_mut_ptr(),
n_options.as_mut_ptr(),
)
.rs()?;
(
backend_id.assume_init(),
options.assume_init(),
n_options.assume_init(),
)
};
let id = unsafe { CStr::from_ptr(backend_id) }
.to_str()
.map_err(|_| Error::InvalidString)?;
let options = options.cast::<BackendOption>();
let options = unsafe { core::slice::from_raw_parts(options, n_options) };
Ok((id, options))
}
pub(crate) fn as_cpp_ptr(&self) -> *const sys::ET_LoadBackendOptionsMap {
&self.0
}
}
impl Default for LoadBackendOptionsMap<'_> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backend_option_bool() {
let opt = BackendOption::new_bool("enable_profiling", true).unwrap();
assert_eq!(opt.key(), "enable_profiling");
assert!(opt.is_bool());
assert_eq!(opt.as_bool(), Some(true));
assert_eq!(opt.as_int(), None);
assert_eq!(opt.as_str(), None);
}
#[test]
fn backend_option_int() {
let opt = BackendOption::new_int("num_threads", 4).unwrap();
assert_eq!(opt.key(), "num_threads");
assert!(opt.is_int());
assert_eq!(opt.as_int(), Some(4));
assert_eq!(opt.as_bool(), None);
}
#[test]
fn backend_option_str() {
let opt = BackendOption::new_str("compute_unit", "cpu_and_gpu").unwrap();
assert!(opt.is_str());
assert_eq!(opt.as_str(), Some("cpu_and_gpu"));
assert_eq!(opt.as_int(), None);
}
#[test]
fn backend_option_key_too_long() {
let long_key = core::str::from_utf8(&[b'k'; 64]).unwrap();
assert!(BackendOption::new_bool(long_key, true).is_err());
}
#[test]
fn backend_option_int_out_of_range() {
assert!(BackendOption::new_int("x", i64::from(i32::MAX) + 1).is_err());
}
#[test]
fn options_map_set_get() {
let opts = [
BackendOption::new_int("num_threads", 4).unwrap(),
BackendOption::new_bool("enable_profiling", true).unwrap(),
];
let mut map = LoadBackendOptionsMap::new();
assert!(map.is_empty());
map.set_options("XnnpackBackend", &opts).unwrap();
assert_eq!(map.len(), 1);
assert!(!map.is_empty());
let (id, got) = map.get(0).unwrap();
assert_eq!(id, "XnnpackBackend");
assert_eq!(got.len(), 2);
assert_eq!(got[0].key(), "num_threads");
assert_eq!(got[0].as_int(), Some(4));
assert_eq!(got[1].as_bool(), Some(true));
assert!(map.get(1).is_err());
}
}