use std::collections::HashMap;
use std::ffi::OsString;
use std::path::PathBuf;
use std::process::Stdio;
use std::time::Duration;
use tokio::io::BufReader;
use tokio::process::Command;
use tokio::sync::{mpsc, watch};
use tokio::task::JoinHandle;
use crate::signal::{
send_signal_to_process, signal_exit_code, Delivery, DEFAULT_KILL_GRACE_MS, DEFAULT_KILL_SIGNAL,
};
use crate::trace::trace_lazy;
use crate::{CommandResult, Result};
const DEFAULT_EXIT_PUMP_GRACE_MS: u64 = 100;
#[derive(Debug, Clone)]
pub enum OutputChunk {
Stdout(Vec<u8>),
Stderr(Vec<u8>),
Exit(i32),
}
pub struct StreamingRunner {
command: StreamingCommand,
cwd: Option<PathBuf>,
env: Option<HashMap<String, String>>,
stdin_content: Option<String>,
kill_signal: String,
kill_grace_ms: u64,
exit_pump_grace_ms: u64,
}
#[derive(Clone)]
enum StreamingCommand {
Shell(String),
Argv {
program: OsString,
args: Vec<OsString>,
},
}
impl StreamingRunner {
pub fn new(command: impl Into<String>) -> Self {
Self::with_command(StreamingCommand::Shell(command.into()))
}
pub fn from_argv<P, I, S>(program: P, args: I) -> Self
where
P: Into<OsString>,
I: IntoIterator<Item = S>,
S: Into<OsString>,
{
Self::with_command(StreamingCommand::Argv {
program: program.into(),
args: args.into_iter().map(Into::into).collect(),
})
}
fn with_command(command: StreamingCommand) -> Self {
StreamingRunner {
command,
cwd: None,
env: None,
stdin_content: None,
kill_signal: DEFAULT_KILL_SIGNAL.to_string(),
kill_grace_ms: DEFAULT_KILL_GRACE_MS,
exit_pump_grace_ms: DEFAULT_EXIT_PUMP_GRACE_MS,
}
}
pub fn cwd(mut self, path: impl Into<PathBuf>) -> Self {
self.cwd = Some(path.into());
self
}
pub fn env(mut self, env: HashMap<String, String>) -> Self {
self.env = Some(env);
self
}
pub fn stdin(mut self, content: impl Into<String>) -> Self {
self.stdin_content = Some(content.into());
self
}
pub fn kill_signal(mut self, signal: impl Into<String>) -> Self {
self.kill_signal = signal.into();
self
}
pub fn kill_grace_ms(mut self, ms: u64) -> Self {
self.kill_grace_ms = ms;
self
}
pub fn exit_pump_grace_ms(mut self, ms: u64) -> Self {
self.exit_pump_grace_ms = ms;
self
}
fn spawn(mut self) -> (OutputStream, JoinHandle<Result<()>>) {
let (tx, rx) = mpsc::channel(1024);
let (kill_tx, kill_rx) = mpsc::unbounded_channel::<String>();
let (pid_tx, pid_rx) = watch::channel(None);
let command = self.command.clone();
let cwd = self.cwd.take();
let env = self.env.take();
let stdin_content = self.stdin_content.take();
let grace = GraceWindows {
exit_pump_ms: self.exit_pump_grace_ms,
kill_ms: self.kill_grace_ms,
};
let kill_signal = self.kill_signal.clone();
let task = tokio::spawn(async move {
let channels = StreamChannels {
output_tx: tx,
kill_rx,
pid_tx,
};
let result =
run_streaming_process(command, cwd, env, stdin_content, grace, channels).await;
if let Err(error) = &result {
trace_lazy("StreamingRunner", || format!("Error: {error}"));
}
result
});
(
OutputStream {
rx,
kill_tx,
kill_signal,
killed: false,
pid_rx,
},
task,
)
}
pub fn stream(self) -> OutputStream {
self.spawn().0
}
pub async fn collect(self) -> Result<CommandResult> {
let stdin_content = self.stdin_content.clone();
let mut stdout = Vec::new();
let mut stderr = Vec::new();
let mut exit_code = 0;
let (mut stream, task) = self.spawn();
while let Some(chunk) = stream.rx.recv().await {
match chunk {
OutputChunk::Stdout(data) => stdout.extend(data),
OutputChunk::Stderr(data) => stderr.extend(data),
OutputChunk::Exit(code) => exit_code = code,
}
}
task.await.map_err(|error| {
std::io::Error::other(format!("streaming process task failed: {error}"))
})??;
let mut result = CommandResult::new(
String::from_utf8_lossy(&stdout).to_string(),
String::from_utf8_lossy(&stderr).to_string(),
exit_code,
);
if let Some(content) = stdin_content {
result.stdin = crate::result_streams::CapturedInput::new(content.into_bytes());
}
Ok(result)
}
}
pub struct OutputStream {
rx: mpsc::Receiver<OutputChunk>,
kill_tx: mpsc::UnboundedSender<String>,
kill_signal: String,
killed: bool,
pid_rx: watch::Receiver<Option<u32>>,
}
impl OutputStream {
pub async fn next(&mut self) -> Option<OutputChunk> {
self.rx.recv().await
}
pub fn pid(&self) -> Option<u32> {
*self.pid_rx.borrow()
}
pub async fn wait_for_pid(&mut self) -> Option<u32> {
match self.pid_rx.wait_for(|pid| pid.is_some()).await {
Ok(pid) => *pid,
Err(_) => None,
}
}
pub fn kill(&mut self) {
let signal = self.kill_signal.clone();
self.kill_with(&signal);
}
pub fn kill_with(&mut self, signal: &str) {
if self.killed {
return;
}
self.killed = true;
trace_lazy("OutputStream", || format!("kill | signal={}", signal));
let _ = self.kill_tx.send(signal.to_string());
}
pub async fn collect(mut self) -> (Vec<u8>, Vec<u8>, i32) {
let mut stdout = Vec::new();
let mut stderr = Vec::new();
let mut exit_code = 0;
while let Some(chunk) = self.rx.recv().await {
match chunk {
OutputChunk::Stdout(data) => stdout.extend(data),
OutputChunk::Stderr(data) => stderr.extend(data),
OutputChunk::Exit(code) => exit_code = code,
}
}
(stdout, stderr, exit_code)
}
pub async fn collect_stdout(mut self) -> Vec<u8> {
let mut stdout = Vec::new();
while let Some(chunk) = self.rx.recv().await {
if let OutputChunk::Stdout(data) = chunk {
stdout.extend(data);
}
}
stdout
}
}
impl Drop for OutputStream {
fn drop(&mut self) {
if !self.killed {
let _ = self.kill_tx.send(self.kill_signal.clone());
}
}
}
struct StreamChannels {
output_tx: mpsc::Sender<OutputChunk>,
kill_rx: mpsc::UnboundedReceiver<String>,
pid_tx: watch::Sender<Option<u32>>,
}
#[derive(Debug, Clone, Copy)]
struct GraceWindows {
exit_pump_ms: u64,
kill_ms: u64,
}
async fn run_streaming_process(
command: StreamingCommand,
cwd: Option<PathBuf>,
env: Option<HashMap<String, String>>,
stdin_content: Option<String>,
grace: GraceWindows,
channels: StreamChannels,
) -> Result<()> {
let StreamChannels {
output_tx: tx,
mut kill_rx,
pid_tx,
} = channels;
trace_lazy("StreamingRunner", || match &command {
StreamingCommand::Shell(command) => format!("Starting: {command}"),
StreamingCommand::Argv { program, args } => {
format!("Starting argv command: {program:?} {args:?}")
}
});
let mut cmd = match command {
StreamingCommand::Shell(command) => crate::utils::shell_command(&command, env.as_ref()),
StreamingCommand::Argv { program, args } => {
let mut cmd = Command::new(program);
cmd.args(args);
cmd
}
};
if stdin_content.is_some() {
cmd.stdin(Stdio::piped());
} else {
cmd.stdin(Stdio::null());
}
cmd.stdout(Stdio::piped());
cmd.stderr(Stdio::piped());
#[cfg(unix)]
cmd.process_group(0);
if let Some(ref cwd) = cwd {
cmd.current_dir(cwd);
}
if let Some(ref env_vars) = env {
for (key, value) in env_vars {
cmd.env(key, value);
}
}
let mut child = cmd.spawn()?;
let _ = pid_tx.send(child.id());
if let Some(content) = stdin_content {
if let Some(mut stdin) = child.stdin.take() {
use tokio::io::AsyncWriteExt;
let _ = stdin.write_all(content.as_bytes()).await;
let _ = stdin.shutdown().await;
}
}
let stdout = child.stdout.take();
let tx_stdout = tx.clone();
let stdout_handle = stdout.map(|stdout| {
tokio::spawn(async move {
let mut reader = BufReader::new(stdout);
let mut buf = vec![0u8; 8192];
loop {
use tokio::io::AsyncReadExt;
match reader.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if tx_stdout
.send(OutputChunk::Stdout(buf[..n].to_vec()))
.await
.is_err()
{
break;
}
}
Err(_) => break,
}
}
})
});
let stderr = child.stderr.take();
let tx_stderr = tx.clone();
let stderr_handle = stderr.map(|stderr| {
tokio::spawn(async move {
let mut reader = BufReader::new(stderr);
let mut buf = vec![0u8; 8192];
loop {
use tokio::io::AsyncReadExt;
match reader.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
if tx_stderr
.send(OutputChunk::Stderr(buf[..n].to_vec()))
.await
.is_err()
{
break;
}
}
Err(_) => break,
}
}
})
});
let pid = child.id();
let code;
tokio::select! {
status = child.wait() => {
code = status_to_code(status?);
}
maybe_signal = kill_rx.recv() => {
let signal = maybe_signal.unwrap_or_else(|| DEFAULT_KILL_SIGNAL.to_string());
trace_lazy("StreamingRunner", || format!("Kill requested | signal={}", signal));
let survived_grace = if grace.kill_ms == 0 {
true
} else {
if let Some(pid) = pid {
send_signal_to_process(pid, &signal, Delivery::ProcessAndGroup);
}
tokio::time::timeout(Duration::from_millis(grace.kill_ms), child.wait())
.await
.is_err()
};
if survived_grace {
if let Some(pid) = pid {
send_signal_to_process(pid, "SIGKILL", Delivery::ProcessAndGroup);
}
let _ = child.start_kill();
let _ = child.wait().await;
}
code = signal_exit_code(&signal);
}
}
let stdout_abort = stdout_handle.as_ref().map(|h| h.abort_handle());
let stderr_abort = stderr_handle.as_ref().map(|h| h.abort_handle());
let drain = async {
if let Some(handle) = stdout_handle {
let _ = handle.await;
}
if let Some(handle) = stderr_handle {
let _ = handle.await;
}
};
if tokio::time::timeout(Duration::from_millis(grace.exit_pump_ms), drain)
.await
.is_err()
{
if let Some(abort) = stdout_abort {
abort.abort();
}
if let Some(abort) = stderr_abort {
abort.abort();
}
}
let _ = tx.send(OutputChunk::Exit(code)).await;
trace_lazy("StreamingRunner", || format!("Exited with code: {}", code));
Ok(())
}
fn status_to_code(status: std::process::ExitStatus) -> i32 {
if let Some(code) = status.code() {
return code;
}
#[cfg(unix)]
{
use std::os::unix::process::ExitStatusExt;
if let Some(sig) = status.signal() {
return 128 + sig;
}
}
-1
}
#[async_trait::async_trait]
pub trait AsyncIterator {
type Item;
async fn next(&mut self) -> Option<Self::Item>;
}
#[async_trait::async_trait]
impl AsyncIterator for OutputStream {
type Item = OutputChunk;
async fn next(&mut self) -> Option<Self::Item> {
self.rx.recv().await
}
}
pub trait IntoStream {
fn into_stream(self) -> OutputStream;
}
impl IntoStream for crate::ProcessRunner {
fn into_stream(self) -> OutputStream {
let streaming = StreamingRunner::new(self.command().to_string());
streaming.stream()
}
}