use std::os::raw::c_char;
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use super::api::{cstr_to_string, Api};
use super::ffi::*;
use crate::error::{FoundryLocalError, Result};
use crate::types::EpInfo;
static NATIVE_LIFECYCLE: Mutex<()> = Mutex::new(());
fn lock_lifecycle() -> std::sync::MutexGuard<'static, ()> {
NATIVE_LIFECYCLE
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub(crate) struct NativeManager {
api: Arc<Api>,
ptr: *mut flManager,
released: AtomicBool,
}
unsafe impl Send for NativeManager {}
unsafe impl Sync for NativeManager {}
impl NativeManager {
pub(crate) fn create(api: Arc<Api>, config: *const flConfiguration) -> Result<Self> {
let _lifecycle = lock_lifecycle();
let mut ptr: *mut flManager = std::ptr::null_mut();
let status = unsafe { (api.root().Manager_Create)(config, &mut ptr) };
api.check(status)?;
if ptr.is_null() {
return Err(FoundryLocalError::Internal {
reason: "Manager_Create returned a null manager".into(),
});
}
Ok(Self {
api,
ptr,
released: AtomicBool::new(false),
})
}
pub(crate) fn api(&self) -> Arc<Api> {
Arc::clone(&self.api)
}
pub(crate) fn catalog_ptr(&self) -> Result<*mut flCatalog> {
let mut catalog: *mut flCatalog = std::ptr::null_mut();
let status = unsafe { (self.api.root().Manager_GetCatalog)(self.ptr, &mut catalog) };
self.api.check(status)?;
if catalog.is_null() {
return Err(FoundryLocalError::Internal {
reason: "Manager_GetCatalog returned a null catalog".into(),
});
}
Ok(catalog)
}
pub(crate) fn web_service_start(&self) -> Result<()> {
let status = unsafe { (self.api.root().Manager_WebServiceStart)(self.ptr) };
self.api.check(status)
}
pub(crate) fn web_service_stop(&self) -> Result<()> {
let status = unsafe { (self.api.root().Manager_WebServiceStop)(self.ptr) };
self.api.check(status)
}
pub(crate) fn web_service_urls(&self) -> Result<Vec<String>> {
let mut urls: *const *const c_char = std::ptr::null();
let mut count: usize = 0;
let status =
unsafe { (self.api.root().Manager_WebServiceUrls)(self.ptr, &mut urls, &mut count) };
self.api.check(status)?;
if urls.is_null() || count == 0 {
return Ok(Vec::new());
}
let mut out = Vec::with_capacity(count);
for i in 0..count {
if let Some(s) = unsafe { cstr_to_string(*urls.add(i)) } {
out.push(s);
}
}
Ok(out)
}
pub(crate) fn discover_eps(&self) -> Result<Vec<EpInfo>> {
let mut eps: *const flEpInfo = std::ptr::null();
let mut count: usize = 0;
let status =
unsafe { (self.api.root().Manager_GetDiscoverableEps)(self.ptr, &mut eps, &mut count) };
self.api.check(status)?;
if eps.is_null() || count == 0 {
return Ok(Vec::new());
}
let mut out = Vec::with_capacity(count);
for i in 0..count {
let ep = unsafe { &*eps.add(i) };
out.push(EpInfo {
name: unsafe { cstr_to_string(ep.name) }.unwrap_or_default(),
is_registered: ep.is_registered,
});
}
Ok(out)
}
pub(crate) fn download_and_register_eps(
&self,
names: Option<&[&str]>,
progress: Option<EpProgressCallback>,
cancel_flag: Option<Arc<AtomicBool>>,
) -> Result<Option<String>> {
let name_cstrings: Option<Vec<std::ffi::CString>> = match names {
Some(ns) => Some(
ns.iter()
.map(|n| {
std::ffi::CString::new(*n).map_err(|_| FoundryLocalError::Validation {
reason: format!(
"execution provider name contains an interior NUL byte: {n:?}"
),
})
})
.collect::<Result<Vec<_>>>()?,
),
None => None,
};
let name_ptrs: Option<Vec<*const c_char>> = name_cstrings
.as_ref()
.map(|cs| cs.iter().map(|c| c.as_ptr()).collect());
let (names_ptr, names_len) = match &name_ptrs {
Some(p) if !p.is_empty() => (p.as_ptr(), p.len()),
_ => (std::ptr::null(), 0usize),
};
let mut ctx = EpCtx {
progress,
cancel_flag,
cancelled: false,
};
let user_data = &mut ctx as *mut EpCtx as *mut std::ffi::c_void;
let callback: flEpProgressCallback = Some(ep_trampoline);
let status = unsafe {
(self.api.root().Manager_DownloadAndRegisterEps)(
self.ptr, names_ptr, names_len, callback, user_data,
)
};
Ok(self.api.status_message(status))
}
pub(crate) fn shutdown(&self) -> Result<()> {
if self.released.load(Ordering::Acquire) {
return Ok(());
}
let status = unsafe { (self.api.root().Manager_Shutdown)(self.ptr) };
self.api.check(status)
}
pub(crate) fn teardown(&self) {
if self.released.swap(true, Ordering::AcqRel) {
return;
}
let _lifecycle = lock_lifecycle();
unsafe {
let status = (self.api.root().Manager_Shutdown)(self.ptr);
if !status.is_null() {
(self.api.root().Status_Release)(status);
}
(self.api.root().Manager_Release)(self.ptr);
}
}
}
impl Drop for NativeManager {
fn drop(&mut self) {
self.teardown();
}
}
pub(crate) type EpProgressCallback = Box<dyn FnMut(&str, f64) + Send>;
struct EpCtx {
progress: Option<EpProgressCallback>,
cancel_flag: Option<Arc<AtomicBool>>,
cancelled: bool,
}
unsafe extern "C" fn ep_trampoline(
ep_name: *const c_char,
value: f32,
user_data: *mut std::ffi::c_void,
) -> std::os::raw::c_int {
if user_data.is_null() {
return 0;
}
let result = catch_unwind(AssertUnwindSafe(|| {
let ctx = &mut *(user_data as *mut EpCtx);
if ctx
.cancel_flag
.as_ref()
.is_some_and(|f| f.load(Ordering::Relaxed))
{
ctx.cancelled = true;
return 1;
}
if let Some(cb) = ctx.progress.as_mut() {
let name = cstr_to_string(ep_name).unwrap_or_default();
cb(&name, value as f64);
}
0
}));
result.unwrap_or(1)
}