use core::fmt;
use core::future::Future;
use core::pin::Pin;
use core::task::Context;
use core::task::Poll;
use std::io;
use nix::sys::signal::Signal;
use serde::Deserialize;
use serde::Serialize;
use syscalls::Errno;
use super::Command;
use super::ExitStatus;
use super::Pid;
use super::seccomp::SeccompNotif;
use super::stdio::ChildStderr;
use super::stdio::ChildStdin;
use super::stdio::ChildStdout;
use super::stdio::Stdio;
#[derive(Debug)]
pub struct Child {
pub(super) pid: Pid,
pub(super) exit_status: Option<ExitStatus>,
pub seccomp_notif: Option<SeccompNotif>,
pub stdin: Option<ChildStdin>,
pub stdout: Option<ChildStdout>,
pub stderr: Option<ChildStderr>,
}
#[derive(PartialEq, Eq, Clone, Serialize, Deserialize)]
pub struct Output {
pub status: ExitStatus,
pub stdout: Vec<u8>,
pub stderr: Vec<u8>,
}
impl fmt::Debug for Output {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let stdout = core::str::from_utf8(&self.stdout);
let stdout: &dyn fmt::Debug = match stdout {
Ok(ref s) => s,
Err(_) => &self.stdout,
};
let stderr = core::str::from_utf8(&self.stderr);
let stderr: &dyn fmt::Debug = match stderr {
Ok(ref s) => s,
Err(_) => &self.stderr,
};
f.debug_struct("Output")
.field("status", &self.status)
.field("stdout", stdout)
.field("stderr", stderr)
.finish()
}
}
impl Child {
pub fn id(&self) -> Pid {
self.pid
}
pub fn try_wait(&mut self) -> io::Result<Option<ExitStatus>> {
match self.exit_status {
Some(exit_status) => Ok(Some(exit_status)),
None => {
let mut status = 0;
let ret = Errno::result(unsafe {
libc::waitpid(self.pid.as_raw(), &mut status, libc::WNOHANG)
})?;
if ret == 0 {
Ok(None)
} else {
let exit_status = ExitStatus::from_raw(status);
self.exit_status = Some(exit_status);
Ok(Some(exit_status))
}
}
}
}
pub async fn wait(&mut self) -> io::Result<ExitStatus> {
drop(self.stdin.take());
WaitForChild::new(self)?.await
}
pub fn wait_blocking(&mut self) -> io::Result<ExitStatus> {
drop(self.stdin.take());
let mut status = 0;
let ret = loop {
match Errno::result(unsafe { libc::waitpid(self.pid.as_raw(), &mut status, 0) }) {
Ok(ret) => break ret,
Err(Errno::EINTR) => continue,
Err(err) => return Err(err.into()),
}
};
debug_assert_ne!(ret, 0);
Ok(ExitStatus::from_raw(status))
}
pub async fn wait_with_output(mut self) -> io::Result<Output> {
use futures::future::try_join3;
use tokio::io::AsyncRead;
use tokio::io::AsyncReadExt;
async fn read_to_end<A: AsyncRead + Unpin>(io: Option<A>) -> io::Result<Vec<u8>> {
let mut vec = Vec::new();
if let Some(mut io) = io {
io.read_to_end(&mut vec).await?;
}
Ok(vec)
}
let stdout_fut = read_to_end(self.stdout.take());
let stderr_fut = read_to_end(self.stderr.take());
let (status, stdout, stderr) = try_join3(self.wait(), stdout_fut, stderr_fut).await?;
Ok(Output {
status,
stdout,
stderr,
})
}
pub fn signal(&self, sig: Signal) -> io::Result<()> {
if self.exit_status.is_none() {
Errno::result(unsafe { libc::kill(self.pid.as_raw(), sig as i32) })?;
}
Ok(())
}
}
impl Command {
pub async fn status(&mut self) -> io::Result<ExitStatus> {
let mut child = self.spawn()?;
drop(child.stdin.take());
drop(child.stdout.take());
drop(child.stderr.take());
child.wait().await
}
pub async fn output(&mut self) -> io::Result<Output> {
self.stdout(Stdio::piped());
self.stderr(Stdio::piped());
let child = self.spawn();
child?.wait_with_output().await
}
}
struct WaitForChild<'a> {
signal: tokio::signal::unix::Signal,
child: &'a mut Child,
}
impl<'a> WaitForChild<'a> {
fn new(child: &'a mut Child) -> io::Result<Self> {
use tokio::signal::unix::SignalKind;
use tokio::signal::unix::signal;
Ok(Self {
signal: signal(SignalKind::child())?,
child,
})
}
}
impl<'a> Future for WaitForChild<'a> {
type Output = io::Result<ExitStatus>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
loop {
let sig = self.signal.poll_recv(cx);
if let Some(status) = self.child.try_wait()? {
return Poll::Ready(Ok(status));
}
if sig.is_pending() {
return Poll::Pending;
}
}
}
}