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#[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 callbacks: HashMap<WatchHandle, Box<dyn Send + FnOnce(Pid)>>,
35 pidfd: Option<OwnedFd>,
37 unregistered: bool,
40}
41
42#[derive(Debug)]
43struct Inner {
44 next_handle: WatchHandle,
46 commands: Vec<Command>,
48 pids: HashMap<Pid, PidData>,
50 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 return;
66 };
67 let Some(fd) = piddata.pidfd.take() else {
68 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 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 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 continue;
156 }
157 Err(e) => panic!("epoll_wait: {e:?}"),
158 };
159
160 let mut inner = inner.lock().unwrap();
166
167 for event in events.drain(..) {
168 if event.data.u64() == 0 {
169 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 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 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 break;
213 }
214 }
215 }
216 }
217 }
218
219 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 pub fn register_pid(&self, pid: Pid) {
255 let mut inner = self.inner.lock().unwrap();
256 let pidfd = rustix::process::pidfd_open(pid.into(), PidfdFlags::NONBLOCK)
259 .unwrap_or_else(|e| panic!("pidfd_open failed for {pid:?}: {e:?}"));
260 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 pub fn unregister_pid(&self, pid: Pid) {
297 let mut inner = self.inner.lock().unwrap();
301 inner.send_command(Command::UnregisterPid(pid));
302 }
303
304 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 inner.send_command(Command::RunCallbacks(pid));
325 }
326 handle
327 }
328
329 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 #[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 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 watcher.unregister_pid(child);
419
420 let status = waitpid(Some(child.into()), WaitOptions::NOHANG).unwrap();
422 assert!(status.is_none(), "Unexpected status: {status:?}");
423
424 assert!(!*callback_ran.0.lock().unwrap());
426
427 nix::unistd::write(¬ifier, &1u64.to_ne_bytes()).unwrap();
429
430 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 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 #[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 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 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 watcher.unregister_pid(child);
485
486 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 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 #[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 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 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 #[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 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 watcher.unregister_pid(child);
599
600 watcher.unregister_callback(child, handles[0]);
601
602 nix::unistd::write(¬ifier, &1u64.to_ne_bytes()).unwrap();
604
605 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 assert!(!*cb1_ran.0.lock().unwrap());
613
614 assert_eq!(
619 waitpid(Some(child.into()), WaitOptions::empty())
620 .unwrap()
621 .unwrap()
622 .1
623 .exit_status(),
624 Some(42)
625 );
626 }
627}