use crate::worker::Worker;
use futures::stream::{FuturesUnordered, StreamExt};
use libdd_common::MutexExt;
use std::io;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use tokio::runtime::{Builder, Runtime};
use tracing::{debug, error};
use super::{
pausable_worker::{tokio_spawn_fn, PausableWorker},
BlockingRuntime, BoxedWorker, SharedRuntime, SharedRuntimeError, WorkerEntry, WorkerHandle,
};
fn build_runtime(worker_threads: usize) -> Result<Runtime, io::Error> {
Builder::new_multi_thread()
.worker_threads(worker_threads)
.enable_all()
.build()
}
#[derive(Debug)]
pub struct ForkSafeRuntime {
worker_threads: usize,
runtime: Arc<Mutex<Option<Arc<Runtime>>>>,
workers: Arc<Mutex<Vec<WorkerEntry>>>,
next_worker_id: AtomicU64,
}
impl ForkSafeRuntime {
pub fn with_worker_threads(worker_threads: usize) -> Result<Self, SharedRuntimeError> {
let runtime = Arc::new(build_runtime(worker_threads)?);
Ok(Self {
worker_threads,
runtime: Arc::new(Mutex::new(Some(runtime))),
workers: Arc::new(Mutex::new(Vec::new())),
next_worker_id: AtomicU64::new(1),
})
}
pub fn before_fork(&self) {
debug!("before_fork: pausing all workers");
let mut runtime_lock = self.runtime.lock_or_panic();
let Some(runtime) = runtime_lock.take() else {
return;
};
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(self.worker_threads)?));
}
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() {
if let Err(e) = worker_entry.worker.start(tokio_spawn_fn(&handle)) {
error!(
worker_id = worker_entry.id,
"Worker failed to restart after fork in parent: {:?}", e
)
}
}
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();
if let Err(e) = worker_entry.worker.start(tokio_spawn_fn(&handle)) {
error!(
worker_id = worker_entry.id,
"Worker failed to restart after fork in parent: {:?}", e
)
}
}
Ok(())
}
pub fn shutdown(&self, timeout: Option<std::time::Duration>) -> Result<(), SharedRuntimeError> {
debug!(?timeout, "Shutting down ForkSafeRuntime");
match self.runtime.lock_or_panic().take() {
Some(runtime) => {
if let Some(timeout) = timeout {
match runtime.block_on(async {
tokio::time::timeout(timeout, <Self as SharedRuntime>::shutdown_async(self))
.await
}) {
Ok(()) => Ok(()),
Err(_) => Err(SharedRuntimeError::ShutdownTimedOut(timeout)),
}
} else {
runtime.block_on(<Self as SharedRuntime>::shutdown_async(self));
Ok(())
}
}
None => Ok(()),
}
}
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 ForkSafeRuntime {
fn new() -> Result<Self, SharedRuntimeError> {
Self::with_worker_threads(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 ForkSafeRuntime");
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() {
pausable_worker.start(tokio_spawn_fn(rt.handle()))?;
}
Ok(self.push_worker(&mut workers_guard, pausable_worker, restart_on_fork))
}
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;
}
}
impl BlockingRuntime for ForkSafeRuntime {
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))
}
}
#[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_fork_safe_runtime_creation() {
let shared_runtime = ForkSafeRuntime::new();
assert!(shared_runtime.is_ok());
}
#[test]
fn test_spawn_worker() {
let shared_runtime = ForkSafeRuntime::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 = ForkSafeRuntime::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 = ForkSafeRuntime::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 = ForkSafeRuntime::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 = ForkSafeRuntime::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 = ForkSafeRuntime::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"
);
}
#[test]
fn test_set_fork_restart_drops_worker_without_shutdown() {
let shared_runtime = ForkSafeRuntime::new().unwrap();
let (worker, receiver) = make_test_worker();
let handle = shared_runtime.spawn_worker(worker, true).unwrap();
receiver
.recv_timeout(Duration::from_secs(1))
.expect("worker did not run");
shared_runtime.before_fork();
while receiver.try_recv().is_ok() {}
handle.set_fork_restart(false).unwrap();
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 be dropped without running or shutting down in the fork child"
);
}
#[test]
fn after_fork_parent_skips_invalid_state_workers() {
let runtime = ForkSafeRuntime::new().unwrap();
let (good, good_rx) = make_test_worker();
let _ = runtime.spawn_worker(good, true).unwrap();
let (bad, _bad_rx) = make_test_worker();
let _ = runtime.spawn_worker(bad, true).unwrap();
good_rx
.recv_timeout(Duration::from_secs(1))
.expect("good worker did not run before fork");
{
let mut workers = runtime.workers.lock_or_panic();
workers[1].worker = PausableWorker::InvalidState;
}
runtime.before_fork();
while good_rx.try_recv().is_ok() {}
let result = runtime.after_fork_parent();
assert!(
result.is_ok(),
"after_fork_parent should not bail on a single InvalidState worker"
);
assert!(
good_rx.recv_timeout(Duration::from_secs(1)).is_ok(),
"good worker should resume after fork even if a peer is InvalidState"
);
}
}