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            worker_entry.worker.start(tokio_spawn_fn(&handle))?;
98        }
99
100        Ok(())
101    }
102
103    /// Reinitializes the runtime in the child after forking.
104    /// Workers with `restart_on_fork = true` are reset and restarted; others are dropped
105    /// without shutdown.
106    pub fn after_fork_child(&self) -> Result<(), SharedRuntimeError> {
107        debug!("after_fork_child: reinitializing runtime and workers");
108        self.restart_runtime()?;
109
110        let runtime_lock = self.runtime.lock_or_panic();
111        let handle = runtime_lock
112            .as_ref()
113            .ok_or(SharedRuntimeError::RuntimeUnavailable)?
114            .handle()
115            .clone();
116        drop(runtime_lock);
117
118        let mut workers_lock = self.workers.lock_or_panic();
119
120        workers_lock.retain(|entry| entry.restart_on_fork);
121
122        for worker_entry in workers_lock.iter_mut() {
123            worker_entry.worker.reset();
124            worker_entry.worker.start(tokio_spawn_fn(&handle))?;
125        }
126
127        Ok(())
128    }
129
130    /// Shuts down all workers synchronously. Returns `ShutdownTimedOut` if `timeout` is
131    /// exceeded.
132    pub fn shutdown(&self, timeout: Option<std::time::Duration>) -> Result<(), SharedRuntimeError> {
133        debug!(?timeout, "Shutting down ForkSafeRuntime");
134        match self.runtime.lock_or_panic().take() {
135            Some(runtime) => {
136                if let Some(timeout) = timeout {
137                    match runtime.block_on(async {
138                        tokio::time::timeout(timeout, <Self as SharedRuntime>::shutdown_async(self))
139                            .await
140                    }) {
141                        Ok(()) => Ok(()),
142                        Err(_) => Err(SharedRuntimeError::ShutdownTimedOut(timeout)),
143                    }
144                } else {
145                    runtime.block_on(<Self as SharedRuntime>::shutdown_async(self));
146                    Ok(())
147                }
148            }
149            None => Ok(()),
150        }
151    }
152
153    fn push_worker(
154        &self,
155        workers_guard: &mut std::sync::MutexGuard<Vec<WorkerEntry>>,
156        pausable_worker: PausableWorker<BoxedWorker>,
157        restart_on_fork: bool,
158    ) -> WorkerHandle {
159        let worker_id = self.next_worker_id.fetch_add(1, Ordering::Relaxed);
160        workers_guard.push(WorkerEntry {
161            id: worker_id,
162            restart_on_fork,
163            worker: pausable_worker,
164        });
165        WorkerHandle {
166            worker_id,
167            workers: self.workers.clone(),
168        }
169    }
170}
171
172impl SharedRuntime for ForkSafeRuntime {
173    fn new() -> Result<Self, SharedRuntimeError> {
174        Self::with_worker_threads(1)
175    }
176
177    fn spawn_worker<T: Worker + Sync + 'static>(
178        &self,
179        worker: T,
180        restart_on_fork: bool,
181    ) -> Result<WorkerHandle, SharedRuntimeError> {
182        let boxed_worker: BoxedWorker = Box::new(worker);
183        debug!(?boxed_worker, "Spawning worker on ForkSafeRuntime");
184        let mut pausable_worker = PausableWorker::new(boxed_worker);
185
186        // Hold both locks together (runtime → workers, per struct lock order) so
187        // before_fork cannot interleave between start and push. If runtime is already
188        // None (fork window), skip start; after_fork_* will pick it up.
189        let runtime_guard = self.runtime.lock_or_panic();
190        let mut workers_guard = self.workers.lock_or_panic();
191
192        if let Some(rt) = runtime_guard.as_ref() {
193            pausable_worker.start(tokio_spawn_fn(rt.handle()))?;
194        }
195
196        Ok(self.push_worker(&mut workers_guard, pausable_worker, restart_on_fork))
197    }
198
199    async fn shutdown_async(&self) {
200        debug!("Shutting down all workers asynchronously");
201        let workers = {
202            let mut workers_lock = self.workers.lock_or_panic();
203            std::mem::take(&mut *workers_lock)
204        };
205
206        let futures: FuturesUnordered<_> = workers
207            .into_iter()
208            .map(|mut worker_entry| async move {
209                if let Err(e) = worker_entry.worker.pause().await {
210                    error!("Worker failed to shutdown: {:?}", e);
211                    return;
212                }
213                worker_entry.worker.shutdown().await;
214            })
215            .collect();
216
217        futures.collect::<()>().await;
218    }
219}
220
221impl BlockingRuntime for ForkSafeRuntime {
222    /// Falls back to a temporary current-thread runtime in the fork window.
223    fn block_on<F: std::future::Future>(&self, f: F) -> Result<F::Output, io::Error> {
224        let runtime = match self.runtime.lock_or_panic().as_ref() {
225            None => Arc::new(Builder::new_current_thread().enable_all().build()?),
226            Some(runtime) => runtime.clone(),
227        };
228        Ok(runtime.block_on(f))
229    }
230}
231
232#[cfg(test)]
233mod tests {
234    use super::*;
235    use async_trait::async_trait;
236    use std::sync::mpsc::{channel, Receiver, Sender};
237    use std::time::Duration;
238    use tokio::time::sleep;
239
240    #[derive(Debug)]
241    struct TestWorker {
242        state: i32,
243        sender: Sender<i32>,
244    }
245
246    fn make_test_worker() -> (TestWorker, Receiver<i32>) {
247        let (sender, receiver) = channel::<i32>();
248        (TestWorker { state: 0, sender }, receiver)
249    }
250
251    #[async_trait]
252    impl Worker for TestWorker {
253        async fn run(&mut self) {
254            let _ = self.sender.send(self.state);
255            self.state += 1;
256        }
257
258        async fn trigger(&mut self) {
259            sleep(Duration::from_millis(100)).await;
260        }
261
262        fn reset(&mut self) {
263            self.state = 0;
264        }
265
266        async fn shutdown(&mut self) {
267            self.state = -1;
268            let _ = self.sender.send(self.state);
269        }
270    }
271
272    #[test]
273    fn test_fork_safe_runtime_creation() {
274        let shared_runtime = ForkSafeRuntime::new();
275        assert!(shared_runtime.is_ok());
276    }
277
278    #[test]
279    fn test_spawn_worker() {
280        let shared_runtime = ForkSafeRuntime::new().unwrap();
281        let (worker, receiver) = make_test_worker();
282
283        let result = shared_runtime.spawn_worker(worker, true);
284        assert!(result.is_ok());
285        assert_eq!(shared_runtime.workers.lock_or_panic().len(), 1);
286
287        assert_eq!(
288            receiver
289                .recv_timeout(Duration::from_secs(1))
290                .expect("worker did not run"),
291            0
292        );
293    }
294
295    #[test]
296    fn test_worker_handle_stop() {
297        let rt = tokio::runtime::Runtime::new().unwrap();
298        let shared_runtime = ForkSafeRuntime::new().unwrap();
299        let (worker, receiver) = make_test_worker();
300
301        let handle = shared_runtime.spawn_worker(worker, true).unwrap();
302        assert_eq!(shared_runtime.workers.lock_or_panic().len(), 1);
303
304        receiver
305            .recv_timeout(Duration::from_secs(1))
306            .expect("worker did not run");
307
308        rt.block_on(async {
309            assert!(handle.stop().await.is_ok());
310        });
311
312        assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
313
314        let mut last = receiver
315            .recv_timeout(Duration::from_secs(1))
316            .expect("shutdown did not send a value");
317        while let Ok(v) = receiver.try_recv() {
318            last = v;
319        }
320        assert_eq!(last, -1);
321    }
322
323    #[test]
324    fn test_before_and_after_fork_parent() {
325        let shared_runtime = ForkSafeRuntime::new().unwrap();
326        let (worker, receiver) = make_test_worker();
327
328        let _ = shared_runtime.spawn_worker(worker, true).unwrap();
329
330        let mut state_before_fork = 0;
331        while state_before_fork == 0 {
332            state_before_fork = receiver
333                .recv_timeout(Duration::from_secs(1))
334                .expect("worker did not advance state before fork");
335        }
336
337        shared_runtime.before_fork();
338        while receiver.try_recv().is_ok() {}
339
340        assert!(shared_runtime.after_fork_parent().is_ok());
341
342        let after_fork_value = receiver
343            .recv_timeout(Duration::from_secs(1))
344            .expect("worker did not resume after fork");
345        assert!(
346            after_fork_value > state_before_fork,
347            "after_fork_parent should preserve state: got {after_fork_value}, expected > {state_before_fork}"
348        );
349    }
350
351    #[test]
352    fn test_after_fork_child() {
353        let shared_runtime = ForkSafeRuntime::new().unwrap();
354        let (worker, receiver) = make_test_worker();
355
356        let _ = shared_runtime.spawn_worker(worker, true).unwrap();
357
358        let mut state_before_fork = 0;
359        while state_before_fork == 0 {
360            state_before_fork = receiver
361                .recv_timeout(Duration::from_secs(1))
362                .expect("worker did not advance state before fork");
363        }
364
365        shared_runtime.before_fork();
366        while receiver.try_recv().is_ok() {}
367
368        assert!(shared_runtime.after_fork_child().is_ok());
369
370        let after_fork_value = receiver
371            .recv_timeout(Duration::from_secs(1))
372            .expect("worker did not resume after fork child");
373        assert_eq!(
374            after_fork_value, 0,
375            "after_fork_child should reset state to 0, got {after_fork_value}"
376        );
377    }
378
379    #[test]
380    fn test_shutdown() {
381        let shared_runtime = ForkSafeRuntime::new().unwrap();
382        let (worker, receiver) = make_test_worker();
383
384        let _ = shared_runtime.spawn_worker(worker, true).unwrap();
385
386        receiver
387            .recv_timeout(Duration::from_secs(1))
388            .expect("worker did not run");
389
390        shared_runtime.shutdown(None).unwrap();
391
392        let mut last = receiver
393            .recv_timeout(Duration::from_secs(1))
394            .expect("shutdown did not send a value");
395        while let Ok(v) = receiver.try_recv() {
396            last = v;
397        }
398        assert_eq!(last, -1);
399    }
400
401    #[test]
402    fn test_after_fork_child_drops_worker_not_restart_on_fork() {
403        let shared_runtime = ForkSafeRuntime::new().unwrap();
404        let (worker, receiver) = make_test_worker();
405
406        let _ = shared_runtime.spawn_worker(worker, false).unwrap();
407
408        receiver
409            .recv_timeout(Duration::from_secs(1))
410            .expect("worker did not run");
411
412        shared_runtime.before_fork();
413        while receiver.try_recv().is_ok() {}
414
415        assert!(shared_runtime.after_fork_child().is_ok());
416
417        assert_eq!(shared_runtime.workers.lock_or_panic().len(), 0);
418
419        assert!(
420            receiver.recv_timeout(Duration::from_millis(200)).is_err(),
421            "worker should not run or shut down after fork in child when restart_on_fork is false"
422        );
423    }
424}