use std::ffi::OsString;
use std::io::{self, BufRead, IoSlice, IoSliceMut, Write};
use std::os::fd::{AsRawFd, FromRawFd, OwnedFd, RawFd};
use std::os::unix::net::UnixStream;
use std::os::unix::process::CommandExt;
use std::process::{Command, ExitStatus, Output};
use std::time::Duration;
use nix::sys::signal::{Signal, kill};
use nix::sys::socket::{ControlMessage, ControlMessageOwned, MsgFlags, recvmsg, sendmsg};
use nix::unistd::Pid;
pub fn args<S: AsRef<str>>(items: &[S]) -> Vec<OsString> {
items.iter().map(|s| OsString::from(s.as_ref())).collect()
}
pub fn exec_replace(bin: &str, args: &[OsString]) -> anyhow::Error {
let mut cmd = Command::new(bin);
cmd.args(args);
let err = cmd.exec();
anyhow::Error::from(err).context(format!("failed to exec {bin}"))
}
pub fn run_piped(bin: &str, args: &[OsString]) -> anyhow::Result<Output> {
let output = Command::new(bin)
.args(args)
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped())
.spawn()?
.wait_with_output()?;
Ok(output)
}
pub fn spawn_interactive(bin: &str, args: &[OsString]) -> anyhow::Result<ExitStatus> {
let status = Command::new(bin).args(args).status()?;
Ok(status)
}
pub fn run_with_log(
bin: &str,
args: &[OsString],
log: &mut std::fs::File,
mirror: bool,
) -> anyhow::Result<ExitStatus> {
let mut child = Command::new(bin)
.args(args)
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped())
.spawn()
.map_err(|e| anyhow::Error::new(e).context(format!("failed to execute {bin}")))?;
let mut out = child.stdout.take().expect("stdout piped");
let mut err = child.stderr.take().expect("stderr piped");
let mut log_out = log.try_clone()?;
let mut log_err = log.try_clone()?;
let t_out = std::thread::spawn(move || tee_stream(&mut out, &mut log_out, false, mirror));
let t_err = std::thread::spawn(move || tee_stream(&mut err, &mut log_err, true, mirror));
let status = child.wait()?;
let _ = t_out.join();
let _ = t_err.join();
let _ = log.flush();
Ok(status)
}
fn tee_stream<R: io::Read, W: io::Write>(
src: &mut R,
dst: &mut W,
to_stderr: bool,
mirror: bool,
) {
let mut reader = io::BufReader::new(src);
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line) {
Ok(0) | Err(_) => break,
Ok(_) => {
if dst.write_all(line.as_bytes()).is_err() {
break;
}
if mirror {
let _ = if to_stderr {
io::stderr().write_all(line.as_bytes())
} else {
io::stdout().write_all(line.as_bytes())
};
}
}
}
}
}
pub fn run_piped_timeout(
bin: &str,
args: &[OsString],
timeout: Duration,
) -> anyhow::Result<Output> {
let child = Command::new(bin)
.args(args)
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::piped())
.spawn()?;
wait_child_timeout(child, timeout)
}
pub fn spawn_interactive_timeout(
bin: &str,
args: &[OsString],
timeout: Duration,
) -> anyhow::Result<ExitStatus> {
let mut child = Command::new(bin).args(args).spawn()?;
let pid = Pid::from_raw(child.id().cast_signed());
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
if rx.recv_timeout(timeout).is_err() {
let _ = kill(pid, Signal::SIGTERM);
std::thread::sleep(Duration::from_secs(5));
let _ = kill(pid, Signal::SIGKILL);
}
});
let status = child.wait()?;
let _ = tx.send(());
Ok(status)
}
pub fn wait_child_timeout(child: std::process::Child, timeout: Duration) -> anyhow::Result<Output> {
let pid = Pid::from_raw(child.id().cast_signed());
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
if rx.recv_timeout(timeout).is_err() {
let _ = kill(pid, Signal::SIGKILL);
}
});
let output = child.wait_with_output()?;
let _ = tx.send(());
Ok(output)
}
pub fn open_pidfd(pid: i32) -> io::Result<OwnedFd> {
let pid = rustix::process::Pid::from_raw(pid)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "invalid PID"))?;
rustix::process::pidfd_open(pid, rustix::process::PidfdFlags::empty()).map_err(io::Error::from)
}
pub fn send_fd(stream: &UnixStream, fd: RawFd) -> io::Result<()> {
let raw_fd = stream.as_raw_fd();
let cmsg = ControlMessage::ScmRights(&[fd]);
let iov = [IoSlice::new(&[0u8])];
loop {
match sendmsg::<()>(raw_fd, &iov, &[cmsg], MsgFlags::empty(), None) {
Ok(_) => return Ok(()),
Err(nix::errno::Errno::EINTR) => {}
Err(e) => return Err(io::Error::from(e)),
}
}
}
pub fn adopt_scm_fd(raw: RawFd) -> OwnedFd {
#[allow(unsafe_code)]
unsafe {
OwnedFd::from_raw_fd(raw)
}
}
pub fn recv_fd(stream: &UnixStream) -> io::Result<Option<RawFd>> {
let raw_fd = stream.as_raw_fd();
let mut buf = [0u8; 1];
let mut iov = [IoSliceMut::new(&mut buf)];
let mut cmsg_buf = vec![0u8; 256];
let msg = loop {
match recvmsg::<()>(raw_fd, &mut iov, Some(&mut cmsg_buf), MsgFlags::empty()) {
Ok(m) => break m,
Err(nix::errno::Errno::EINTR) => {}
Err(e) => return Err(io::Error::from(e)),
}
};
if msg.bytes == 0 {
return Ok(None);
}
if let Ok(cmsgs) = msg.cmsgs() {
for cmsg in cmsgs {
if let ControlMessageOwned::ScmRights(fds) = cmsg {
if let Some(&fd) = fds.first() {
return Ok(Some(fd));
}
}
}
}
Ok(None)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn args_builds_osstring_vec() {
let v = args(&["foo", "bar", "baz"]);
assert_eq!(v.len(), 3);
assert_eq!(v[0], "foo");
assert_eq!(v[1], "bar");
assert_eq!(v[2], "baz");
}
#[test]
fn args_accepts_mixed_types() {
let s = String::from("hello");
let v = args(&["a", &s, "c"]);
assert_eq!(v[1], "hello");
}
}