use std::collections::HashMap;
use std::collections::HashSet;
use std::os::fd::AsRawFd;
use std::os::fd::BorrowedFd;
use std::process;
use nix::errno::Errno;
use nix::poll::PollFd;
use nix::poll::PollFlags;
use nix::poll::PollTimeout;
use nix::poll::{self};
use ptools::proc::pidfd::PidFd;
struct Args {
verbose: bool,
pid: Vec<u64>,
}
fn print_usage() {
eprintln!("Usage: pwait [-v] PID...");
eprintln!("Wait for processes to terminate. A /proc/pid path may also be used.");
eprintln!();
eprintln!("Options:");
eprintln!(" -v Report terminations to standard output");
eprintln!(" -h, --help Print help");
eprintln!(" -V, --version Print version");
}
fn parse_args() -> Args {
use lexopt::prelude::*;
let mut args = Args {
verbose: false,
pid: Vec::new(),
};
let mut parser = lexopt::Parser::from_env();
while let Some(arg) = parser.next().unwrap_or_else(|e| {
eprintln!("pwait: {e}");
process::exit(2);
}) {
match arg {
Short('h') | Long("help") => {
print_usage();
process::exit(0);
}
Short('V') | Long("version") => {
println!("pwait {}", env!("CARGO_PKG_VERSION"));
process::exit(0);
}
Short('v') => args.verbose = true,
Value(val) => {
let s = val.to_string_lossy();
match ptools::proc::parse_pid_arg(&s) {
Ok(ptools::proc::PidArg::Pid(pid)) => args.pid.push(pid),
Ok(ptools::proc::PidArg::Skip) => {}
Err(msg) => {
eprintln!("pwait: {msg}");
process::exit(2);
}
}
}
_ => {
eprintln!("pwait: unexpected argument: {arg:?}");
process::exit(2);
}
}
}
if args.pid.is_empty() {
eprintln!("pwait: at least one PID required");
process::exit(2);
}
args
}
fn main() {
ptools::reset_sigpipe();
let args = parse_args();
let mut failed = false;
let mut seen = HashSet::new();
let pids: Vec<u64> = args.pid.into_iter().filter(|p| seen.insert(*p)).collect();
let mut entries: HashMap<i32, (PidFd, u64)> = HashMap::new();
let my_pid = process::id() as u64;
for pid in &pids {
if *pid == my_pid {
eprintln!("pwait: skipping self PID {pid}");
failed = true;
continue;
}
match PidFd::open(*pid) {
Ok(fd) => {
entries.insert(fd.as_raw_fd(), (fd, *pid));
}
Err(e) => {
eprintln!("pwait: failed to open pidfd for {pid}: {e}");
failed = true;
}
}
}
while !entries.is_empty() {
let raw_fds: Vec<i32> = entries.keys().copied().collect();
let mut pollfds: Vec<PollFd> = raw_fds
.iter()
.map(|&raw_fd| {
PollFd::new(unsafe { BorrowedFd::borrow_raw(raw_fd) }, PollFlags::POLLIN)
})
.collect();
match poll::poll(&mut pollfds, PollTimeout::NONE) {
Err(Errno::EINTR) => continue,
Err(e) => {
eprintln!("pwait: poll: {e}");
process::exit(1);
}
Ok(_) => {}
}
let ready_fds: HashSet<i32> = pollfds
.iter()
.zip(raw_fds.iter())
.filter(|(pfd, _)| {
pfd.revents().is_some_and(|r| {
r.intersects(PollFlags::POLLIN | PollFlags::POLLHUP | PollFlags::POLLERR)
})
})
.map(|(_, &raw_fd)| raw_fd)
.collect();
for &raw_fd in &ready_fds {
let (pidfd, pid) = &entries[&raw_fd];
if args.verbose {
match pidfd.wait_status() {
Ok(status) => {
println!("{pid}: terminated, wait status {status:#06x}");
}
Err(Errno::ECHILD) => {
println!("{pid}: terminated");
}
Err(e) => {
eprintln!("pwait: waitid for {pid}: {e}");
println!("{pid}: terminated");
}
}
}
entries.remove(&raw_fd);
}
}
if failed {
process::exit(1);
}
}