use std::sync::atomic::AtomicBool;
use std::sync::{Arc, Mutex, OnceLock, Weak};
use tokio::sync::Mutex as AsyncMutex;
use crate::catalog::Catalog;
use crate::configuration::{FoundryLocalConfig, Logger};
use crate::detail::api::Api;
use crate::detail::manager::{EpProgressCallback, NativeManager};
use crate::detail::task::spawn_blocking;
use crate::error::{FoundryLocalError, Result};
use crate::types::{EpDownloadResult, EpInfo};
#[derive(Default)]
struct SharedInstance {
outer: Weak<FoundryLocalManager>,
native: Weak<NativeManager>,
}
static INSTANCE: OnceLock<Mutex<SharedInstance>> = OnceLock::new();
fn instance_slot() -> &'static Mutex<SharedInstance> {
INSTANCE.get_or_init(|| Mutex::new(SharedInstance::default()))
}
pub struct FoundryLocalManager {
native: Arc<NativeManager>,
catalog: Catalog,
urls: Mutex<Vec<String>>,
web_service_lock: AsyncMutex<()>,
_logger: Option<Box<dyn Logger>>,
}
type EpDownloadProgressCallback = Box<dyn FnMut(&str, f64) + Send + 'static>;
pub struct EpDownloadBuilder<'a> {
manager: &'a FoundryLocalManager,
names: Option<Vec<String>>,
progress_callback: Option<EpDownloadProgressCallback>,
cancel_flag: Option<Arc<AtomicBool>>,
}
impl<'a> EpDownloadBuilder<'a> {
fn new(manager: &'a FoundryLocalManager) -> Self {
Self {
manager,
names: None,
progress_callback: None,
cancel_flag: None,
}
}
pub fn names<I, S>(mut self, names: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.names = Some(names.into_iter().map(Into::into).collect());
self
}
pub fn progress<F>(mut self, callback: F) -> Self
where
F: FnMut(&str, f64) + Send + 'static,
{
self.progress_callback = Some(Box::new(callback));
self
}
pub fn cancel(mut self, cancel_flag: Arc<AtomicBool>) -> Self {
self.cancel_flag = Some(cancel_flag);
self
}
pub async fn run(self) -> Result<EpDownloadResult> {
self.manager
.download_and_register_eps_impl(self.names, self.progress_callback, self.cancel_flag)
.await
}
}
impl FoundryLocalManager {
pub fn create(config: FoundryLocalConfig) -> Result<Arc<Self>> {
let mut slot = instance_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(existing) = slot.outer.upgrade() {
return Ok(existing);
}
if let Some(native) = slot.native.upgrade() {
let manager = Self::wrap_existing(native, config)?;
slot.outer = Arc::downgrade(&manager);
return Ok(manager);
}
let manager = Self::initialise(config)?;
slot.native = Arc::downgrade(&manager.native);
slot.outer = Arc::downgrade(&manager);
Ok(manager)
}
fn initialise(mut config: FoundryLocalConfig) -> Result<Arc<Self>> {
let api = Arc::new(Api::load(config.library_path_ref())?);
let logger = config.take_logger();
let native_config = config.build_native(&api)?;
let native = Arc::new(NativeManager::create(
Arc::clone(&api),
native_config.as_ptr(),
)?);
let catalog_ptr = native.catalog_ptr()?;
let catalog = Catalog::new(Arc::clone(&api), catalog_ptr, Arc::clone(&native))?;
Ok(Arc::new(FoundryLocalManager {
native,
catalog,
urls: Mutex::new(Vec::new()),
web_service_lock: AsyncMutex::new(()),
_logger: logger,
}))
}
fn wrap_existing(
native: Arc<NativeManager>,
mut config: FoundryLocalConfig,
) -> Result<Arc<Self>> {
let logger = config.take_logger();
let catalog_ptr = native.catalog_ptr()?;
let catalog = Catalog::new(native.api(), catalog_ptr, Arc::clone(&native))?;
let urls = native.web_service_urls().unwrap_or_default();
Ok(Arc::new(FoundryLocalManager {
native,
catalog,
urls: Mutex::new(urls),
web_service_lock: AsyncMutex::new(()),
_logger: logger,
}))
}
pub fn catalog(&self) -> &Catalog {
&self.catalog
}
pub fn shutdown(&self) -> Result<()> {
self.native.shutdown()
}
pub fn urls(&self) -> Result<Vec<String>> {
let lock = self.urls.lock().map_err(|_| FoundryLocalError::Internal {
reason: "Failed to acquire urls lock".into(),
})?;
Ok(lock.clone())
}
pub async fn start_web_service(&self) -> Result<()> {
let _guard = self.web_service_lock.lock().await;
let native = Arc::clone(&self.native);
let urls = spawn_blocking(move || {
native.web_service_start()?;
native.web_service_urls()
})
.await?;
*self.urls.lock().map_err(|_| FoundryLocalError::Internal {
reason: "Failed to acquire urls lock".into(),
})? = urls;
Ok(())
}
pub async fn stop_web_service(&self) -> Result<()> {
let _guard = self.web_service_lock.lock().await;
let native = Arc::clone(&self.native);
spawn_blocking(move || native.web_service_stop()).await?;
self.urls
.lock()
.map_err(|_| FoundryLocalError::Internal {
reason: "Failed to acquire urls lock".into(),
})?
.clear();
Ok(())
}
pub fn discover_eps(&self) -> Result<Vec<EpInfo>> {
self.native.discover_eps()
}
pub async fn download_and_register_eps(
&self,
names: Option<&[&str]>,
) -> Result<EpDownloadResult> {
let names = names.map(|n| n.iter().map(|s| s.to_string()).collect::<Vec<_>>());
self.download_and_register_eps_impl(names, None, None).await
}
pub async fn download_and_register_eps_with_progress<F>(
&self,
names: Option<&[&str]>,
progress_callback: F,
) -> Result<EpDownloadResult>
where
F: FnMut(&str, f64) + Send + 'static,
{
let names = names.map(|n| n.iter().map(|s| s.to_string()).collect::<Vec<_>>());
self.download_and_register_eps_impl(names, Some(Box::new(progress_callback)), None)
.await
}
pub fn download_and_register_eps_builder(&self) -> EpDownloadBuilder<'_> {
EpDownloadBuilder::new(self)
}
async fn download_and_register_eps_impl(
&self,
names: Option<Vec<String>>,
progress_callback: Option<EpDownloadProgressCallback>,
cancel_flag: Option<Arc<AtomicBool>>,
) -> Result<EpDownloadResult> {
let native = Arc::clone(&self.native);
let requested: Vec<String> = match &names {
Some(n) if !n.is_empty() => n.clone(),
_ => native.discover_eps()?.into_iter().map(|e| e.name).collect(),
};
let (message, after) = spawn_blocking(move || {
let name_refs: Option<Vec<&str>> = names
.as_ref()
.map(|n| n.iter().map(String::as_str).collect());
let progress: Option<EpProgressCallback> =
progress_callback.map(|cb| cb as EpProgressCallback);
let message =
native.download_and_register_eps(name_refs.as_deref(), progress, cancel_flag)?;
let after = native.discover_eps()?;
Ok::<(Option<String>, Vec<EpInfo>), FoundryLocalError>((message, after))
})
.await?;
let registered_eps: Vec<String> = requested
.iter()
.filter(|name| after.iter().any(|e| &e.name == *name && e.is_registered))
.cloned()
.collect();
let failed_eps: Vec<String> = requested
.iter()
.filter(|name| !registered_eps.contains(*name))
.cloned()
.collect();
let success = message.is_none() && failed_eps.is_empty();
let status = match &message {
None => "All requested execution providers were registered successfully.".to_string(),
Some(msg) if msg.is_empty() => {
"One or more execution providers failed to register.".to_string()
}
Some(msg) => msg.clone(),
};
let result = EpDownloadResult {
success,
status,
registered_eps,
failed_eps,
};
if result.success || !result.registered_eps.is_empty() {
let _ = self.catalog.update_models().await;
}
Ok(result)
}
}