use std::ffi::CString;
use std::path::Path;
use std::sync::Arc;
use libloading::{Library, Symbol};
use crate::error::{Error, Result};
use crate::ffi::*;
pub const DEFAULT_LIBRARY_PATHS: &[&str] = &[
"libtensorflowlite_c.so",
"libtensorflow-lite.so",
"libtensorflowlite_c.dylib",
"libtensorflow-lite.dylib",
"tensorflowlite_c.dll",
"/usr/lib/libtensorflowlite_c.so",
"/usr/lib/libtensorflow-lite.so",
"/usr/local/lib/libtensorflowlite_c.so",
"/usr/local/lib/libtensorflow-lite.so",
"/usr/local/lib/libtensorflowlite_c.dylib",
"/opt/homebrew/lib/libtensorflowlite_c.dylib",
];
pub struct TfLiteLibrary {
_lib: Arc<Library>,
pub(crate) model_create_from_file: Symbol<'static, ModelCreateFromFileFn>,
pub(crate) model_delete: Symbol<'static, ModelDeleteFn>,
pub(crate) options_create: Symbol<'static, OptionsCreateFn>,
pub(crate) options_delete: Symbol<'static, OptionsDeleteFn>,
pub(crate) options_set_num_threads: Symbol<'static, OptionsSetNumThreadsFn>,
pub(crate) options_add_delegate: Symbol<'static, OptionsAddDelegateFn>,
pub(crate) interpreter_create: Symbol<'static, InterpreterCreateFn>,
pub(crate) interpreter_delete: Symbol<'static, InterpreterDeleteFn>,
pub(crate) interpreter_allocate_tensors: Symbol<'static, InterpreterAllocateTensorsFn>,
pub(crate) interpreter_invoke: Symbol<'static, InterpreterInvokeFn>,
pub(crate) interpreter_get_input_tensor_count: Symbol<'static, InterpreterGetTensorCountFn>,
pub(crate) interpreter_get_output_tensor_count: Symbol<'static, InterpreterGetTensorCountFn>,
pub(crate) interpreter_get_input_tensor: Symbol<'static, InterpreterGetInputTensorFn>,
pub(crate) interpreter_get_output_tensor: Symbol<'static, InterpreterGetOutputTensorFn>,
pub(crate) tensor_type: Symbol<'static, TensorTypeFn>,
pub(crate) tensor_num_dims: Symbol<'static, TensorNumDimsFn>,
pub(crate) tensor_dim: Symbol<'static, TensorDimFn>,
pub(crate) tensor_byte_size: Symbol<'static, TensorByteSizeFn>,
pub(crate) tensor_data: Symbol<'static, TensorDataFn>,
pub(crate) tensor_name: Symbol<'static, TensorNameFn>,
pub(crate) tensor_quantization_params: Symbol<'static, TensorQuantizationParamsFn>,
pub(crate) external_delegate: Option<ExternalDelegateSymbols>,
}
pub(crate) struct ExternalDelegateSymbols {
pub options_create: Symbol<'static, ExternalDelegateOptionsCreateFn>,
pub options_delete: Symbol<'static, ExternalDelegateOptionsDeleteFn>,
pub options_set_library_path: Symbol<'static, ExternalDelegateOptionsSetLibraryPathFn>,
pub options_insert: Symbol<'static, ExternalDelegateOptionsInsertFn>,
pub create: Symbol<'static, ExternalDelegateCreateFn>,
pub delete: Symbol<'static, ExternalDelegateDeleteFn>,
}
unsafe impl Send for TfLiteLibrary {}
unsafe impl Sync for TfLiteLibrary {}
impl TfLiteLibrary {
pub fn load_default() -> Result<Arc<Self>> {
let mut last_err: Option<libloading::Error> = None;
for candidate in DEFAULT_LIBRARY_PATHS {
match Self::try_load(candidate) {
Ok(lib) => return Ok(lib),
Err(Error::LoadFailed { source, .. }) => last_err = Some(source),
Err(other) => return Err(other),
}
}
Err(Error::LoadFailed {
source: last_err.unwrap_or_else(|| {
libloading::Error::DlOpenUnknown
}),
paths: DEFAULT_LIBRARY_PATHS.iter().map(|s| s.to_string()).collect(),
})
}
pub fn load_from_path<P: AsRef<Path>>(path: P) -> Result<Arc<Self>> {
let path_ref = path.as_ref();
Self::try_load(path_ref).map_err(|err| match err {
Error::LoadFailed { source, .. } => Error::LoadFailed {
source,
paths: vec![path_ref.display().to_string()],
},
other => other,
})
}
fn try_load<P: AsRef<Path>>(path: P) -> Result<Arc<Self>> {
let path_ref = path.as_ref();
let lib = unsafe { Library::new(path_ref) }.map_err(|source| Error::LoadFailed {
source,
paths: vec![path_ref.display().to_string()],
})?;
let lib = Arc::new(lib);
let leaked: &'static Arc<Library> = Box::leak(Box::new(Arc::clone(&lib)));
macro_rules! sym {
($name:literal) => {{
unsafe {
leaked.get(concat!($name, "\0").as_bytes()).map_err(|source| {
Error::SymbolNotFound { name: $name, source }
})?
}
}};
}
Ok(Arc::new(TfLiteLibrary {
_lib: lib,
model_create_from_file: sym!("TfLiteModelCreateFromFile"),
model_delete: sym!("TfLiteModelDelete"),
options_create: sym!("TfLiteInterpreterOptionsCreate"),
options_delete: sym!("TfLiteInterpreterOptionsDelete"),
options_set_num_threads: sym!("TfLiteInterpreterOptionsSetNumThreads"),
options_add_delegate: sym!("TfLiteInterpreterOptionsAddDelegate"),
interpreter_create: sym!("TfLiteInterpreterCreate"),
interpreter_delete: sym!("TfLiteInterpreterDelete"),
interpreter_allocate_tensors: sym!("TfLiteInterpreterAllocateTensors"),
interpreter_invoke: sym!("TfLiteInterpreterInvoke"),
interpreter_get_input_tensor_count: sym!("TfLiteInterpreterGetInputTensorCount"),
interpreter_get_output_tensor_count: sym!("TfLiteInterpreterGetOutputTensorCount"),
interpreter_get_input_tensor: sym!("TfLiteInterpreterGetInputTensor"),
interpreter_get_output_tensor: sym!("TfLiteInterpreterGetOutputTensor"),
tensor_type: sym!("TfLiteTensorType"),
tensor_num_dims: sym!("TfLiteTensorNumDims"),
tensor_dim: sym!("TfLiteTensorDim"),
tensor_byte_size: sym!("TfLiteTensorByteSize"),
tensor_data: sym!("TfLiteTensorData"),
tensor_name: sym!("TfLiteTensorName"),
tensor_quantization_params: sym!("TfLiteTensorQuantizationParams"),
external_delegate: Self::resolve_external_delegate(leaked),
}))
}
fn resolve_external_delegate(
leaked: &'static Arc<Library>,
) -> Option<ExternalDelegateSymbols> {
unsafe {
let options_create = leaked
.get::<ExternalDelegateOptionsCreateFn>(b"TfLiteExternalDelegateOptionsCreate\0")
.ok()?;
let options_delete = leaked
.get::<ExternalDelegateOptionsDeleteFn>(b"TfLiteExternalDelegateOptionsDelete\0")
.ok()?;
let options_set_library_path = leaked
.get::<ExternalDelegateOptionsSetLibraryPathFn>(
b"TfLiteExternalDelegateOptionsSetLibraryPath\0",
)
.ok()?;
let options_insert = leaked
.get::<ExternalDelegateOptionsInsertFn>(b"TfLiteExternalDelegateOptionsInsert\0")
.ok()?;
let create = leaked
.get::<ExternalDelegateCreateFn>(b"TfLiteExternalDelegateCreate\0")
.ok()?;
let delete = leaked
.get::<ExternalDelegateDeleteFn>(b"TfLiteExternalDelegateDelete\0")
.ok()?;
Some(ExternalDelegateSymbols {
options_create,
options_delete,
options_set_library_path,
options_insert,
create,
delete,
})
}
}
}
pub struct ExternalDelegate {
pub(crate) raw: *mut TfLiteDelegate,
pub(crate) lib: Arc<TfLiteLibrary>,
}
unsafe impl Send for ExternalDelegate {}
unsafe impl Sync for ExternalDelegate {}
impl ExternalDelegate {
pub fn builder<P: AsRef<Path>>(
lib: Arc<TfLiteLibrary>,
library_path: P,
) -> Result<ExternalDelegateBuilder> {
let path_str = library_path
.as_ref()
.to_str()
.ok_or_else(|| Error::InvalidPath(c_string_nul_error()))?;
let c_path = CString::new(path_str).map_err(Error::InvalidPath)?;
Ok(ExternalDelegateBuilder {
lib,
library_path: c_path,
options: Vec::new(),
})
}
pub fn as_ptr(&self) -> *mut TfLiteDelegate {
self.raw
}
}
impl Drop for ExternalDelegate {
fn drop(&mut self) {
if let Some(ext) = self.lib.external_delegate.as_ref() {
unsafe {
(ext.delete)(self.raw);
}
}
}
}
pub struct ExternalDelegateBuilder {
lib: Arc<TfLiteLibrary>,
library_path: CString,
options: Vec<(CString, CString)>,
}
impl ExternalDelegateBuilder {
pub fn with_option(mut self, key: &str, value: &str) -> Result<Self> {
let k = CString::new(key).map_err(Error::InvalidOption)?;
let v = CString::new(value).map_err(Error::InvalidOption)?;
self.options.push((k, v));
Ok(self)
}
pub fn build(self) -> Result<ExternalDelegate> {
let lib = self.lib;
let ext = lib
.external_delegate
.as_ref()
.ok_or(Error::ExternalDelegateApiUnavailable)?;
unsafe {
let opts = (ext.options_create)();
if opts.is_null() {
return Err(Error::DelegateCreateFailed);
}
let set_status = (ext.options_set_library_path)(opts, self.library_path.as_ptr());
if set_status != TfLiteStatus::Ok {
(ext.options_delete)(opts);
return Err(Error::DelegateOptionRejected { status: set_status });
}
for (k, v) in &self.options {
let status = (ext.options_insert)(opts, k.as_ptr(), v.as_ptr());
if status != TfLiteStatus::Ok {
(ext.options_delete)(opts);
return Err(Error::DelegateOptionRejected { status });
}
}
let raw = (ext.create)(opts);
(ext.options_delete)(opts);
if raw.is_null() {
return Err(Error::DelegateCreateFailed);
}
Ok(ExternalDelegate { raw, lib: lib.clone() })
}
}
}
fn c_string_nul_error() -> std::ffi::NulError {
CString::new(vec![0u8]).unwrap_err()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Error;
#[test]
fn default_paths_are_nonempty_and_named() {
let len = DEFAULT_LIBRARY_PATHS.len();
assert!(len > 0, "default path list collapsed to empty");
assert!(DEFAULT_LIBRARY_PATHS
.iter()
.any(|p| p.ends_with("libtensorflow-lite.so")));
assert!(DEFAULT_LIBRARY_PATHS
.iter()
.any(|p| p.ends_with("libtensorflowlite_c.so")));
}
#[test]
fn load_failed_error_displays_paths() {
let err = Error::LoadFailed {
source: libloading::Error::DlOpenUnknown,
paths: vec!["/nope/libtensorflow-lite.so".to_string()],
};
let rendered = format!("{err}");
assert!(rendered.contains("/nope/libtensorflow-lite.so"));
}
#[test]
fn symbol_not_found_error_includes_symbol_name() {
let err = Error::SymbolNotFound {
name: "TfLiteSomeFunction",
source: libloading::Error::DlSymUnknown,
};
assert!(format!("{err}").contains("TfLiteSomeFunction"));
}
}