use async_stream::stream;
use color_eyre::{eyre::bail, Result};
use log::debug;
use nix::pty::openpty;
use nix::sys::signal::Signal;
use nix::sys::wait::{waitpid, WaitStatus};
use nix::unistd::dup2;
use nix::unistd::execve;
use nix::unistd::ForkResult;
use nix::unistd::{close, fork, Pid};
use std::env;
use std::ffi::CString;
use std::os::fd::AsRawFd;
use std::os::unix::prelude::RawFd;
use std::path::Path;
use std::pin::Pin;
use tokio::io::AsyncBufReadExt;
use tokio::io::BufReader;
use tokio_fd::AsyncFd;
use tokio_stream::wrappers::LinesStream;
use tokio_stream::{Stream, StreamExt, StreamMap};
use which::which;
fn str2cstring(s: &str) -> Result<CString> {
if let Ok(c) = CString::new(s) {
return Ok(c);
}
bail!("Unable to convert str '{}' to cstring", s);
}
fn vec_slice_of_string_2_vec_of_cstring(input: &[String]) -> Result<Vec<CString>> {
let output = input.iter().map(|s| str2cstring(s).unwrap()).collect();
Ok(output)
}
fn path2cstring(p: &Path) -> Result<CString> {
if let Some(s) = p.to_str() {
if let Ok(c) = str2cstring(s) {
return Ok(c);
};
};
bail!("Unable to convert pathbuf '{}' to str", p.display());
}
#[derive(Debug, Eq, PartialEq, Hash, Copy, Clone)]
pub enum StdioType {
Stdout,
Stderr,
}
pub enum XStatus {
Exited(i32),
Signaled(Signal),
}
pub struct XChildHandle {
pid: Pid,
stdout_raw_fd: RawFd,
stderr_raw_fd: RawFd,
}
impl XChildHandle {
fn new(pid: Pid, stdout_raw_fd: RawFd, stderr_raw_fd: RawFd) -> Result<Self> {
Ok(XChildHandle {
pid,
stdout_raw_fd,
stderr_raw_fd,
})
}
pub fn pid(&self) -> Pid {
self.pid
}
pub fn stream(&self) -> impl Stream<Item = Result<(StdioType, String), std::io::Error>> + '_ {
stream! {
let child_pid = self.pid;
let mut join = tokio::task::spawn_blocking(move || {
let Ok(status) = waitpid(child_pid, None) else {
panic!("Error waiting for child to complete");
};
debug!("Child status: {:?}", status);
status
});
let stdout = AsyncFd::try_from(self.stdout_raw_fd).unwrap();
let stderr = AsyncFd::try_from(self.stderr_raw_fd).unwrap();
let mut stdout_reader = LinesStream::new(BufReader::new(stdout).lines());
let mut stderr_reader = LinesStream::new(BufReader::new(stderr).lines());
let stdout_stream = Box::pin(stream! {
while let Some(Ok(item)) = stdout_reader.next().await {
yield item;
}
})
as Pin<Box<dyn Stream<Item = String> + Send>>;
let stderr_stream = Box::pin(stream! {
while let Some(Ok(item)) = stderr_reader.next().await {
yield item;
}
})
as Pin<Box<dyn Stream<Item = String> + Send>>;
let mut map = StreamMap::with_capacity(2);
map.insert(StdioType::Stdout, stdout_stream);
map.insert(StdioType::Stderr, stderr_stream);
loop {
tokio::select! {
biased;
Some(output) = map.next() => {
yield Ok(output);
},
status = &mut join => {
debug!("status");
let status = status.unwrap();
while let Some(output) = map.next().await {
yield Ok(output);
}
close(self.stdout_raw_fd).unwrap();
close(self.stderr_raw_fd).unwrap();
match status {
WaitStatus::Exited(pid, return_code) => {
debug!("Child exited with return code {}", return_code);
return;
},
WaitStatus::Signaled(pid, signal, _) => {
debug!("Child was killed by signal {:?}", signal);
return;
},
_ => {
panic!("Child process in unexpected state: '{:?}'", status);
},
}
},
}
}
}
}
}
pub struct TTYCommand<'a> {
command: &'a str,
args: &'a [String],
env: Vec<String>,
}
impl<'a> TTYCommand<'a> {
pub fn new(command: &'a str, args: &'a [String]) -> Self {
let env: Vec<String> = env::vars()
.into_iter()
.map(|x| format!("{}={}", x.0, x.1))
.collect();
TTYCommand { command, args, env }
}
async fn exec(&self) -> Result<()> {
let Ok(command) = which(self.command) else {
bail!("Unable to find '{}' on path", self.command);
};
let mut fixed_args = Vec::new();
fixed_args.push(command.to_str().unwrap().to_owned());
fixed_args.extend_from_slice(self.args);
let command = path2cstring(&command).unwrap();
let args = vec_slice_of_string_2_vec_of_cstring(&fixed_args).unwrap();
let env = vec_slice_of_string_2_vec_of_cstring(&self.env).unwrap();
if execve(&command, &args, &env).is_err() {
bail!(
"Unable to execve command '{:?}' with args {:?}",
command,
args
);
}
Ok(())
}
pub async fn spawn(&self) -> Result<XChildHandle> {
let Ok(stdout_pty) = openpty(None, None) else {
bail!("Unable to create pty for stdout");
};
let stdout_read_side = stdout_pty.master;
let stdout_write_side = stdout_pty.slave;
let Ok(stderr_pty) = openpty(None, None) else {
bail!("Unable to create pty for stderr");
};
let stderr_read_side = stderr_pty.master;
let stderr_write_side = stderr_pty.slave;
let Ok(res) = (unsafe { fork() }) else {
bail!("fork() failed");
};
match res {
ForkResult::Parent { child } => {
close(stdout_write_side.as_raw_fd()).unwrap();
close(stderr_write_side.as_raw_fd()).unwrap();
Ok(XChildHandle::new(
child,
stdout_read_side.as_raw_fd(),
stderr_read_side.as_raw_fd(),
)
.unwrap())
}
ForkResult::Child => {
close(stdout_read_side.as_raw_fd()).unwrap();
close(stderr_read_side.as_raw_fd()).unwrap();
dup2(stdout_write_side.as_raw_fd(), libc::STDOUT_FILENO).unwrap();
dup2(stderr_write_side.as_raw_fd(), libc::STDERR_FILENO).unwrap();
self.exec().await.unwrap(); unreachable!();
}
}
}
}