#[cfg(all(test, target_os = "macos"))]
mod tests;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use arcbox_connect::v1::{
MachineExecOutput, MachineExecRequest, MachineExecSignal, MachineExecWindow,
MachineTcpConnectRequest, TerminalSize,
};
use arcbox_constants::wire::MessageType;
use arcbox_transport::vsock::{VsockReceiver, VsockSender};
use buffa::Message;
use bytes::Bytes;
use tokio::sync::{Semaphore, mpsc};
use super::{AgentClient, wire};
use crate::error::{EngineError, Result};
const OUTPUT_WINDOW: u32 = 1024 * 1024;
const STDIN_FRAME: usize = 32 * 1024;
#[derive(Debug)]
pub enum ExecSessionInput {
Stdin(Vec<u8>),
Resize {
width: u16,
height: u16,
},
Signal(String),
}
impl ExecSessionInput {
fn frame(&self) -> Bytes {
match self {
Self::Stdin(data) => wire::build_message(MessageType::MachineExecInput, "", data),
Self::Resize { width, height } => {
let size = TerminalSize {
width: u32::from(*width),
height: u32::from(*height),
..Default::default()
};
wire::build_message(MessageType::MachineExecResize, "", &size.encode_to_vec())
}
Self::Signal(name) => {
let signal = MachineExecSignal {
name: name.clone(),
..Default::default()
};
wire::build_message(MessageType::MachineExecSignal, "", &signal.encode_to_vec())
}
}
}
}
pub struct ExecSessionOutput {
frames: mpsc::UnboundedReceiver<(usize, Result<MachineExecOutput>)>,
window: OutputWindow,
}
impl ExecSessionOutput {
pub async fn recv(&mut self) -> Option<Result<MachineExecOutput>> {
let (cost, item) = self.frames.recv().await?;
self.window.consumed(cost);
Some(item)
}
}
struct OutputWindow {
available: Arc<AtomicUsize>,
unreturned: usize,
frames: mpsc::UnboundedSender<Bytes>,
}
impl OutputWindow {
fn consumed(&mut self, len: usize) {
self.unreturned += len;
if self.unreturned < OUTPUT_WINDOW as usize / 2 {
return;
}
let bytes = std::mem::take(&mut self.unreturned);
self.available.fetch_add(bytes, Ordering::AcqRel);
let grant = MachineExecWindow {
bytes: bytes as u32,
..Default::default()
};
let frame = wire::build_message(
MessageType::MachineExecOutputWindow,
"",
&grant.encode_to_vec(),
);
let _ = self.frames.send(frame);
}
}
fn take_output_window(available: &AtomicUsize, cost: usize) -> Result<()> {
available
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |left| {
left.checked_sub(cost)
})
.map(drop)
.map_err(|_| EngineError::Machine("guest agent overran the exec output window".into()))
}
struct StdinWindow {
credit: Semaphore,
limit: usize,
frame: usize,
}
impl StdinWindow {
fn new(granted: u32) -> Result<Self> {
if granted == 0 {
return Err(EngineError::Machine(
"guest agent granted no exec stdin window".into(),
));
}
let limit = granted as usize;
Ok(Self {
credit: Semaphore::new(limit),
limit,
frame: STDIN_FRAME.min(limit),
})
}
async fn reserve(&self, len: usize) {
if let Ok(permit) = self.credit.acquire_many(len as u32).await {
permit.forget();
}
}
fn grant(&self, bytes: u32) -> Result<()> {
if self.credit.available_permits() + bytes as usize > self.limit {
return Err(EngineError::Machine(
"guest agent returned more exec stdin window than it granted".into(),
));
}
self.credit.add_permits(bytes as usize);
Ok(())
}
}
impl AgentClient {
pub async fn machine_exec(self, req: MachineExecRequest) -> Result<ExecSessionOutput> {
let (_, no_input) = mpsc::channel(1);
self.machine_exec_session(req, no_input).await
}
pub async fn machine_exec_session(
self,
req: MachineExecRequest,
input: mpsc::Receiver<ExecSessionInput>,
) -> Result<ExecSessionOutput> {
self.exec_session_with(
MessageType::MachineExecRequest,
MessageType::MachineExecOutput,
req,
input,
)
.await
}
pub async fn machine_debug_session(
self,
req: MachineExecRequest,
input: mpsc::Receiver<ExecSessionInput>,
) -> Result<ExecSessionOutput> {
self.exec_session_with(
MessageType::DebugExecRequest,
MessageType::DebugExecResponse,
req,
input,
)
.await
}
async fn exec_session_with(
self,
request_type: MessageType,
response_type: MessageType,
mut req: MachineExecRequest,
input: mpsc::Receiver<ExecSessionInput>,
) -> Result<ExecSessionOutput> {
req.output_window = OUTPUT_WINDOW;
let request = wire::build_message(request_type, "", &req.encode_to_vec());
self.open_session(request, response_type, input).await
}
pub async fn machine_tcp_connect(
self,
host: &str,
port: u16,
input: mpsc::Receiver<ExecSessionInput>,
) -> Result<ExecSessionOutput> {
let req = MachineTcpConnectRequest {
host: host.to_owned(),
port: port.into(),
output_window: OUTPUT_WINDOW,
..Default::default()
};
let request = wire::build_message(
MessageType::MachineTcpConnectRequest,
"",
&req.encode_to_vec(),
);
self.open_session(request, MessageType::MachineExecOutput, input)
.await
}
async fn open_session(
mut self,
request: Bytes,
response_type: MessageType,
input: mpsc::Receiver<ExecSessionInput>,
) -> Result<ExecSessionOutput> {
if !self.connected {
self.connect().await?;
}
self.transport
.async_send(request)
.await
.map_err(|source| EngineError::Transport {
context: "failed to send exec session request",
source,
})?;
let (sender, mut receiver) =
self.transport
.into_split()
.map_err(|source| EngineError::Transport {
context: "failed to split exec session transport",
source,
})?;
let stdin_window = match next_frame(&mut receiver, response_type).await? {
Frame::StdinWindow(bytes) => Arc::new(StdinWindow::new(bytes)?),
Frame::Output { .. } => {
return Err(EngineError::Machine(
"the machine's guest agent predates flow-controlled exec sessions; \
restart the machine to update it"
.into(),
));
}
};
let (frames_tx, frames) = mpsc::unbounded_channel();
let writer = tokio::spawn(write_frames(sender, frames));
let input_pump = tokio::spawn(pump_input(
input,
frames_tx.clone(),
Arc::clone(&stdin_window),
));
let available = Arc::new(AtomicUsize::new(OUTPUT_WINDOW as usize));
let (out_tx, out_rx) = mpsc::unbounded_channel();
tokio::spawn({
let available = Arc::clone(&available);
async move {
tokio::select! {
() = pump_output(&mut receiver, response_type, &out_tx, &available, &stdin_window) => {}
() = out_tx.closed() => {}
}
input_pump.abort();
writer.abort();
}
});
Ok(ExecSessionOutput {
frames: out_rx,
window: OutputWindow {
available,
unreturned: 0,
frames: frames_tx,
},
})
}
}
async fn write_frames(mut sender: VsockSender, mut frames: mpsc::UnboundedReceiver<Bytes>) {
while let Some(frame) = frames.recv().await {
if sender.send(frame).await.is_err() {
return;
}
}
}
async fn pump_input(
mut input: mpsc::Receiver<ExecSessionInput>,
frames: mpsc::UnboundedSender<Bytes>,
window: Arc<StdinWindow>,
) {
while let Some(item) = input.recv().await {
match item {
ExecSessionInput::Stdin(data) if !data.is_empty() => {
for chunk in data.chunks(window.frame) {
window.reserve(chunk.len()).await;
let frame = wire::build_message(MessageType::MachineExecInput, "", chunk);
if frames.send(frame).is_err() {
return;
}
}
}
other => {
if frames.send(other.frame()).is_err() {
return;
}
}
}
}
let _ = frames.send(ExecSessionInput::Stdin(Vec::new()).frame());
}
async fn pump_output(
receiver: &mut VsockReceiver,
expected: MessageType,
out: &mpsc::UnboundedSender<(usize, Result<MachineExecOutput>)>,
output_window: &AtomicUsize,
stdin_window: &StdinWindow,
) {
loop {
let (cost, item) = match next_frame(receiver, expected).await {
Ok(Frame::StdinWindow(bytes)) => match stdin_window.grant(bytes) {
Ok(()) => continue,
Err(e) => (0, Err(e)),
},
Ok(Frame::Output { output, cost }) => match take_output_window(output_window, cost) {
Ok(()) => (cost, Ok(output)),
Err(e) => (0, Err(e)),
},
Err(e) => (0, Err(e)),
};
let last = item.as_ref().map_or(true, |output| output.done);
if out.send((cost, item)).is_err() || last {
return;
}
}
}
enum Frame {
Output {
output: MachineExecOutput,
cost: usize,
},
StdinWindow(u32),
}
async fn next_frame(receiver: &mut VsockReceiver, expected: MessageType) -> Result<Frame> {
let raw = receiver
.recv()
.await
.map_err(|source| EngineError::Transport {
context: "failed to receive exec session output",
source,
})?;
let (resp_type, _, payload) = wire::parse_response(&raw)?;
if resp_type == MessageType::Error as u32 {
let (code, message) = wire::parse_error_response(&payload)
.unwrap_or_else(|_| (500, "unknown error".to_string()));
return Err(EngineError::Agent { code, message });
}
if resp_type == MessageType::MachineExecInputWindow as u32 {
let window = MachineExecWindow::decode_from_slice(&payload).map_err(decode_error)?;
return Ok(Frame::StdinWindow(window.bytes));
}
AgentClient::expect_response_type(resp_type, expected)?;
let output = MachineExecOutput::decode_from_slice(&payload).map_err(decode_error)?;
let cost = if output.done { 0 } else { payload.len() };
Ok(Frame::Output { output, cost })
}
fn decode_error(e: impl std::fmt::Display) -> EngineError {
EngineError::Machine(format!("decode error: {e}"))
}