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 worker_entry.worker.start(tokio_spawn_fn(&handle))?;
98 }
99
100 Ok(())
101 }
102
103 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 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 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 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}