use std::io::Write as _;
use std::os::fd::OwnedFd;
use std::os::unix::process::ExitStatusExt as _;
use std::pin::pin;
use std::process::ExitStatus;
use anyhow::Context as _;
use arcbox_connect::v1::{MachineExecOutput, MachineExecWindow};
use buffa::Message as _;
use nix::sys::signal::Signal;
use tokio::io::{AsyncRead, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _};
use tokio::process::Child;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use super::control::{self, Control};
use super::flow::Flow;
use crate::rpc::{ErrorResponse, MessageType, write_message};
pub(super) const OUTPUT_CHANNEL_CAPACITY: usize = 16;
pub(super) const OUTPUT_CHUNK: usize = 8192;
pub(super) struct Chunk {
pub(super) stream: &'static str,
pub(super) data: Vec<u8>,
}
pub(super) enum StdinItem {
Data(Vec<u8>),
Eof,
}
pub(super) struct Streams {
pub(super) output: mpsc::Receiver<Chunk>,
pub(super) stdin: mpsc::UnboundedSender<StdinItem>,
pub(super) delivered: mpsc::UnboundedReceiver<usize>,
}
#[derive(Clone, Copy)]
pub(super) struct OutFrame<'a> {
pub(super) trace_id: &'a str,
pub(super) msg_type: MessageType,
}
pub(super) enum Ended {
Exited(ExitStatus),
Closed,
}
pub(super) async fn run<S>(
stream: &mut S,
out: OutFrame<'_>,
flow: &Flow,
streams: Streams,
terminal: Option<&OwnedFd>,
child: &mut Child,
) -> anyhow::Result<()>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let pid = child.id();
let exited = async {
let status = child
.wait()
.await
.context("failed to wait for exec child")?;
Ok(Ended::Exited(status))
};
if !relay(stream, out, flow, streams, terminal, pid, exited).await? {
kill_session(pid);
let _ = child.wait().await;
}
Ok(())
}
pub(super) async fn relay<S>(
stream: &mut S,
out: OutFrame<'_>,
flow: &Flow,
streams: Streams,
terminal: Option<&OwnedFd>,
pid: Option<u32>,
ended: impl Future<Output = anyhow::Result<Ended>>,
) -> anyhow::Result<bool>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let (mut conn_rd, mut conn_wr) = tokio::io::split(stream);
let ended = {
let output = pump(
&mut conn_wr,
out,
flow,
streams.output,
streams.delivered,
ended,
);
let control = serve_control(&mut conn_rd, streams.stdin, flow, terminal, pid);
tokio::select! {
ended = output => ended.ok(),
() = control => None,
}
};
match ended {
Some(ended) => write_end(&mut conn_wr, out, ended).await.map(|()| true),
None => Ok(false),
}
}
async fn pump<W>(
conn: &mut W,
out: OutFrame<'_>,
flow: &Flow,
mut output: mpsc::Receiver<Chunk>,
mut delivered: mpsc::UnboundedReceiver<usize>,
ended: impl Future<Output = anyhow::Result<Ended>>,
) -> anyhow::Result<Ended>
where
W: AsyncWrite + Unpin,
{
let mut ended = pin!(ended);
let mut pending: Option<Vec<u8>> = None;
let mut output_open = true;
loop {
let pending_len = pending.as_ref().map_or(0, Vec::len);
tokio::select! {
Some(mut len) = delivered.recv() => {
while let Ok(more) = delivered.try_recv() {
len += more;
}
if let Some(bytes) = flow.stdin_delivered(len) {
write_window(conn, out.trace_id, bytes).await?;
}
}
() = flow.reserve_output(pending_len), if pending.is_some() => {
if let Some(frame) = pending.take() {
write_message(conn, out.msg_type, out.trace_id, &frame).await?;
}
}
chunk = output.recv(), if output_open && pending.is_none() => match chunk {
Some(chunk) => pending = Some(encode_output(chunk)),
None => {
output_open = false;
let eof = MachineExecOutput {
eof: true,
..Default::default()
};
pending = Some(eof.encode_to_vec());
}
},
ended = ended.as_mut(), if !output_open && pending.is_none() => return ended,
}
}
}
async fn serve_control<R>(
conn: &mut R,
stdin: mpsc::UnboundedSender<StdinItem>,
flow: &Flow,
terminal: Option<&OwnedFd>,
pid: Option<u32>,
) where
R: AsyncRead + Unpin,
{
while let Some(control) = control::next(conn).await {
let applied = match control {
Control::Stdin(data) => flow.admit_stdin(data.len()).map(|()| {
let _ = stdin.send(StdinItem::Data(data));
}),
Control::Eof => {
let _ = stdin.send(StdinItem::Eof);
Ok(())
}
Control::Resize(size) => {
if let Some(master) = terminal {
if let Err(e) = arcbox_pty::resize(master, size) {
tracing::debug!(error = %e, "pty resize failed");
}
}
Ok(())
}
Control::Signal(signal) => {
signal_process(pid, signal);
Ok(())
}
Control::OutputWindow(bytes) => flow.return_output(bytes),
};
if let Err(e) = applied {
tracing::warn!(error = %format!("{e:#}"), "machine exec host broke flow control");
return;
}
}
}
pub(super) fn read_output<R>(
mut pipe: R,
stream: &'static str,
chunks: mpsc::Sender<Chunk>,
) -> JoinHandle<()>
where
R: AsyncRead + Unpin + Send + 'static,
{
tokio::spawn(async move {
let mut buf = vec![0u8; OUTPUT_CHUNK];
loop {
match pipe.read(&mut buf).await {
Ok(0) => break,
Ok(n) => {
let chunk = Chunk {
stream,
data: buf[..n].to_vec(),
};
if chunks.send(chunk).await.is_err() {
break;
}
}
Err(e) => {
tracing::warn!(error = %e, stream, "machine exec output read error");
break;
}
}
}
})
}
pub(super) async fn write_stdin<W>(
mut pipe: Option<W>,
mut items: mpsc::UnboundedReceiver<StdinItem>,
delivered: mpsc::UnboundedSender<usize>,
) where
W: AsyncWrite + Unpin,
{
while let Some(item) = items.recv().await {
match item {
StdinItem::Data(data) => {
if let Some(open) = pipe.as_mut() {
if open.write_all(&data).await.is_err() {
pipe = None;
}
}
let _ = delivered.send(data.len());
}
StdinItem::Eof => {
if let Some(mut open) = pipe {
let _ = open.shutdown().await;
}
return;
}
}
}
}
pub(super) fn write_terminal(
mut master: std::fs::File,
mut items: mpsc::UnboundedReceiver<StdinItem>,
delivered: mpsc::UnboundedSender<usize>,
) {
let mut open = true;
while let Some(item) = items.blocking_recv() {
let StdinItem::Data(data) = item else {
continue;
};
if open {
if let Err(e) = master.write_all(&data) {
tracing::debug!(error = %e, "pty stdin write ended");
open = false;
}
}
let _ = delivered.send(data.len());
}
}
pub(super) async fn write_window<W>(
writer: &mut W,
trace_id: &str,
bytes: u32,
) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let frame = MachineExecWindow {
bytes,
..Default::default()
};
write_message(
writer,
MessageType::MachineExecInputWindow,
trace_id,
&frame.encode_to_vec(),
)
.await
}
pub(super) async fn write_error<W>(
writer: &mut W,
trace_id: &str,
err: &ErrorResponse,
) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
write_message(writer, MessageType::Error, trace_id, &err.encode()).await
}
fn encode_output(chunk: Chunk) -> Vec<u8> {
MachineExecOutput {
stream: chunk.stream.to_owned(),
data: chunk.data,
..Default::default()
}
.encode_to_vec()
}
async fn write_end<W>(writer: &mut W, out: OutFrame<'_>, ended: Ended) -> anyhow::Result<()>
where
W: AsyncWrite + Unpin,
{
let frame = match ended {
Ended::Exited(status) => MachineExecOutput {
done: true,
exit_code: status.code().unwrap_or(-1),
exit_signal: status.signal().map(signal_name).unwrap_or_default(),
..Default::default()
},
Ended::Closed => MachineExecOutput {
done: true,
..Default::default()
},
};
write_message(writer, out.msg_type, out.trace_id, &frame.encode_to_vec()).await
}
fn signal_name(signal: i32) -> String {
Signal::try_from(signal).map_or_else(
|_| signal.to_string(),
|s| s.as_str().trim_start_matches("SIG").to_owned(),
)
}
fn signal_process(pid: Option<u32>, signal: Signal) {
if let Some(pid) = pid.and_then(|p| i32::try_from(p).ok()) {
if let Err(e) = nix::sys::signal::kill(nix::unistd::Pid::from_raw(pid), signal) {
tracing::debug!(error = %e, ?signal, "machine exec signal not delivered");
}
}
}
fn kill_session(pid: Option<u32>) {
if let Some(pid) = pid.and_then(|p| i32::try_from(p).ok()) {
let _ = nix::sys::signal::killpg(nix::unistd::Pid::from_raw(pid), Signal::SIGKILL);
}
}