libdd_shared_runtime/shared_runtime/
fork_safe.rs1use crate::worker::Worker;
5use futures::stream::{FuturesUnordered, StreamExt};
6use libdd_common::MutexExt;
7use std::io;
8use std::sync::atomic::{AtomicU64, Ordering};
9use std::sync::{Arc, Mutex};
10use tokio::runtime::{Builder, Runtime};
11use tracing::{debug, error};
12
13use super::{
14 pausable_worker::{tokio_spawn_fn, PausableWorker},
15 BlockingRuntime, BoxedWorker, SharedRuntime, SharedRuntimeError, WorkerEntry, WorkerHandle,
16};
17
18fn build_runtime(worker_threads: usize) -> Result<Runtime, io::Error> {
19 Builder::new_multi_thread()
20 .worker_threads(worker_threads)
21 .enable_all()
22 .build()
23}
24
25#[derive(Debug)]
31pub struct ForkSafeRuntime {
32 worker_threads: usize,
33 runtime: Arc<Mutex<Option<Arc<Runtime>>>>,
35 workers: Arc<Mutex<Vec<WorkerEntry>>>,
36 next_worker_id: AtomicU64,
37}
38
39impl ForkSafeRuntime {
40 pub fn with_worker_threads(worker_threads: usize) -> Result<Self, SharedRuntimeError> {
42 let runtime = Arc::new(build_runtime(worker_threads)?);
43 Ok(Self {
44 worker_threads,
45 runtime: Arc::new(Mutex::new(Some(runtime))),
46 workers: Arc::new(Mutex::new(Vec::new())),
47 next_worker_id: AtomicU64::new(1),
48 })
49 }
50
51 pub fn before_fork(&self) {
53 debug!("before_fork: pausing all workers");
54 let mut runtime_lock = self.runtime.lock_or_panic();
55 let Some(runtime) = runtime_lock.take() else {
56 return;
57 };
58 let mut workers_lock = self.workers.lock_or_panic();
59 runtime.block_on(async {
60 let futures: FuturesUnordered<_> = workers_lock
61 .iter_mut()
62 .map(|worker_entry| async {
63 if let Err(e) = worker_entry.worker.pause().await {
64 error!("Worker failed to pause before fork: {:?}", e);
65 }
66 })
67 .collect();
68
69 futures.collect::<()>().await;
70 });
71 }
72
73 fn restart_runtime(&self) -> Result<(), SharedRuntimeError> {
74 let mut runtime_lock = self.runtime.lock_or_panic();
75 if runtime_lock.is_none() {
76 *runtime_lock = Some(Arc::new(build_runtime(self.worker_threads)?));
77 }
78 Ok(())
79 }
80
81 pub fn after_fork_parent(&self) -> Result<(), SharedRuntimeError> {
83 debug!("after_fork_parent: restarting runtime and workers");
84 self.restart_runtime()?;
85
86 let runtime_lock = self.runtime.lock_or_panic();
87 let handle = runtime_lock
88 .as_ref()
89 .ok_or(SharedRuntimeError::RuntimeUnavailable)?
90 .handle()
91 .clone();
92 drop(runtime_lock);
93
94 let mut workers_lock = self.workers.lock_or_panic();
95
96 for worker_entry in workers_lock.iter_mut() {
97 if let Err(e) = worker_entry.worker.start(tokio_spawn_fn(&handle)) {
98 error!(
99 worker_id = worker_entry.id,
100 "Worker failed to restart after fork in parent: {:?}", e
101 )
102 }
103 }
104
105 Ok(())
106 }
107
108 pub fn after_fork_child(&self) -> Result<(), SharedRuntimeError> {
112 debug!("after_fork_child: reinitializing runtime and workers");
113 self.restart_runtime()?;
114
115 let runtime_lock = self.runtime.lock_or_panic();
116 let handle = runtime_lock
117 .as_ref()
118 .ok_or(SharedRuntimeError::RuntimeUnavailable)?
119 .handle()
120 .clone();
121 drop(runtime_lock);
122
123 let mut workers_lock = self.workers.lock_or_panic();
124
125 workers_lock.retain(|entry| entry.restart_on_fork);
126
127 for worker_entry in workers_lock.iter_mut() {
128 worker_entry.worker.reset();
129 if let Err(e) = worker_entry.worker.start(tokio_spawn_fn(&handle)) {
130 error!(
131 worker_id = worker_entry.id,
132 "Worker failed to restart after fork in parent: {:?}", e
133 )
134 }
135 }
136
137 Ok(())
138 }
139
140 pub fn shutdown(&self, timeout: Option<std::time::Duration>) -> Result<(), SharedRuntimeError> {
143 debug!(?timeout, "Shutting down ForkSafeRuntime");
144 match self.runtime.lock_or_panic().take() {
145 Some(runtime) => {
146 if let Some(timeout) = timeout {
147 match runtime.block_on(async {
148 tokio::time::timeout(timeout, <Self as SharedRuntime>::shutdown_async(self))
149 .await
150 }) {
151 Ok(()) => Ok(()),
152 Err(_) => Err(SharedRuntimeError::ShutdownTimedOut(timeout)),
153 }
154 } else {
155 runtime.block_on(<Self as SharedRuntime>::shutdown_async(self));
156 Ok(())
157 }
158 }
159 None => Ok(()),
160 }
161 }
162
163 fn push_worker(
164 &self,
165 workers_guard: &mut std::sync::MutexGuard<Vec<WorkerEntry>>,
166 pausable_worker: PausableWorker<BoxedWorker>,
167 restart_on_fork: bool,
168 ) -> WorkerHandle {
169 let worker_id = self.next_worker_id.fetch_add(1, Ordering::Relaxed);
170 workers_guard.push(WorkerEntry {
171 id: worker_id,
172 restart_on_fork,
173 worker: pausable_worker,
174 });
175 WorkerHandle {
176 worker_id,
177 workers: self.workers.clone(),
178 }
179 }
180}
181
182impl SharedRuntime for ForkSafeRuntime {
183 fn new() -> Result<Self, SharedRuntimeError> {
184 Self::with_worker_threads(1)
185 }
186
187 fn spawn_worker<T: Worker + Sync + 'static>(
188 &self,
189 worker: T,
190 restart_on_fork: bool,
191 ) -> Result<WorkerHandle, SharedRuntimeError> {
192 let boxed_worker: BoxedWorker = Box::new(worker);
193 debug!(?boxed_worker, "Spawning worker on ForkSafeRuntime");
194 let mut pausable_worker = PausableWorker::new(boxed_worker);
195
196 let runtime_guard = self.runtime.lock_or_panic();
200 let mut workers_guard = self.workers.lock_or_panic();
201
202 if let Some(rt) = runtime_guard.as_ref() {
203 pausable_worker.start(tokio_spawn_fn(rt.handle()))?;
204 }
205
206 Ok(self.push_worker(&mut workers_guard, pausable_worker, restart_on_fork))
207 }
208
209 async fn shutdown_async(&self) {
210 debug!("Shutting down all workers asynchronously");
211 let workers = {
212 let mut workers_lock = self.workers.lock_or_panic();
213 std::mem::take(&mut *workers_lock)
214 };
215
216 let futures: FuturesUnordered<_> = workers
217 .into_iter()
218 .map(|mut worker_entry| async move {
219 if let Err(e) = worker_entry.worker.pause().await {
220 error!("Worker failed to shutdown: {:?}", e);
221 return;
222 }
223 worker_entry.worker.shutdown().await;
224 })
225 .collect();
226
227 futures.collect::<()>().await;
228 }
229}
230
231impl BlockingRuntime for ForkSafeRuntime {
232 fn block_on<F: std::future::Future>(&self, f: F) -> Result<F::Output, io::Error> {
234 let runtime = match self.runtime.lock_or_panic().as_ref() {
235 None => Arc::new(Builder::new_current_thread().enable_all().build()?),
236 Some(runtime) => runtime.clone(),
237 };
238 Ok(runtime.block_on(f))
239 }
240}
241
242#[cfg(test)]
243mod tests {
244 use super::*;
245 use async_trait::async_trait;
246 use std::sync::mpsc::{channel, Receiver, Sender};
247 use std::time::Duration;
248 use tokio::time::sleep;
249
250 #[derive(Debug)]
251 struct TestWorker {
252 state: i32,
253 sender: Sender<i32>,
254 }
255
256 fn make_test_worker() -> (TestWorker, Receiver<i32>) {
257 let (sender, receiver) = channel::<i32>();
258 (TestWorker { state: 0, sender }, receiver)
259 }
260
261 #[async_trait]
262 impl Worker for TestWorker {
263 async fn run(&mut self) {
264 let _ = self.sender.send(self.state);
265 self.state += 1;
266 }
267
268 async fn trigger(&mut self) {
269 sleep(Duration::from_millis(100)).await;
270 }
271
272 fn reset(&mut self) {
273 self.state = 0;
274 }
275
276 async fn shutdown(&mut self) {
277 self.state = -1;
278 let _ = self.sender.send(self.state);
279 }
280 }
281
282 #[test]
283 fn test_fork_safe_runtime_creation() {
284 let shared_runtime = ForkSafeRuntime::new();
285 assert!(shared_runtime.is_ok());
286 }
287
288 #[test]
289 fn test_spawn_worker() {
290 let shared_runtime = ForkSafeRuntime::new().unwrap();
291 let (worker, receiver) = make_test_worker();
292
293 let result = shared_runtime.spawn_worker(worker, true);
294 assert!(result.is_ok());
295 assert_eq!(shared_runtime.workers.lock_or_panic().len(), 1);
296
297 assert_eq!(
298 receiver
299 .recv_timeout(Duration::from_secs(1))
300 .expect("worker did not run"),
301 0
302 );
303 }
304
305 #[test]
306 fn test_worker_handle_stop() {
307 let rt = tokio::runtime::Runtime::new().unwrap();
308 let shared_runtime = ForkSafeRuntime::new().unwrap();
309 let (worker, receiver) = make_test_worker();
310
311 let handle = shared_runtime.spawn_worker(worker, true).unwrap();
312 assert_eq!(shared_runtime.workers.lock_or_panic().len(), 1);
313
314 receiver
315 .recv_timeout(Duration::from_secs(1))
316 .expect("worker did not run");
317
318 rt.block_on(async {
319 assert!(handle.stop().await.is_ok());
320 });
321
322 assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
323
324 let mut last = receiver
325 .recv_timeout(Duration::from_secs(1))
326 .expect("shutdown did not send a value");
327 while let Ok(v) = receiver.try_recv() {
328 last = v;
329 }
330 assert_eq!(last, -1);
331 }
332
333 #[test]
334 fn test_before_and_after_fork_parent() {
335 let shared_runtime = ForkSafeRuntime::new().unwrap();
336 let (worker, receiver) = make_test_worker();
337
338 let _ = shared_runtime.spawn_worker(worker, true).unwrap();
339
340 let mut state_before_fork = 0;
341 while state_before_fork == 0 {
342 state_before_fork = receiver
343 .recv_timeout(Duration::from_secs(1))
344 .expect("worker did not advance state before fork");
345 }
346
347 shared_runtime.before_fork();
348 while receiver.try_recv().is_ok() {}
349
350 assert!(shared_runtime.after_fork_parent().is_ok());
351
352 let after_fork_value = receiver
353 .recv_timeout(Duration::from_secs(1))
354 .expect("worker did not resume after fork");
355 assert!(
356 after_fork_value > state_before_fork,
357 "after_fork_parent should preserve state: got {after_fork_value}, expected > {state_before_fork}"
358 );
359 }
360
361 #[test]
362 fn test_after_fork_child() {
363 let shared_runtime = ForkSafeRuntime::new().unwrap();
364 let (worker, receiver) = make_test_worker();
365
366 let _ = shared_runtime.spawn_worker(worker, true).unwrap();
367
368 let mut state_before_fork = 0;
369 while state_before_fork == 0 {
370 state_before_fork = receiver
371 .recv_timeout(Duration::from_secs(1))
372 .expect("worker did not advance state before fork");
373 }
374
375 shared_runtime.before_fork();
376 while receiver.try_recv().is_ok() {}
377
378 assert!(shared_runtime.after_fork_child().is_ok());
379
380 let after_fork_value = receiver
381 .recv_timeout(Duration::from_secs(1))
382 .expect("worker did not resume after fork child");
383 assert_eq!(
384 after_fork_value, 0,
385 "after_fork_child should reset state to 0, got {after_fork_value}"
386 );
387 }
388
389 #[test]
390 fn test_shutdown() {
391 let shared_runtime = ForkSafeRuntime::new().unwrap();
392 let (worker, receiver) = make_test_worker();
393
394 let _ = shared_runtime.spawn_worker(worker, true).unwrap();
395
396 receiver
397 .recv_timeout(Duration::from_secs(1))
398 .expect("worker did not run");
399
400 shared_runtime.shutdown(None).unwrap();
401
402 let mut last = receiver
403 .recv_timeout(Duration::from_secs(1))
404 .expect("shutdown did not send a value");
405 while let Ok(v) = receiver.try_recv() {
406 last = v;
407 }
408 assert_eq!(last, -1);
409 }
410
411 #[test]
412 fn test_after_fork_child_drops_worker_not_restart_on_fork() {
413 let shared_runtime = ForkSafeRuntime::new().unwrap();
414 let (worker, receiver) = make_test_worker();
415
416 let _ = shared_runtime.spawn_worker(worker, false).unwrap();
417
418 receiver
419 .recv_timeout(Duration::from_secs(1))
420 .expect("worker did not run");
421
422 shared_runtime.before_fork();
423 while receiver.try_recv().is_ok() {}
424
425 assert!(shared_runtime.after_fork_child().is_ok());
426
427 assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
428
429 assert!(
430 receiver.recv_timeout(Duration::from_millis(200)).is_err(),
431 "worker should not run or shut down after fork in child when restart_on_fork is false"
432 );
433 }
434
435 #[test]
436 fn test_set_fork_restart_drops_worker_without_shutdown() {
437 let shared_runtime = ForkSafeRuntime::new().unwrap();
438 let (worker, receiver) = make_test_worker();
439
440 let handle = shared_runtime.spawn_worker(worker, true).unwrap();
441
442 receiver
443 .recv_timeout(Duration::from_secs(1))
444 .expect("worker did not run");
445
446 shared_runtime.before_fork();
447 while receiver.try_recv().is_ok() {}
448
449 handle.set_fork_restart(false).unwrap();
450 assert!(shared_runtime.after_fork_child().is_ok());
451
452 assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
453 assert!(
454 receiver.recv_timeout(Duration::from_millis(200)).is_err(),
455 "worker should be dropped without running or shutting down in the fork child"
456 );
457 }
458
459 #[test]
462 fn after_fork_parent_skips_invalid_state_workers() {
463 let runtime = ForkSafeRuntime::new().unwrap();
464
465 let (good, good_rx) = make_test_worker();
466 let _ = runtime.spawn_worker(good, true).unwrap();
467
468 let (bad, _bad_rx) = make_test_worker();
471 let _ = runtime.spawn_worker(bad, true).unwrap();
472
473 good_rx
474 .recv_timeout(Duration::from_secs(1))
475 .expect("good worker did not run before fork");
476
477 {
478 let mut workers = runtime.workers.lock_or_panic();
479 workers[1].worker = PausableWorker::InvalidState;
480 }
481
482 runtime.before_fork();
483
484 while good_rx.try_recv().is_ok() {}
486
487 let result = runtime.after_fork_parent();
488
489 assert!(
490 result.is_ok(),
491 "after_fork_parent should not bail on a single InvalidState worker"
492 );
493 assert!(
494 good_rx.recv_timeout(Duration::from_secs(1)).is_ok(),
495 "good worker should resume after fork even if a peer is InvalidState"
496 );
497 }
498}