Skip to main content

shadow_rs/utility/
childpid_watcher.rs

1use std::collections::HashMap;
2use std::sync::Arc;
3use std::sync::Mutex;
4use std::thread;
5
6use linux_api::errno::Errno;
7use linux_api::posix_types::Pid;
8use rustix::event::{self, epoll};
9use rustix::fd::AsFd;
10use rustix::fd::OwnedFd;
11use rustix::io::FdFlags;
12use rustix::process::PidfdFlags;
13
14/// Utility for monitoring a set of child pid's, calling registered callbacks
15/// when one exits or is killed. Starts a background thread, which is shut down
16/// when the object is dropped.
17#[derive(Debug)]
18pub struct ChildPidWatcher {
19    inner: Arc<Mutex<Inner>>,
20    epoll: Arc<OwnedFd>,
21}
22
23pub type WatchHandle = u64;
24
25#[derive(Debug)]
26enum Command {
27    RunCallbacks(Pid),
28    UnregisterPid(Pid),
29    Finish,
30}
31
32struct PidData {
33    // Registered callbacks.
34    callbacks: HashMap<WatchHandle, Box<dyn Send + FnOnce(Pid)>>,
35    // After the pid has exited, this fd is closed and set to None.
36    pidfd: Option<OwnedFd>,
37    // Whether this pid has been unregistered. The whole struct is removed after
38    // both the pid is unregistered, and `callbacks` is empty.
39    unregistered: bool,
40}
41
42#[derive(Debug)]
43struct Inner {
44    // Next unique handle ID.
45    next_handle: WatchHandle,
46    // Pending commands for watcher thread.
47    commands: Vec<Command>,
48    // Data for each monitored pid.
49    pids: HashMap<Pid, PidData>,
50    // event_fd used to notify watcher thread via epoll. Calling thread writes a
51    // single byte, which the watcher thread reads to reset.
52    command_notifier: OwnedFd,
53    thread_handle: Option<thread::JoinHandle<()>>,
54}
55
56impl Inner {
57    fn send_command(&mut self, cmd: Command) {
58        self.commands.push(cmd);
59        rustix::io::write(&self.command_notifier, &1u64.to_ne_bytes()).unwrap();
60    }
61
62    fn unwatch_pid(&mut self, epoll: impl AsFd, pid: Pid) {
63        let Some(piddata) = self.pids.get_mut(&pid) else {
64            // Already unregistered the pid
65            return;
66        };
67        let Some(fd) = piddata.pidfd.take() else {
68            // Already unwatched the pid
69            return;
70        };
71        epoll::delete(epoll, fd).unwrap();
72    }
73
74    fn pid_has_exited(&self, pid: Pid) -> bool {
75        self.pids.get(&pid).unwrap().pidfd.is_none()
76    }
77
78    fn remove_pid(&mut self, epoll: impl AsFd, pid: Pid) {
79        debug_assert!(self.should_remove_pid(pid));
80        self.unwatch_pid(epoll, pid);
81        self.pids.remove(&pid);
82    }
83
84    fn run_callbacks_for_pid(&mut self, pid: Pid) {
85        for (_handle, cb) in self.pids.get_mut(&pid).unwrap().callbacks.drain() {
86            cb(pid)
87        }
88    }
89
90    fn should_remove_pid(&mut self, pid: Pid) -> bool {
91        let pid_data = self.pids.get(&pid).unwrap();
92        pid_data.callbacks.is_empty() && pid_data.unregistered
93    }
94
95    fn maybe_remove_pid(&mut self, epoll: impl AsFd, pid: Pid) {
96        if self.should_remove_pid(pid) {
97            self.remove_pid(epoll, pid)
98        }
99    }
100}
101
102impl ChildPidWatcher {
103    /// Create a ChildPidWatcher. Spawns a background thread, which is joined
104    /// when the object is dropped.
105    pub fn new() -> Self {
106        let epoll = Arc::new(epoll::create(epoll::CreateFlags::CLOEXEC).unwrap());
107        let command_notifier = event::eventfd(
108            0,
109            event::EventfdFlags::NONBLOCK | event::EventfdFlags::CLOEXEC,
110        )
111        .unwrap();
112        epoll::add(
113            &epoll,
114            &command_notifier,
115            epoll::EventData::new_u64(0),
116            epoll::EventFlags::IN,
117        )
118        .unwrap();
119        let watcher = ChildPidWatcher {
120            inner: Arc::new(Mutex::new(Inner {
121                next_handle: 1,
122                pids: HashMap::new(),
123                commands: Vec::new(),
124                command_notifier,
125                thread_handle: None,
126            })),
127            epoll,
128        };
129        let thread_handle = {
130            let inner = Arc::clone(&watcher.inner);
131            let epoll = watcher.epoll.clone();
132            thread::Builder::new()
133                .name("child-pid-watcher".into())
134                .spawn(move || ChildPidWatcher::thread_loop(&inner, &epoll))
135                .unwrap()
136        };
137        watcher.inner.lock().unwrap().thread_handle = Some(thread_handle);
138        watcher
139    }
140
141    fn thread_loop(inner: &Mutex<Inner>, epoll: impl AsFd) {
142        let mut commands = Vec::new();
143        let mut done = false;
144        let mut events = Vec::<epoll::Event>::with_capacity(10);
145
146        while !done {
147            // Start each loop iteration with a clear vec.
148            events.clear();
149            let spare_buf = rustix::buffer::spare_capacity(&mut events);
150
151            match epoll::wait(epoll.as_fd(), spare_buf, None) {
152                Ok(_) => (),
153                Err(rustix::io::Errno::INTR) => {
154                    // Just try again.
155                    continue;
156                }
157                Err(e) => panic!("epoll_wait: {e:?}"),
158            };
159
160            // We hold the lock the whole time we're processing events. While it'd
161            // be nice to avoid holding it while executing callbacks (and therefore
162            // not require that callbacks don't call ChildPidWatcher APIs), that'd
163            // make it difficult to guarantee a callback *won't* be run if the
164            // caller unregisters it.
165            let mut inner = inner.lock().unwrap();
166
167            for event in events.drain(..) {
168                if event.data.u64() == 0 {
169                    // We get an event for pid=0 when there's a write to the
170                    // command_notifier; Ignore that here and handle below.
171                    continue;
172                }
173                let pid = Pid::from_raw(i32::try_from(event.data.u64()).unwrap()).unwrap();
174                inner.unwatch_pid(epoll.as_fd(), pid);
175                inner.run_callbacks_for_pid(pid);
176                inner.maybe_remove_pid(epoll.as_fd(), pid);
177            }
178
179            // Reading an eventfd always returns an 8 byte integer. Do so to ensure it's
180            // no longer marked 'readable'.
181            let mut buf = [0; 8];
182            let res = rustix::io::read(&inner.command_notifier, &mut buf);
183            debug_assert!(match res {
184                Ok(8) => true,
185                Ok(i) => panic!("Unexpected read size {i}"),
186                Err(rustix::io::Errno::AGAIN) => true,
187                Err(e) => panic!("Unexpected error {e:?}"),
188            });
189
190            // Run commands
191            std::mem::swap(&mut commands, &mut inner.commands);
192            for cmd in commands.drain(..) {
193                match cmd {
194                    Command::RunCallbacks(pid) => {
195                        debug_assert!(inner.pid_has_exited(pid));
196                        inner.run_callbacks_for_pid(pid);
197                        inner.maybe_remove_pid(epoll.as_fd(), pid);
198                    }
199                    Command::UnregisterPid(pid) => {
200                        if let Some(pid_data) = inner.pids.get_mut(&pid) {
201                            pid_data.unregistered = true;
202                            inner.maybe_remove_pid(epoll.as_fd(), pid);
203                        }
204                    }
205                    Command::Finish => {
206                        done = true;
207                        // There could be more commands queued and/or more epoll
208                        // events ready, but it doesn't matter. We don't
209                        // guarantee to callers whether callbacks have run or
210                        // not after having sent `Finish`; only that no more
211                        // callbacks will run after the thread is joined.
212                        break;
213                    }
214                }
215            }
216        }
217    }
218
219    /// Fork a child and register it. Uses `fork` internally; it `vfork` is desired,
220    /// use `register_pid` instead.
221    ///
222    /// Panics if `child_fn` returns.
223    /// TODO: change the type to `FnOnce() -> !` once that's stabilized in Rust.
224    /// <https://github.com/rust-lang/rust/issues/35121>
225    ///
226    /// # Safety
227    ///
228    /// As for fork in Rust in general. *Probably*, *mostly*, safe, since the
229    /// child process gets its own copy of the address space and OS resources etc.
230    /// Still, there may be some dragons here. Best to call exec before too long
231    /// in the child.
232    pub unsafe fn fork_watchable(&self, child_fn: impl FnOnce()) -> Result<Pid, Errno> {
233        let raw_pid = Errno::result_from_libc_errno(-1, unsafe { libc::syscall(libc::SYS_fork) })?;
234        if raw_pid == 0 {
235            child_fn();
236            panic!("child_fn shouldn't have returned");
237        }
238        let pid = Pid::from_raw(raw_pid.try_into().unwrap()).unwrap();
239        self.register_pid(pid);
240
241        Ok(pid)
242    }
243
244    /// Register interest in `pid`.
245    ///
246    /// Will succeed even if `pid` is already dead, in which case callbacks
247    /// registered for this `pid` will immediately be scheduled to run.
248    ///
249    /// `pid` must refer to some process, but that process may be a zombie (dead
250    /// but not yet reaped). Panics if `pid` doesn't exist at all.  The caller
251    /// should ensure the process has not been reaped before calling this
252    /// function both to avoid such panics, and to avoid accidentally watching
253    /// an unrelated process with a recycled `pid`.
254    pub fn register_pid(&self, pid: Pid) {
255        let mut inner = self.inner.lock().unwrap();
256        // We defensively make the pidfd non-blocking, since we intend to always
257        // use epoll to validate that it's ready before operating on it.
258        let pidfd = rustix::process::pidfd_open(pid.into(), PidfdFlags::NONBLOCK)
259            .unwrap_or_else(|e| panic!("pidfd_open failed for {pid:?}: {e:?}"));
260        // `pidfd_open(2)`: the close-on-exec flag is set on the file descriptor.
261        debug_assert!(
262            rustix::io::fcntl_getfd(&pidfd)
263                .unwrap()
264                .contains(FdFlags::CLOEXEC),
265            "pidfd_open unexpected didn't set CLOEXEC"
266        );
267        epoll::add(
268            &self.epoll,
269            &pidfd,
270            epoll::EventData::new_u64(pid.as_raw_nonzero().get().try_into().unwrap()),
271            epoll::EventFlags::IN,
272        )
273        .unwrap();
274
275        let prev = inner.pids.insert(
276            pid,
277            PidData {
278                callbacks: HashMap::new(),
279                pidfd: Some(pidfd),
280                unregistered: false,
281            },
282        );
283        assert!(prev.is_none());
284    }
285
286    // TODO: Re-enable when Rust supports vfork: https://github.com/rust-lang/rust/issues/58314
287    // pub unsafe fn vfork_watchable(&self, child_fn: impl FnOnce()) -> Result<Pid, nix::Error> {
288    //     unsafe { self.fork_watchable_internal(libc::SYS_vfork, child_fn) }
289    // }
290
291    /// Unregister the pid. After unregistration, no more callbacks may be
292    /// registered for the given pid. Already-registered callbacks will still be
293    /// called if and when the pid exits unless individually unregistered.
294    ///
295    /// Safe to call multiple times.
296    pub fn unregister_pid(&self, pid: Pid) {
297        // Let the worker handle the actual unregistration. This avoids a race
298        // where we unregister a pid at the same time as the worker thread
299        // receives an epoll event for it.
300        let mut inner = self.inner.lock().unwrap();
301        inner.send_command(Command::UnregisterPid(pid));
302    }
303
304    /// Call `callback` from another thread after the child `pid`
305    /// has exited, including if it has already exited. Does *not* reap the
306    /// child itself.
307    ///
308    /// The returned handle is guaranteed to be non-zero.
309    ///
310    /// Panics if `pid` isn't registered.
311    pub fn register_callback(
312        &self,
313        pid: Pid,
314        callback: impl Send + FnOnce(Pid) + 'static,
315    ) -> WatchHandle {
316        let mut inner = self.inner.lock().unwrap();
317        let handle = inner.next_handle;
318        inner.next_handle += 1;
319        let pid_data = inner.pids.get_mut(&pid).unwrap();
320        assert!(!pid_data.unregistered);
321        pid_data.callbacks.insert(handle, Box::new(callback));
322        if pid_data.pidfd.is_none() {
323            // pid is already dead. Run the callback we just registered.
324            inner.send_command(Command::RunCallbacks(pid));
325        }
326        handle
327    }
328
329    /// Unregisters a callback. After returning, the corresponding callback is
330    /// guaranteed either to have already run, or to never run. i.e. it's safe to
331    /// free data that the callback might otherwise access.
332    ///
333    /// No-op if `pid` isn't registered.
334    pub fn unregister_callback(&self, pid: Pid, handle: WatchHandle) {
335        let mut inner = self.inner.lock().unwrap();
336        if let Some(pid_data) = inner.pids.get_mut(&pid) {
337            pid_data.callbacks.remove(&handle);
338            inner.maybe_remove_pid(&self.epoll, pid);
339        }
340    }
341}
342
343impl Default for ChildPidWatcher {
344    fn default() -> Self {
345        Self::new()
346    }
347}
348
349impl Drop for ChildPidWatcher {
350    fn drop(&mut self) {
351        let handle = {
352            let mut inner = self.inner.lock().unwrap();
353            inner.send_command(Command::Finish);
354            inner.thread_handle.take().unwrap()
355        };
356        handle.join().unwrap();
357    }
358}
359
360impl std::fmt::Debug for PidData {
361    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
362        f.debug_struct("PidData")
363            .field("fd", &self.pidfd)
364            .field("unregistered", &self.unregistered)
365            .finish_non_exhaustive()
366    }
367}
368
369#[cfg(test)]
370mod tests {
371    use std::sync::{Arc, Condvar};
372
373    use nix::sys::eventfd::EventFd;
374    use rustix::fd::AsRawFd;
375    use rustix::process::{WaitOptions, waitpid};
376
377    use super::*;
378
379    fn is_zombie(pid: Pid) -> bool {
380        let stat_name = format!("/proc/{}/stat", pid.as_raw_nonzero().get());
381        let contents = std::fs::read_to_string(stat_name).unwrap();
382        contents.contains(") Z")
383    }
384
385    #[test]
386    // can't call foreign function: pipe
387    #[cfg_attr(miri, ignore)]
388    fn register_before_exit() {
389        let notifier = EventFd::new().unwrap();
390
391        let watcher = ChildPidWatcher::new();
392        let child = unsafe {
393            watcher.fork_watchable(|| {
394                let mut buf = [0; 8];
395                // Wait for parent to register its callback.
396                nix::unistd::read(notifier.as_raw_fd(), &mut buf).unwrap();
397                libc::_exit(42);
398            })
399        }
400        .unwrap();
401
402        let callback_ran = Arc::new((Mutex::new(false), Condvar::new()));
403        {
404            let callback_ran = callback_ran.clone();
405            watcher.register_callback(
406                child,
407                Box::new(move |pid| {
408                    assert_eq!(pid, child);
409                    *callback_ran.0.lock().unwrap() = true;
410                    callback_ran.1.notify_all();
411                }),
412            );
413        }
414
415        // Should be safe to unregister the pid now.
416        // We don't be able to register any more callbacks, but existing one
417        // should still work.
418        watcher.unregister_pid(child);
419
420        // Child should still be alive.
421        let status = waitpid(Some(child.into()), WaitOptions::NOHANG).unwrap();
422        assert!(status.is_none(), "Unexpected status: {status:?}");
423
424        // Callback shouldn't have run yet.
425        assert!(!*callback_ran.0.lock().unwrap());
426
427        // Let the child exit.
428        nix::unistd::write(&notifier, &1u64.to_ne_bytes()).unwrap();
429
430        // Wait for our callback to run.
431        let mut callback_ran_lock = callback_ran.0.lock().unwrap();
432        while !*callback_ran_lock {
433            callback_ran_lock = callback_ran.1.wait(callback_ran_lock).unwrap();
434        }
435
436        // Child should be ready to be reaped.
437        // TODO: use WNOHANG here if we go back to a pidfd-based implementation.
438        // With the current fd-based implementation we may be notified before kernel
439        // marks the child reapable.
440        let status = waitpid(Some(child.into()), WaitOptions::empty())
441            .unwrap()
442            .unwrap();
443        assert_eq!(status.1.exit_status(), Some(42));
444    }
445
446    #[test]
447    // can't call foreign functions
448    #[cfg_attr(miri, ignore)]
449    fn register_after_exit() {
450        let child = match unsafe { libc::fork() } {
451            0 => {
452                unsafe { libc::_exit(42) };
453            }
454            child => Pid::from_raw(child).unwrap(),
455        };
456
457        // Wait until child is dead, but don't reap it yet.
458        while !is_zombie(child) {
459            unsafe {
460                libc::sched_yield();
461            }
462        }
463
464        let watcher = ChildPidWatcher::new();
465        watcher.register_pid(child);
466
467        // Used to wait until after the ChildPidWatcher has ran our callback
468        let callback_ran = Arc::new((Mutex::new(false), Condvar::new()));
469        {
470            let callback_ran = callback_ran.clone();
471            watcher.register_callback(
472                child,
473                Box::new(move |pid| {
474                    assert_eq!(pid, child);
475                    *callback_ran.0.lock().unwrap() = true;
476                    callback_ran.1.notify_all();
477                }),
478            );
479        }
480
481        // Should be safe to unregister the pid now.
482        // We don't be able to register any more callbacks, but existing one
483        // should still work.
484        watcher.unregister_pid(child);
485
486        // Wait for our callback to run.
487        let mut callback_ran_lock = callback_ran.0.lock().unwrap();
488        while !*callback_ran_lock {
489            callback_ran_lock = callback_ran.1.wait(callback_ran_lock).unwrap();
490        }
491
492        // Child should be ready to be reaped.
493        // TODO: use WNOHANG here if we go back to a pidfd-based implementation.
494        // With the current fd-based implementation we may be notified before kernel
495        // marks the child reapable.
496        assert_eq!(
497            waitpid(Some(child.into()), WaitOptions::empty())
498                .unwrap()
499                .unwrap()
500                .1
501                .exit_status(),
502            Some(42)
503        );
504    }
505
506    #[test]
507    // can't call foreign function: pipe
508    #[cfg_attr(miri, ignore)]
509    fn register_multiple() {
510        let cb1_ran = Arc::new((Mutex::new(false), Condvar::new()));
511        let cb2_ran = Arc::new((Mutex::new(false), Condvar::new()));
512
513        let watcher = ChildPidWatcher::new();
514        let child = unsafe {
515            watcher.fork_watchable(|| {
516                libc::_exit(42);
517            })
518        }
519        .unwrap();
520
521        for cb_ran in vec![cb1_ran.clone(), cb2_ran.clone()].drain(..) {
522            let cb_ran = cb_ran.clone();
523            watcher.register_callback(
524                child,
525                Box::new(move |pid| {
526                    assert_eq!(pid, child);
527                    *cb_ran.0.lock().unwrap() = true;
528                    cb_ran.1.notify_all();
529                }),
530            );
531        }
532
533        // Should be safe to unregister the pid now.
534        // We don't be able to register any more callbacks, but existing one
535        // should still work.
536        watcher.unregister_pid(child);
537
538        for cb_ran in vec![cb1_ran, cb2_ran].drain(..) {
539            let mut cb_ran_lock = cb_ran.0.lock().unwrap();
540            while !*cb_ran_lock {
541                cb_ran_lock = cb_ran.1.wait(cb_ran_lock).unwrap();
542            }
543        }
544
545        // Child should be ready to be reaped.
546        // TODO: use WNOHANG here if we go back to a pidfd-based implementation.
547        // With the current fd-based implementation we may be notified before kernel
548        // marks the child reapable.
549        assert_eq!(
550            waitpid(Some(child.into()), WaitOptions::empty())
551                .unwrap()
552                .unwrap()
553                .1
554                .exit_status(),
555            Some(42)
556        );
557    }
558
559    #[test]
560    // can't call foreign function
561    #[cfg_attr(miri, ignore)]
562    fn unregister_one() {
563        let cb1_ran = Arc::new((Mutex::new(false), Condvar::new()));
564        let cb2_ran = Arc::new((Mutex::new(false), Condvar::new()));
565
566        let notifier = EventFd::new().unwrap();
567
568        let watcher = ChildPidWatcher::new();
569        let child = unsafe {
570            watcher.fork_watchable(|| {
571                let mut buf = [0; 8];
572                // Wait for parent to register its callback.
573                nix::unistd::read(notifier.as_raw_fd(), &mut buf).unwrap();
574                libc::_exit(42);
575            })
576        }
577        .unwrap();
578
579        let handles: Vec<WatchHandle> = [&cb1_ran, &cb2_ran]
580            .iter()
581            .cloned()
582            .map(|cb_ran| {
583                let cb_ran = cb_ran.clone();
584                watcher.register_callback(
585                    child,
586                    Box::new(move |pid| {
587                        assert_eq!(pid, child);
588                        *cb_ran.0.lock().unwrap() = true;
589                        cb_ran.1.notify_all();
590                    }),
591                )
592            })
593            .collect();
594
595        // Should be safe to unregister the pid now.
596        // We don't be able to register any more callbacks, but existing one
597        // should still work.
598        watcher.unregister_pid(child);
599
600        watcher.unregister_callback(child, handles[0]);
601
602        // Let the child exit.
603        nix::unistd::write(&notifier, &1u64.to_ne_bytes()).unwrap();
604
605        // Wait for the still-registered callback to run.
606        let mut cb_ran_lock = cb2_ran.0.lock().unwrap();
607        while !*cb_ran_lock {
608            cb_ran_lock = cb2_ran.1.wait(cb_ran_lock).unwrap();
609        }
610
611        // The unregistered cb should *not* have run.
612        assert!(!*cb1_ran.0.lock().unwrap());
613
614        // Child should be ready to be reaped.
615        // TODO: use WNOHANG here if we go back to a pidfd-based implementation.
616        // With the current fd-based implementation we may be notified before kernel
617        // marks the child reapable.
618        assert_eq!(
619            waitpid(Some(child.into()), WaitOptions::empty())
620                .unwrap()
621                .unwrap()
622                .1
623                .exit_status(),
624            Some(42)
625        );
626    }
627}