Skip to main content

libdd_shared_runtime/shared_runtime/
fork_safe.rs

1// Copyright 2026-Present Datadog, Inc. https://www.datadoghq.com/
2// SPDX-License-Identifier: Apache-2.0
3
4use 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/// Owns a tokio runtime and manages [`PausableWorker`]s on it.
26///
27/// Supports the full fork protocol ([`before_fork`](Self::before_fork) /
28/// [`after_fork_parent`](Self::after_fork_parent) /
29/// [`after_fork_child`](Self::after_fork_child)) and synchronous [`shutdown`](Self::shutdown).
30#[derive(Debug)]
31pub struct ForkSafeRuntime {
32    worker_threads: usize,
33    // Lock order: `runtime` must be acquired before `workers`.
34    runtime: Arc<Mutex<Option<Arc<Runtime>>>>,
35    workers: Arc<Mutex<Vec<WorkerEntry>>>,
36    next_worker_id: AtomicU64,
37}
38
39impl ForkSafeRuntime {
40    /// Creates a `ForkSafeRuntime` with the given number of tokio worker threads.
41    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    /// Pauses all workers before `fork()`. Worker pause errors are logged, not propagated.
52    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    /// Restarts the runtime and workers in the parent after forking; worker state is preserved.
82    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    /// Reinitializes the runtime in the child after forking.
109    /// Workers with `restart_on_fork = true` are reset and restarted; others are dropped
110    /// without shutdown.
111    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    /// Shuts down all workers synchronously. Returns `ShutdownTimedOut` if `timeout` is
141    /// exceeded.
142    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        // Hold both locks together (runtime → workers, per struct lock order) so
197        // before_fork cannot interleave between start and push. If runtime is already
198        // None (fork window), skip start; after_fork_* will pick it up.
199        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    /// Falls back to a temporary current-thread runtime in the fork window.
233    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    /// A single `PausableWorker` in `InvalidState` must
460    /// not abort the whole restart loop in `after_fork_parent`
461    #[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        // Second worker — we'll corrupt its entry into InvalidState below,
469        // simulating a previously-aborted task.
470        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        // Drain good worker queue
485        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}