pub(crate) mod pausable_worker;
use crate::worker::Worker;
use futures::stream::{FuturesUnordered, StreamExt};
use libdd_common::MutexExt;
use pausable_worker::{PausableWorker, PausableWorkerError};
use std::sync::atomic::AtomicU64;
use std::sync::{Arc, Mutex};
use std::{fmt, io};
use tracing::{debug, error};
#[cfg(not(target_arch = "wasm32"))]
mod native {
use super::*;
use pausable_worker::tokio_spawn_fn;
use std::sync::atomic::Ordering;
use tokio::runtime::{Builder, Runtime};
fn build_runtime() -> Result<Runtime, io::Error> {
Builder::new_multi_thread()
.worker_threads(1)
.enable_all()
.build()
}
impl SharedRuntime {
pub(in super::super) fn new_native() -> Result<Self, SharedRuntimeError> {
Ok(Self {
runtime: Arc::new(Mutex::new(Some(Arc::new(build_runtime()?)))),
workers: Arc::new(Mutex::new(Vec::new())),
next_worker_id: AtomicU64::new(1),
})
}
pub fn runtime_handle(&self) -> Result<tokio::runtime::Handle, SharedRuntimeError> {
Ok(self
.runtime
.lock_or_panic()
.as_ref()
.ok_or(SharedRuntimeError::RuntimeUnavailable)?
.handle()
.clone())
}
pub 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 SharedRuntime");
let mut pausable_worker = PausableWorker::new(boxed_worker);
let runtime_guard = self.runtime.lock_or_panic();
let mut workers_guard = self.workers.lock_or_panic();
if let Some(rt) = runtime_guard.as_ref() {
if let Err(e) = pausable_worker.start(tokio_spawn_fn(rt.handle())) {
return Err(e.into());
}
}
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,
});
Ok(WorkerHandle {
worker_id,
workers: self.workers.clone(),
})
}
pub fn before_fork(&self) {
debug!("before_fork: pausing all workers");
if let Some(runtime) = self.runtime.lock_or_panic().take() {
let mut workers_lock = self.workers.lock_or_panic();
runtime.block_on(async {
let futures: FuturesUnordered<_> = workers_lock
.iter_mut()
.map(|worker_entry| async {
if let Err(e) = worker_entry.worker.pause().await {
error!("Worker failed to pause before fork: {:?}", e);
}
})
.collect();
futures.collect::<()>().await;
});
}
}
fn restart_runtime(&self) -> Result<(), SharedRuntimeError> {
let mut runtime_lock = self.runtime.lock_or_panic();
if runtime_lock.is_none() {
*runtime_lock = Some(Arc::new(build_runtime()?));
}
Ok(())
}
pub fn after_fork_parent(&self) -> Result<(), SharedRuntimeError> {
debug!("after_fork_parent: restarting runtime and workers");
self.restart_runtime()?;
let runtime_lock = self.runtime.lock_or_panic();
let handle = runtime_lock
.as_ref()
.ok_or(SharedRuntimeError::RuntimeUnavailable)?
.handle()
.clone();
drop(runtime_lock);
let mut workers_lock = self.workers.lock_or_panic();
for worker_entry in workers_lock.iter_mut() {
worker_entry.worker.start(tokio_spawn_fn(&handle))?;
}
Ok(())
}
pub fn after_fork_child(&self) -> Result<(), SharedRuntimeError> {
debug!("after_fork_child: reinitializing runtime and workers");
self.restart_runtime()?;
let runtime_lock = self.runtime.lock_or_panic();
let handle = runtime_lock
.as_ref()
.ok_or(SharedRuntimeError::RuntimeUnavailable)?
.handle()
.clone();
drop(runtime_lock);
let mut workers_lock = self.workers.lock_or_panic();
workers_lock.retain(|entry| entry.restart_on_fork);
for worker_entry in workers_lock.iter_mut() {
worker_entry.worker.reset();
worker_entry.worker.start(tokio_spawn_fn(&handle))?;
}
Ok(())
}
pub fn block_on<F: std::future::Future>(&self, f: F) -> Result<F::Output, io::Error> {
let runtime = match self.runtime.lock_or_panic().as_ref() {
None => Arc::new(Builder::new_current_thread().enable_all().build()?),
Some(runtime) => runtime.clone(),
};
Ok(runtime.block_on(f))
}
pub fn shutdown(
&self,
timeout: Option<std::time::Duration>,
) -> Result<(), SharedRuntimeError> {
debug!(?timeout, "Shutting down SharedRuntime");
match self.runtime.lock_or_panic().take() {
Some(runtime) => {
if let Some(timeout) = timeout {
match runtime.block_on(async {
tokio::time::timeout(timeout, self.shutdown_async()).await
}) {
Ok(()) => Ok(()),
Err(_) => Err(SharedRuntimeError::ShutdownTimedOut(timeout)),
}
} else {
runtime.block_on(self.shutdown_async());
Ok(())
}
}
None => Ok(()),
}
}
}
}
type BoxedWorker = Box<dyn Worker + Sync>;
#[derive(Debug)]
struct WorkerEntry {
id: u64,
restart_on_fork: bool,
worker: PausableWorker<BoxedWorker>,
}
#[must_use = "dropping a WorkerHandle without calling stop() leaks the worker until the SharedRuntime is shut down"]
#[derive(Clone, Debug)]
pub struct WorkerHandle {
worker_id: u64,
workers: Arc<Mutex<Vec<WorkerEntry>>>,
}
#[derive(Debug)]
pub enum WorkerHandleError {
AlreadyStopped,
WorkerError(PausableWorkerError),
}
impl fmt::Display for WorkerHandleError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::AlreadyStopped => {
write!(f, "Worker has already been stopped")
}
Self::WorkerError(err) => write!(f, "Worker error: {}", err),
}
}
}
impl std::error::Error for WorkerHandleError {}
impl From<PausableWorkerError> for WorkerHandleError {
fn from(err: PausableWorkerError) -> Self {
Self::WorkerError(err)
}
}
impl WorkerHandle {
pub async fn stop(self) -> Result<(), WorkerHandleError> {
let mut worker = {
let mut workers_lock = self.workers.lock_or_panic();
let Some(position) = workers_lock
.iter()
.position(|entry| entry.id == self.worker_id)
else {
return Err(WorkerHandleError::AlreadyStopped);
};
let WorkerEntry { worker, .. } = workers_lock.swap_remove(position);
worker
};
worker.pause().await?;
worker.shutdown().await;
Ok(())
}
}
#[derive(Debug)]
pub enum SharedRuntimeError {
RuntimeUnavailable,
LockFailed(String),
WorkerError(PausableWorkerError),
RuntimeCreation(io::Error),
ShutdownTimedOut(std::time::Duration),
}
impl fmt::Display for SharedRuntimeError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::RuntimeUnavailable => {
write!(f, "Runtime is not available or in an invalid state")
}
Self::LockFailed(msg) => write!(f, "Failed to acquire lock: {}", msg),
Self::WorkerError(err) => write!(f, "Worker error: {}", err),
Self::RuntimeCreation(err) => {
write!(f, "Failed to create runtime: {}", err)
}
Self::ShutdownTimedOut(duration) => {
write!(f, "Shutdown timed out after {:?}", duration)
}
}
}
}
impl std::error::Error for SharedRuntimeError {}
impl From<PausableWorkerError> for SharedRuntimeError {
fn from(err: PausableWorkerError) -> Self {
SharedRuntimeError::WorkerError(err)
}
}
impl From<io::Error> for SharedRuntimeError {
fn from(err: io::Error) -> Self {
SharedRuntimeError::RuntimeCreation(err)
}
}
#[derive(Debug)]
pub struct SharedRuntime {
#[cfg(not(target_arch = "wasm32"))]
runtime: Arc<Mutex<Option<Arc<tokio::runtime::Runtime>>>>,
workers: Arc<Mutex<Vec<WorkerEntry>>>,
next_worker_id: AtomicU64,
}
impl SharedRuntime {
pub fn new() -> Result<Self, SharedRuntimeError> {
debug!("Creating new SharedRuntime");
#[cfg(not(target_arch = "wasm32"))]
{
Self::new_native()
}
#[cfg(target_arch = "wasm32")]
{
Ok(Self {
workers: Arc::new(Mutex::new(Vec::new())),
next_worker_id: AtomicU64::new(1),
})
}
}
#[cfg(target_arch = "wasm32")]
pub fn spawn_worker<T: Worker + Sync + 'static>(
&self,
worker: T,
restart_on_fork: bool,
) -> Result<WorkerHandle, SharedRuntimeError> {
use std::sync::atomic::Ordering;
let boxed_worker: BoxedWorker = Box::new(worker);
debug!(?boxed_worker, "Spawning worker on SharedRuntime");
let mut pausable_worker = PausableWorker::new(boxed_worker);
let mut workers_guard = self.workers.lock_or_panic();
if let Err(e) = 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) })
}) {
return Err(e.into());
}
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,
});
Ok(WorkerHandle {
worker_id,
workers: self.workers.clone(),
})
}
pub async fn shutdown_async(&self) {
debug!("Shutting down all workers asynchronously");
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;
}
}
#[cfg(test)]
mod tests {
use super::*;
use async_trait::async_trait;
use std::sync::mpsc::{channel, Receiver, Sender};
use std::time::Duration;
use tokio::time::sleep;
#[derive(Debug)]
struct TestWorker {
state: i32,
sender: Sender<i32>,
}
fn make_test_worker() -> (TestWorker, Receiver<i32>) {
let (sender, receiver) = channel::<i32>();
(TestWorker { state: 0, sender }, receiver)
}
#[async_trait]
impl Worker for TestWorker {
async fn run(&mut self) {
let _ = self.sender.send(self.state);
self.state += 1;
}
async fn trigger(&mut self) {
sleep(Duration::from_millis(100)).await;
}
fn reset(&mut self) {
self.state = 0;
}
async fn shutdown(&mut self) {
self.state = -1;
let _ = self.sender.send(self.state);
}
}
#[test]
fn test_shared_runtime_creation() {
let shared_runtime = SharedRuntime::new();
assert!(shared_runtime.is_ok());
}
#[test]
fn test_spawn_worker() {
let shared_runtime = SharedRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let result = shared_runtime.spawn_worker(worker, true);
assert!(result.is_ok());
assert_eq!(shared_runtime.workers.lock_or_panic().len(), 1);
assert_eq!(
receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not run"),
0
);
}
#[test]
fn test_worker_handle_stop() {
let rt = tokio::runtime::Runtime::new().unwrap();
let shared_runtime = SharedRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let handle = shared_runtime.spawn_worker(worker, true).unwrap();
assert_eq!(shared_runtime.workers.lock_or_panic().len(), 1);
receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not run");
rt.block_on(async {
assert!(handle.stop().await.is_ok());
});
assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
let mut last = receiver
.recv_timeout(Duration::from_secs(1))
.expect("shutdown did not send a value");
while let Ok(v) = receiver.try_recv() {
last = v;
}
assert_eq!(last, -1);
}
#[test]
fn test_before_and_after_fork_parent() {
let shared_runtime = SharedRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let _ = shared_runtime.spawn_worker(worker, true).unwrap();
let mut state_before_fork = 0;
while state_before_fork == 0 {
state_before_fork = receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not advance state before fork");
}
shared_runtime.before_fork();
while receiver.try_recv().is_ok() {}
assert!(shared_runtime.after_fork_parent().is_ok());
let after_fork_value = receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not resume after fork");
assert!(
after_fork_value > state_before_fork,
"after_fork_parent should preserve state: got {after_fork_value}, expected > {state_before_fork}"
);
}
#[test]
fn test_after_fork_child() {
let shared_runtime = SharedRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let _ = shared_runtime.spawn_worker(worker, true).unwrap();
let mut state_before_fork = 0;
while state_before_fork == 0 {
state_before_fork = receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not advance state before fork");
}
shared_runtime.before_fork();
while receiver.try_recv().is_ok() {}
assert!(shared_runtime.after_fork_child().is_ok());
let after_fork_value = receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not resume after fork child");
assert_eq!(
after_fork_value, 0,
"after_fork_child should reset state to 0, got {after_fork_value}"
);
}
#[test]
fn test_shutdown() {
let shared_runtime = SharedRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let _ = shared_runtime.spawn_worker(worker, true).unwrap();
receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not run");
shared_runtime.shutdown(None).unwrap();
let mut last = receiver
.recv_timeout(Duration::from_secs(1))
.expect("shutdown did not send a value");
while let Ok(v) = receiver.try_recv() {
last = v;
}
assert_eq!(last, -1);
}
#[test]
fn test_after_fork_child_drops_worker_not_restart_on_fork() {
let shared_runtime = SharedRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let _ = shared_runtime.spawn_worker(worker, false).unwrap();
receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not run");
shared_runtime.before_fork();
while receiver.try_recv().is_ok() {}
assert!(shared_runtime.after_fork_child().is_ok());
assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
assert!(
receiver.recv_timeout(Duration::from_millis(200)).is_err(),
"worker should not run or shut down after fork in child when restart_on_fork is false"
);
}
}