use crate::worker::Worker;
use futures::stream::{FuturesUnordered, StreamExt};
use libdd_common::MutexExt;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use tracing::{debug, error};
use super::{
pausable_worker::PausableWorker, BoxedWorker, SharedRuntime, SharedRuntimeError, WorkerEntry,
WorkerHandle,
};
#[derive(Debug)]
pub struct LocalRuntime {
workers: Arc<Mutex<Vec<WorkerEntry>>>,
next_worker_id: AtomicU64,
}
impl LocalRuntime {
fn push_worker(
&self,
workers_guard: &mut std::sync::MutexGuard<Vec<WorkerEntry>>,
pausable_worker: PausableWorker<BoxedWorker>,
restart_on_fork: bool,
) -> WorkerHandle {
let worker_id = self.next_worker_id.fetch_add(1, Ordering::Relaxed);
workers_guard.push(WorkerEntry {
id: worker_id,
restart_on_fork,
worker: pausable_worker,
});
WorkerHandle {
worker_id,
workers: self.workers.clone(),
}
}
}
impl SharedRuntime for LocalRuntime {
fn new() -> Result<Self, SharedRuntimeError> {
Ok(Self {
workers: Arc::new(Mutex::new(Vec::new())),
next_worker_id: AtomicU64::new(1),
})
}
fn spawn_worker<T: Worker + Sync + 'static>(
&self,
worker: T,
restart_on_fork: bool,
) -> Result<WorkerHandle, SharedRuntimeError> {
let boxed_worker: BoxedWorker = Box::new(worker);
debug!(?boxed_worker, "Spawning worker on LocalRuntime");
let mut pausable_worker = PausableWorker::new(boxed_worker);
let mut workers_guard = self.workers.lock_or_panic();
pausable_worker.start(|future| {
use futures_util::FutureExt;
let (remote, handle) = future.remote_handle();
wasm_bindgen_futures::spawn_local(remote);
Box::pin(async { Ok(handle.await) })
})?;
Ok(self.push_worker(&mut workers_guard, pausable_worker, restart_on_fork))
}
async fn shutdown_async(&self) {
debug!("Shutting down all workers on LocalRuntime");
let workers = {
let mut workers_lock = self.workers.lock_or_panic();
std::mem::take(&mut *workers_lock)
};
let futures: FuturesUnordered<_> = workers
.into_iter()
.map(|mut worker_entry| async move {
if let Err(e) = worker_entry.worker.pause().await {
error!("Worker failed to shutdown: {:?}", e);
return;
}
worker_entry.worker.shutdown().await;
})
.collect();
futures.collect::<()>().await;
}
}