use std::{
collections::{HashMap, VecDeque},
fmt,
io::{self, Read},
os::fd::AsFd,
process::{Child, ChildStderr, ChildStdin, ChildStdout, Command, ExitStatus},
time::{Duration, Instant},
};
mod sys;
pub type ProcId = u64;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PsetEvent {
Stdout(Vec<u8>),
Stderr(Vec<u8>),
ProcessExited(ExitStatus),
}
pub struct ProcessSet<T> {
poller: sys::Poller,
procs: HashMap<u64, Proc<T>>,
pending: VecDeque<(T, PsetEvent)>,
tokens: Vec<u64>,
buf: Vec<u8>,
next_id: u64,
}
struct Proc<T> {
tag: T,
child: Child,
exit_watch: sys::ExitWatch,
stdout: Option<ChildStdout>,
stderr: Option<ChildStderr>,
status: Option<ExitStatus>,
}
impl<T> Proc<T> {
fn finished(&self) -> bool {
self.status.is_some() && self.stdout.is_none() && self.stderr.is_none()
}
}
const READ_CHUNK: usize = 64 * 1024;
const STREAM_STDOUT: u64 = 0;
const STREAM_STDERR: u64 = 1;
const STREAM_EXIT: u64 = 2;
fn token(id: u64, stream: u64) -> u64 {
(id << 2) | stream
}
fn untoken(token: u64) -> (u64, u64) {
(token >> 2, token & 0b11)
}
pub fn create<T: Clone>() -> io::Result<ProcessSet<T>> {
ProcessSet::new()
}
impl<T: Clone> ProcessSet<T> {
pub fn new() -> io::Result<Self> {
Ok(Self {
poller: sys::Poller::new()?,
procs: HashMap::new(),
pending: VecDeque::new(),
tokens: Vec::new(),
buf: vec![0; READ_CHUNK],
next_id: 0,
})
}
pub fn len(&self) -> usize {
self.procs.len()
}
pub fn is_empty(&self) -> bool {
self.procs.is_empty() && self.pending.is_empty()
}
pub fn kill(&self, id: ProcId, signal: i32) -> io::Result<()> {
let proc = self
.procs
.get(&id)
.ok_or_else(|| io::Error::new(io::ErrorKind::NotFound, "no such process in the set"))?;
if proc.status.is_some() {
return Ok(());
}
proc.exit_watch.send_signal(signal)
}
pub fn kill_all(&self, signal: i32) -> io::Result<()> {
let mut first_err = None;
for proc in self.procs.values() {
if proc.status.is_some() {
continue;
}
if let Err(e) = proc.exit_watch.send_signal(signal) {
first_err.get_or_insert(e);
}
}
match first_err {
Some(e) => Err(e),
None => Ok(()),
}
}
pub fn spawn(
&mut self,
tag: T,
mut command: Command,
) -> io::Result<(ProcId, Option<ChildStdin>)> {
let mut child = command.spawn()?;
let stdin = child.stdin.take();
let id = self.next_id;
match self.attach(id, &mut child) {
Ok((exit_watch, stdout, stderr)) => {
self.next_id += 1;
self.procs.insert(
id,
Proc {
tag,
child,
exit_watch,
stdout,
stderr,
status: None,
},
);
Ok((id, stdin))
}
Err(e) => {
let _ = child.kill();
let _ = child.wait();
Err(e)
}
}
}
fn attach(
&mut self,
id: u64,
child: &mut Child,
) -> io::Result<(sys::ExitWatch, Option<ChildStdout>, Option<ChildStderr>)> {
let stdout = child.stdout.take();
let stderr = child.stderr.take();
if let Some(stdout) = stdout.as_ref() {
sys::set_nonblocking(stdout.as_fd())?;
self.poller
.watch_read(stdout.as_fd(), token(id, STREAM_STDOUT))?;
}
if let Some(stderr) = stderr.as_ref() {
sys::set_nonblocking(stderr.as_fd())?;
self.poller
.watch_read(stderr.as_fd(), token(id, STREAM_STDERR))?;
}
let exit_watch = self
.poller
.watch_exit(child.id() as libc::pid_t, token(id, STREAM_EXIT))?;
Ok((exit_watch, stdout, stderr))
}
pub fn wait_next(&mut self) -> io::Result<Option<(T, PsetEvent)>> {
self.next_event(None)
}
pub fn wait_next_timeout(&mut self, timeout: Duration) -> io::Result<Option<(T, PsetEvent)>> {
self.next_event(Some(Instant::now() + timeout))
}
fn next_event(&mut self, deadline: Option<Instant>) -> io::Result<Option<(T, PsetEvent)>> {
loop {
if let Some(event) = self.pending.pop_front() {
return Ok(Some(event));
}
if self.procs.is_empty() {
return Ok(None);
}
let timeout = deadline.map(|d| d.saturating_duration_since(Instant::now()));
let n = match self.poller.wait(&mut self.tokens, timeout) {
Ok(n) => n,
Err(e) if e.kind() == io::ErrorKind::Interrupted => continue,
Err(e) => return Err(e),
};
if n == 0 {
return Ok(None); }
self.dispatch(n)?;
}
}
fn dispatch(&mut self, n: usize) -> io::Result<()> {
for i in 0..n {
let (id, stream) = untoken(self.tokens[i]);
match stream {
STREAM_STDOUT | STREAM_STDERR => self.read_stream(id, stream)?,
STREAM_EXIT => self.reap(id)?,
_ => unreachable!("token carries one of three streams"),
}
}
self.collect_finished();
Ok(())
}
fn read_stream(&mut self, id: u64, stream: u64) -> io::Result<()> {
let Some(proc) = self.procs.get_mut(&id) else {
return Ok(());
};
let reader: Option<&mut dyn Read> = match stream {
STREAM_STDOUT => match proc.stdout.as_mut() {
Some(stdout) => Some(stdout),
None => None,
},
_ => match proc.stderr.as_mut() {
Some(stderr) => Some(stderr),
None => None,
},
};
let Some(reader) = reader else {
return Ok(());
};
let read = match reader.read(&mut self.buf) {
Ok(read) => read,
Err(e) if e.kind() == io::ErrorKind::WouldBlock => return Ok(()),
Err(e) if e.kind() == io::ErrorKind::Interrupted => return Ok(()),
Err(_) => 0,
};
if read == 0 {
match stream {
STREAM_STDOUT => {
let stdout = proc.stdout.take().expect("checked above");
self.poller.unwatch_read(stdout.as_fd())?;
}
_ => {
let stderr = proc.stderr.take().expect("checked above");
self.poller.unwatch_read(stderr.as_fd())?;
}
}
return Ok(());
}
let data = self.buf[..read].to_vec();
let event = match stream {
STREAM_STDOUT => PsetEvent::Stdout(data),
_ => PsetEvent::Stderr(data),
};
self.pending.push_back((proc.tag.clone(), event));
Ok(())
}
fn reap(&mut self, id: u64) -> io::Result<()> {
let Some(proc) = self.procs.get_mut(&id) else {
return Ok(());
};
if proc.status.is_some() {
return Ok(());
}
self.poller.unwatch_exit(&proc.exit_watch)?;
proc.status = Some(proc.child.wait()?);
Ok(())
}
fn collect_finished(&mut self) {
let mut done: Vec<u64> = self
.procs
.iter()
.filter(|(_, proc)| proc.finished())
.map(|(id, _)| *id)
.collect();
done.sort_unstable();
for id in done {
let proc = self.procs.remove(&id).expect("just listed");
let status = proc.status.expect("finished implies reaped");
self.pending
.push_back((proc.tag, PsetEvent::ProcessExited(status)));
}
}
}
impl<T> fmt::Debug for ProcessSet<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ProcessSet")
.field("processes", &self.procs.len())
.field("pending_events", &self.pending.len())
.finish()
}
}
impl<T> Drop for ProcessSet<T> {
fn drop(&mut self) {
for proc in self.procs.values_mut() {
if proc.status.is_some() {
continue;
}
let _ = proc.exit_watch.send_signal(libc::SIGKILL);
let _ = proc.child.wait();
}
}
}
macro_rules! cvt {
($e:expr) => {{
let n = $e;
if n < 0 {
return Err(io::Error::last_os_error());
}
n
}};
}
pub(crate) use cvt;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn token_round_trips() {
for id in [0u64, 1, 7, 4096, u64::MAX >> 2] {
for stream in [STREAM_STDOUT, STREAM_STDERR, STREAM_EXIT] {
assert_eq!(untoken(token(id, stream)), (id, stream));
}
}
}
#[test]
fn tokens_of_one_process_differ() {
let id = 3;
let tokens = [
token(id, STREAM_STDOUT),
token(id, STREAM_STDERR),
token(id, STREAM_EXIT),
];
assert_eq!(
tokens
.iter()
.collect::<std::collections::HashSet<_>>()
.len(),
3
);
}
#[test]
fn a_fresh_set_is_empty() {
let set = create::<()>().unwrap();
assert!(set.is_empty());
assert_eq!(set.len(), 0);
}
#[test]
fn waiting_on_an_empty_set_returns_nothing() {
let mut set = create::<()>().unwrap();
assert!(set.wait_next().unwrap().is_none());
}
#[test]
fn killing_an_unknown_process_is_not_found() {
let set = create::<()>().unwrap();
let err = set.kill(42, libc::SIGTERM).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::NotFound);
}
}