use std::{collections::HashMap, usize};
use libc::{EPOLLIN, EPOLL_CTL_ADD, EPOLL_CTL_DEL};
type FD = i32;
type PID = u32;
type FDPidsMap = HashMap<PID, FD>;
pub struct PidSet {
fd_pids: FDPidsMap,
epoll_fd: Option<FD>,
}
#[derive(Debug, thiserror::Error)]
pub enum PidSetError {
#[error("Error while creating epoll file instance:`{0}`")]
EpollCreate(std::io::Error),
#[error("Error on pidfd_open syscall for pid `{0}`: `{1}")]
PidFdOpenSyscall(u32, std::io::Error),
#[error("Error on epoll_ctl: `{0}")]
EpollCtl(std::io::Error),
#[error("Error on epoll_wait: `{0}")]
EpollWait(std::io::Error),
#[error("PID not found: `{0}")]
PidNotFound(u32),
#[error("Error while closing epoll file descriptor: `{0}")]
EpollClose(std::io::Error),
}
impl PidSet {
pub fn new<P: IntoIterator<Item = PID>>(pids: P) -> Self {
let fd_pids: FDPidsMap = pids.into_iter().map(|pid| (pid, 0)).collect();
Self {
fd_pids,
epoll_fd: None,
}
}
fn register_pid(epoll_fd: i32, pid: u32, token: u32) -> Result<FD, PidSetError> {
let cfd = unsafe { syscallerr(libc::syscall(libc::SYS_pidfd_open, pid, 0)) }
.map_err(|err| PidSetError::PidFdOpenSyscall(pid, err))?;
unsafe {
syserr(libc::epoll_ctl(
epoll_fd,
EPOLL_CTL_ADD,
cfd as i32,
&mut libc::epoll_event {
events: EPOLLIN as u32,
u64: token as u64,
} as *mut _ as *mut libc::epoll_event,
))
}
.map_err(PidSetError::EpollCtl)?;
Ok(cfd as i32)
}
fn deregister_pid(epoll_fd: i32, fd: i32) -> Result<(), PidSetError> {
let _ = unsafe {
syserr(libc::epoll_ctl(
epoll_fd,
EPOLL_CTL_DEL,
fd,
std::ptr::null_mut(),
))
}
.map_err(PidSetError::EpollWait)?;
Ok(())
}
fn init_epoll(&mut self) -> Result<FD, PidSetError> {
let epoll_fd =
unsafe { syserr(libc::epoll_create1(0)) }.map_err(PidSetError::EpollCreate)?;
for (pid, fd) in &mut self.fd_pids {
*fd = PidSet::register_pid(epoll_fd, *pid, *pid)?;
}
self.epoll_fd = Some(epoll_fd);
Ok(epoll_fd)
}
}
fn syserr(status_code: libc::c_int) -> std::io::Result<libc::c_int> {
if status_code < 0 {
return Err(std::io::Error::from_raw_os_error(status_code));
}
Ok(status_code)
}
fn syscallerr(status_code: libc::c_long) -> std::io::Result<libc::c_long> {
if status_code < 0 {
return Err(std::io::Error::last_os_error());
}
Ok(status_code)
}
impl PidSet {
fn wait(&mut self, n: usize) -> Result<usize, PidSetError> {
let max_events = self.fd_pids.len();
let mut total_events: usize = 0;
let epoll_fd = self.epoll_fd.unwrap_or(self.init_epoll()?);
while total_events < n {
let mut events: Vec<libc::epoll_event> = Vec::with_capacity(max_events);
let event_count = syserr(unsafe {
libc::epoll_wait(epoll_fd, events.as_mut_ptr(), max_events as i32, -1)
})
.map_err(PidSetError::EpollWait)? as usize;
unsafe { events.set_len(event_count as usize) };
total_events += event_count;
for event in events {
let cdata = event.u64 as u32;
let fd = self
.fd_pids
.get(&cdata)
.ok_or(PidSetError::PidNotFound(cdata))?;
PidSet::deregister_pid(epoll_fd, *fd)?;
self.fd_pids.remove(&cdata);
}
}
Ok(total_events)
}
pub fn wait_all(&mut self) -> Result<(), PidSetError> {
self.wait(self.fd_pids.len())?;
Ok(())
}
pub fn wait_any(&mut self) -> Result<(), PidSetError> {
self.wait(1)?;
Ok(())
}
pub fn close(mut self) -> Result<(), PidSetError> {
let epoll_fd = self.epoll_fd.unwrap_or(self.init_epoll()?);
unsafe { syserr(libc::close(epoll_fd)) }.map_err(PidSetError::EpollClose)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::{Duration, Instant};
fn sleep_cmd(duration: &str) -> std::process::Command {
let mut cmd1 = std::process::Command::new("sleep");
cmd1.arg(duration);
cmd1
}
#[test]
fn wait_all() {
let mut pid_set = PidSet::new([
sleep_cmd("0.1").spawn().unwrap().id(),
sleep_cmd("0.2").spawn().unwrap().id(),
sleep_cmd("0.3").spawn().unwrap().id(),
sleep_cmd("0.4").spawn().unwrap().id(),
sleep_cmd("0.5").spawn().unwrap().id(),
]);
assert!(pid_set.wait_all().is_ok());
}
#[test]
fn wait_any() {
let start_time = Instant::now();
let mut pid_set = PidSet::new([
sleep_cmd("0.2").spawn().unwrap().id(),
sleep_cmd("3").spawn().unwrap().id(),
sleep_cmd("3").spawn().unwrap().id(),
sleep_cmd("3").spawn().unwrap().id(),
sleep_cmd("3").spawn().unwrap().id(),
]);
assert!(pid_set.wait_any().is_ok());
assert!(
start_time.elapsed() < Duration::from_secs(3),
"Expected wait_any() to return in less than 3 seconds, but it took {:?}",
start_time.elapsed()
);
}
}