use std::collections::HashMap;
#[cfg(feature = "stream")]
use std::future::Future;
#[cfg(any(all(feature = "named-pipe", windows), all(feature = "uds", unix)))]
use std::path::Path;
#[cfg(feature = "stream")]
use std::pin::Pin;
use std::sync::{Arc, atomic::AtomicU32};
#[cfg(feature = "stream")]
use std::time::Duration;
use microsandbox_protocol::message::FLAG_BULK;
#[cfg(feature = "stream")]
use microsandbox_protocol::message::FLAG_TERMINAL;
use microsandbox_protocol::{
bulk::{BULK_PROTOCOL_VERSION, BulkCancel, BulkRecord, MAX_BULK_RECORD_PAYLOAD},
codec::{self, RawFrame},
core::Ready,
message::{Message, MessageType, PROTOCOL_VERSION},
};
#[cfg(feature = "stream")]
use microsandbox_protocol::{codec::MAX_FRAME_SIZE, message::FRAME_HEADER_SIZE};
use serde::Serialize;
#[cfg(feature = "stream")]
use tokio::io::{AsyncRead, AsyncWrite};
#[cfg(all(feature = "uds", unix))]
use tokio::net::UnixStream;
#[cfg(all(feature = "named-pipe", windows))]
use tokio::net::windows::named_pipe::ClientOptions;
#[cfg(all(feature = "stream", feature = "uds", unix))]
use tokio::sync::watch;
use tokio::sync::{Mutex, mpsc, oneshot};
use tokio::task::JoinHandle;
#[cfg(feature = "stream")]
use tokio::time::Instant;
use super::error::{AgentClientError, AgentClientResult};
#[cfg(all(feature = "uds", unix))]
use super::local_shm::{
LOCAL_SHM_FORMAT_V1, LocalBulkRelease, LocalShmClient, LocalShmFrame, LocalShmUpgrade,
PreparedLocalBulk, SharedArenaConsumer, SharedArenaProducer, decode_local_body,
encode_local_bulk_ref, encode_local_bulk_release, local_upgrade_request_frame,
receive_local_shm_upgrade,
};
#[cfg(feature = "stream")]
const DEFAULT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
#[cfg(all(feature = "named-pipe", windows))]
const WINDOWS_PIPE_CONNECT_RETRY: Duration = Duration::from_millis(10);
#[cfg(feature = "stream")]
const WRITER_QUEUE_CAPACITY: usize = 8;
const REQUEST_QUEUE_CAPACITY: usize = 1;
const STREAM_QUEUE_CAPACITY: usize = 2;
const LEGACY_PROTOCOL_VERSION: u8 = 1;
#[cfg(feature = "stream")]
const LEGACY_RELAY_ID_RANGE_STEP: u32 = u32::MAX / 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AgentProtocol {
Current,
LegacyV1,
}
#[derive(Debug, Clone)]
pub enum AgentFrame {
Control(Message),
Bulk(BulkRecord),
}
pub struct AgentClient {
writer: mpsc::Sender<WriterCommand>,
next_id: AtomicU32,
id_min: u32,
id_max: u32,
protocol: AgentProtocol,
negotiated_version: u8,
pending: Arc<Mutex<HashMap<u32, CorrelationRoute>>>,
reader_handle: JoinHandle<()>,
writer_handle: JoinHandle<()>,
ready_body: Vec<u8>,
ready: Ready,
#[cfg(all(feature = "uds", unix))]
local_outbound: Option<SharedArenaProducer>,
}
#[cfg(feature = "stream")]
struct AgentHandshake {
id_min: u32,
id_max: u32,
protocol: AgentProtocol,
negotiated_version: u8,
ready_body: Vec<u8>,
ready: Ready,
}
#[cfg_attr(not(feature = "stream"), allow(dead_code))]
struct WriterCommand {
frame: WriterFrame,
ack: oneshot::Sender<AgentClientResult<()>>,
}
#[cfg_attr(not(feature = "stream"), allow(dead_code))]
enum WriterFrame {
Control(RawFrame),
Bulk(BulkRecord),
#[cfg(all(feature = "uds", unix))]
LocalBulk(PreparedLocalBulk),
}
enum InboundFrame {
Raw(RawFrame),
#[cfg(all(feature = "uds", unix))]
Bulk(BulkRecord),
}
struct CorrelationRoute {
tx: mpsc::Sender<InboundFrame>,
state: CorrelationState,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum CorrelationState {
Active,
Cancelling,
}
#[cfg(feature = "stream")]
trait HandshakeReader {
fn read_exact_handshake<'a>(
&'a mut self,
out: &'a mut [u8],
) -> Pin<Box<dyn Future<Output = AgentClientResult<()>> + Send + 'a>>;
fn read_frame_handshake<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = AgentClientResult<RawFrame>> + Send + 'a>>;
}
impl AgentProtocol {
fn version(self) -> u8 {
match self {
Self::Current => PROTOCOL_VERSION,
Self::LegacyV1 => LEGACY_PROTOCOL_VERSION,
}
}
}
impl InboundFrame {
fn id(&self) -> u32 {
match self {
Self::Raw(frame) => frame.id,
#[cfg(all(feature = "uds", unix))]
Self::Bulk(record) => record.id,
}
}
fn flags(&self) -> u8 {
match self {
Self::Raw(frame) => frame.flags,
#[cfg(all(feature = "uds", unix))]
Self::Bulk(_) => FLAG_BULK,
}
}
fn into_raw_frame(self) -> AgentClientResult<RawFrame> {
match self {
Self::Raw(frame) => Ok(frame),
#[cfg(all(feature = "uds", unix))]
Self::Bulk(record) => {
let mut body = Vec::with_capacity(12 + record.payload.len());
body.push(record.kind as u8);
body.push(record.flow as u8);
body.extend_from_slice(&[0, 0]);
body.extend_from_slice(&record.offset.to_be_bytes());
body.extend_from_slice(&record.payload);
Ok(RawFrame {
id: record.id,
flags: FLAG_BULK,
body,
})
}
}
}
}
impl AgentClient {
#[cfg(any(all(feature = "named-pipe", windows), all(feature = "uds", unix)))]
pub async fn connect(sock_path: impl AsRef<Path>) -> AgentClientResult<Self> {
Self::connect_with_timeout(sock_path, DEFAULT_HANDSHAKE_TIMEOUT).await
}
#[cfg(any(all(feature = "named-pipe", windows), all(feature = "uds", unix)))]
pub async fn connect_with_timeout(
sock_path: impl AsRef<Path>,
timeout: Duration,
) -> AgentClientResult<Self> {
let deadline = Instant::now() + timeout;
Self::connect_with_deadline(sock_path, deadline).await
}
#[cfg(any(all(feature = "named-pipe", windows), all(feature = "uds", unix)))]
pub async fn connect_with_deadline(
sock_path: impl AsRef<Path>,
deadline: Instant,
) -> AgentClientResult<Self> {
let sock_path = sock_path.as_ref();
#[cfg(all(feature = "uds", unix))]
{
let stream = connect_local_stream(sock_path, deadline).await?;
match Self::connect_uds_stream_with_deadline(stream, deadline, true).await {
Ok(client) => Ok(client),
Err(AgentClientError::LocalTransport(error)) if Instant::now() < deadline => {
tracing::warn!(%error, "agent client: local shared-arena upgrade failed; reconnecting in-band");
let stream = connect_local_stream(sock_path, deadline).await?;
Self::connect_uds_stream_with_deadline(stream, deadline, false).await
}
Err(error) => Err(error),
}
}
#[cfg(all(feature = "named-pipe", windows))]
{
let stream = connect_local_stream(sock_path, deadline).await?;
Self::connect_stream_with_deadline(stream, deadline).await
}
}
#[cfg(feature = "stream")]
pub async fn connect_stream<S>(stream: S) -> AgentClientResult<Self>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
Self::connect_stream_with_timeout(stream, DEFAULT_HANDSHAKE_TIMEOUT).await
}
#[cfg(feature = "stream")]
pub async fn connect_stream_with_timeout<S>(
stream: S,
timeout: Duration,
) -> AgentClientResult<Self>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let deadline = Instant::now() + timeout;
Self::connect_stream_with_deadline(stream, deadline).await
}
#[cfg(feature = "stream")]
pub async fn connect_stream_with_deadline<S>(
stream: S,
deadline: Instant,
) -> AgentClientResult<Self>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let (mut reader, writer) = tokio::io::split(stream);
let handshake = perform_handshake(&mut reader, deadline).await?;
#[cfg(all(feature = "uds", unix))]
{
finish_connection(reader, writer, handshake, None).await
}
#[cfg(not(all(feature = "uds", unix)))]
{
finish_connection(reader, writer, handshake).await
}
}
#[cfg(all(feature = "uds", unix))]
async fn connect_uds_stream_with_deadline(
mut stream: UnixStream,
deadline: Instant,
allow_local: bool,
) -> AgentClientResult<Self> {
let handshake = perform_handshake(&mut stream, deadline).await?;
let selected = allow_local
&& handshake
.ready
.local_transport
.as_ref()
.and_then(|capability| capability.select_supported(LOCAL_SHM_FORMAT_V1))
== Some(LOCAL_SHM_FORMAT_V1);
let local = if selected {
tokio::time::timeout_at(
deadline,
codec::write_raw_frame(&mut stream, &local_upgrade_request_frame()),
)
.await
.map_err(|_| {
AgentClientError::LocalTransport("upgrade request write timed out".into())
})?
.map_err(|error| AgentClientError::LocalTransport(error.to_string()))?;
match tokio::time::timeout_at(deadline, receive_local_shm_upgrade(&stream))
.await
.map_err(|_| {
AgentClientError::LocalTransport("descriptor acknowledgement timed out".into())
})?
.map_err(|error| AgentClientError::LocalTransport(error.to_string()))?
{
LocalShmUpgrade::Accepted(fds) => Some(
LocalShmClient::from_fds(fds)
.map_err(|error| AgentClientError::LocalTransport(error.to_string()))?,
),
LocalShmUpgrade::Rejected => None,
}
} else {
None
};
let std_reader = stream.into_std()?;
let std_writer = std_reader.try_clone()?;
let reader = UnixStream::from_std(std_reader)?;
let writer = UnixStream::from_std(std_writer)?;
finish_connection(reader, writer, handshake, local).await
}
pub async fn close(self) {
}
}
impl AgentClient {
pub async fn request_raw(&self, flags: u8, body: Vec<u8>) -> AgentClientResult<RawFrame> {
let (tx, mut rx) = mpsc::channel(REQUEST_QUEUE_CAPACITY);
let id = self.reserve_id(tx).await?;
if let Err(e) = self.write_frame_owned(id, flags, body).await {
self.pending.lock().await.remove(&id);
return Err(e);
}
let frame = rx
.recv()
.await
.ok_or(AgentClientError::ReaderClosed(id))?
.into_raw_frame()?;
self.pending.lock().await.remove(&id);
Ok(frame)
}
pub async fn stream_raw(
&self,
flags: u8,
body: Vec<u8>,
) -> AgentClientResult<(u32, mpsc::Receiver<RawFrame>)> {
let (id, inbound_rx) = self.open_inbound_stream(flags, body).await?;
let (tx, rx) = mpsc::channel(STREAM_QUEUE_CAPACITY);
tokio::spawn(materialize_raw_stream_task(inbound_rx, tx));
Ok((id, rx))
}
async fn open_inbound_stream(
&self,
flags: u8,
body: Vec<u8>,
) -> AgentClientResult<(u32, mpsc::Receiver<InboundFrame>)> {
let (tx, rx) = mpsc::channel(STREAM_QUEUE_CAPACITY);
let id = self.reserve_id(tx).await?;
if let Err(e) = self.write_frame_owned(id, flags, body).await {
self.pending.lock().await.remove(&id);
return Err(e);
}
Ok((id, rx))
}
pub async fn send_raw(&self, id: u32, flags: u8, body: &[u8]) -> AgentClientResult<()> {
self.write_frame(id, flags, body).await
}
pub async fn forget_stream(&self, id: u32) {
self.pending.lock().await.remove(&id);
}
pub fn ready_bytes(&self) -> &[u8] {
&self.ready_body
}
pub fn protocol(&self) -> AgentProtocol {
self.protocol
}
pub fn is_legacy_protocol(&self) -> bool {
self.protocol == AgentProtocol::LegacyV1
}
pub fn negotiated_version(&self) -> u8 {
self.negotiated_version
}
pub fn agent_version(&self) -> &str {
&self.ready.agent_version
}
pub fn supports(&self, t: MessageType) -> bool {
t.min_protocol_version() <= self.negotiated_version
}
pub fn ensure_version_compat(&self, t: MessageType) -> AgentClientResult<()> {
Self::ensure_version_compat_for(t, self.negotiated_version)
}
pub fn ensure_version_compat_for(t: MessageType, negotiated: u8) -> AgentClientResult<()> {
if t.is_available_at(negotiated) {
return Ok(());
}
Err(AgentClientError::UnsupportedOperation {
msg_type: t.as_str(),
needs: t.min_protocol_version(),
peer: negotiated,
})
}
}
impl AgentClient {
pub async fn request<T: Serialize>(
&self,
t: MessageType,
payload: &T,
) -> AgentClientResult<Message> {
self.ensure_version_compat(t)?;
let flags = t.flags();
let body = encode_message_body(self.protocol.version(), t, payload)?;
let frame = self.request_raw(flags, body).await?;
Ok(codec::raw_frame_to_message(frame)?)
}
pub async fn stream<T: Serialize>(
&self,
t: MessageType,
payload: &T,
) -> AgentClientResult<(u32, mpsc::Receiver<Message>)> {
self.ensure_version_compat(t)?;
let flags = t.flags();
let body = encode_message_body(self.protocol.version(), t, payload)?;
let (id, raw_rx) = self.open_inbound_stream(flags, body).await?;
let (tx, rx) = mpsc::channel(STREAM_QUEUE_CAPACITY);
tokio::spawn(decode_stream_task(raw_rx, tx));
Ok((id, rx))
}
pub async fn stream_frames<T: Serialize>(
&self,
t: MessageType,
payload: &T,
) -> AgentClientResult<(u32, mpsc::Receiver<AgentFrame>)> {
self.ensure_version_compat(t)?;
let flags = t.flags();
let body = encode_message_body(self.protocol.version(), t, payload)?;
let (id, raw_rx) = self.open_inbound_stream(flags, body).await?;
let (tx, rx) = mpsc::channel(STREAM_QUEUE_CAPACITY);
tokio::spawn(decode_frame_stream_task(raw_rx, tx));
Ok((id, rx))
}
pub async fn send<T: Serialize>(
&self,
id: u32,
t: MessageType,
payload: &T,
) -> AgentClientResult<()> {
self.ensure_version_compat(t)?;
let flags = t.flags();
let body = encode_message_body(self.protocol.version(), t, payload)?;
self.write_frame_owned(id, flags, body).await
}
pub async fn send_bulk(&self, record: BulkRecord) -> AgentClientResult<()> {
self.ensure_bulk_supported()?;
#[cfg(all(feature = "uds", unix))]
if let Some(producer) = self.local_outbound.as_ref() {
let prepared = producer
.prepare(&record)
.await
.map_err(|error| AgentClientError::LocalTransport(error.to_string()))?;
return self
.write_writer_frame(WriterFrame::LocalBulk(prepared))
.await;
}
self.write_writer_frame(WriterFrame::Bulk(record)).await
}
fn ensure_bulk_supported(&self) -> AgentClientResult<()> {
if self.negotiated_version < BULK_PROTOCOL_VERSION {
return Err(AgentClientError::UnsupportedOperation {
msg_type: "raw bulk record",
needs: BULK_PROTOCOL_VERSION,
peer: self.negotiated_version,
});
}
Ok(())
}
async fn write_writer_frame(&self, frame: WriterFrame) -> AgentClientResult<()> {
let (ack, written) = oneshot::channel();
self.writer
.send(WriterCommand { frame, ack })
.await
.map_err(|_| AgentClientError::Closed)?;
written.await.map_err(|_| AgentClientError::Closed)?
}
pub async fn cancel_bulk(&self, id: u32, cancel: &BulkCancel) -> AgentClientResult<()> {
if let Some(route) = self.pending.lock().await.get_mut(&id) {
route.state = CorrelationState::Cancelling;
}
self.send(id, MessageType::BulkCancel, cancel).await
}
pub fn ready(&self) -> AgentClientResult<Ready> {
Ok(self.ready.clone())
}
}
impl AgentClient {
async fn reserve_id(&self, tx: mpsc::Sender<InboundFrame>) -> AgentClientResult<u32> {
let id = self
.next_id
.fetch_update(
std::sync::atomic::Ordering::Relaxed,
std::sync::atomic::Ordering::Relaxed,
|next| (next < self.id_max).then_some(next.saturating_add(1)),
)
.map_err(|_| AgentClientError::IdRangeExhausted)?;
if id == 0 || id < self.id_min {
return Err(AgentClientError::IdRangeExhausted);
}
let replaced = self.pending.lock().await.insert(
id,
CorrelationRoute {
tx,
state: CorrelationState::Active,
},
);
debug_assert!(
replaced.is_none(),
"single-use correlation was already routed"
);
Ok(id)
}
async fn write_frame(&self, id: u32, flags: u8, body: &[u8]) -> AgentClientResult<()> {
self.write_frame_owned(id, flags, body.to_vec()).await
}
async fn write_frame_owned(&self, id: u32, flags: u8, body: Vec<u8>) -> AgentClientResult<()> {
let (ack, written) = oneshot::channel();
self.writer
.send(WriterCommand {
frame: WriterFrame::Control(RawFrame { id, flags, body }),
ack,
})
.await
.map_err(|_| AgentClientError::Closed)?;
written.await.map_err(|_| AgentClientError::Closed)?
}
}
#[cfg(all(feature = "uds", unix))]
async fn connect_local_stream(
sock_path: &Path,
_deadline: Instant,
) -> AgentClientResult<UnixStream> {
UnixStream::connect(sock_path)
.await
.map_err(|source| AgentClientError::Connect {
path: sock_path.to_path_buf(),
source,
})
}
#[cfg(all(feature = "named-pipe", windows))]
async fn connect_local_stream(
pipe_path: &Path,
deadline: Instant,
) -> AgentClientResult<tokio::net::windows::named_pipe::NamedPipeClient> {
loop {
match ClientOptions::new().open(pipe_path) {
Ok(stream) => return Ok(stream),
Err(source)
if is_retryable_named_pipe_connect_error(&source) && Instant::now() < deadline =>
{
tokio::time::sleep(WINDOWS_PIPE_CONNECT_RETRY).await;
}
Err(source) => {
return Err(AgentClientError::Connect {
path: pipe_path.to_path_buf(),
source,
});
}
}
}
}
#[cfg(all(feature = "named-pipe", windows))]
fn is_retryable_named_pipe_connect_error(error: &std::io::Error) -> bool {
const ERROR_PIPE_BUSY: i32 = 231;
error.kind() == std::io::ErrorKind::NotFound || error.raw_os_error() == Some(ERROR_PIPE_BUSY)
}
#[cfg(feature = "stream")]
async fn finish_connection<R, W>(
reader: R,
writer: W,
handshake: AgentHandshake,
#[cfg(all(feature = "uds", unix))] local: Option<LocalShmClient>,
) -> AgentClientResult<AgentClient>
where
R: AsyncRead + Unpin + Send + 'static,
W: AsyncWrite + Unpin + Send + 'static,
{
#[cfg(all(feature = "uds", unix))]
let local_shm = local.is_some();
#[cfg(not(all(feature = "uds", unix)))]
let local_shm = false;
tracing::info!(
id_min = handshake.id_min,
id_max = handshake.id_max,
protocol = ?handshake.protocol,
ready_bytes = handshake.ready_body.len(),
boot_time_ns = handshake.ready.boot_time_ns,
local_shm,
"agent client: connected to relay"
);
if handshake.protocol == AgentProtocol::LegacyV1 {
tracing::warn!(
"agent client: connected to a sandbox started before microsandbox 0.5; exec compatibility is temporary and filesystem/SFTP require stop/start"
);
}
let pending: Arc<Mutex<HashMap<u32, CorrelationRoute>>> = Arc::new(Mutex::new(HashMap::new()));
let (writer_tx, writer_rx) = mpsc::channel(WRITER_QUEUE_CAPACITY);
#[cfg(all(feature = "uds", unix))]
let (local_release_tx, local_release_rx) = mpsc::unbounded_channel();
#[cfg(all(feature = "uds", unix))]
let (connection_shutdown_tx, connection_shutdown_rx) = watch::channel(false);
#[cfg(all(feature = "uds", unix))]
let (local_inbound, local_outbound) = match local {
Some(local) => (Some(local.inbound), Some(local.outbound)),
None => (None, None),
};
#[cfg(all(feature = "uds", unix))]
let reader_handle = tokio::spawn(reader_loop(
reader,
Arc::clone(&pending),
local_inbound,
local_outbound.clone(),
local_release_tx,
connection_shutdown_tx.clone(),
connection_shutdown_rx.clone(),
));
#[cfg(not(all(feature = "uds", unix)))]
let reader_handle = tokio::spawn(reader_loop(reader, Arc::clone(&pending)));
#[cfg(all(feature = "uds", unix))]
let writer_handle = tokio::spawn(stream_writer_loop(
writer,
writer_rx,
local_release_rx,
local_outbound.clone(),
connection_shutdown_tx,
connection_shutdown_rx,
));
#[cfg(not(all(feature = "uds", unix)))]
let writer_handle = tokio::spawn(stream_writer_loop(writer, writer_rx));
Ok(AgentClient {
writer: writer_tx,
next_id: AtomicU32::new(first_request_id(handshake.id_min)),
id_min: handshake.id_min,
id_max: handshake.id_max,
protocol: handshake.protocol,
negotiated_version: handshake.negotiated_version,
pending,
reader_handle,
writer_handle,
ready_body: handshake.ready_body,
ready: handshake.ready,
#[cfg(all(feature = "uds", unix))]
local_outbound,
})
}
#[cfg(feature = "stream")]
async fn perform_handshake<R>(
reader: &mut R,
deadline: Instant,
) -> AgentClientResult<AgentHandshake>
where
R: HandshakeReader + ?Sized,
{
let mut range_buf = [0u8; 8];
tokio::time::timeout_at(deadline, reader.read_exact_handshake(&mut range_buf))
.await
.map_err(|_| {
AgentClientError::Handshake("read id range: timed out before relay sent bytes".into())
})??;
let id_start_or_offset = u32::from_be_bytes(range_buf[0..4].try_into().unwrap());
let id_max_or_frame_len = u32::from_be_bytes(range_buf[4..8].try_into().unwrap());
let legacy_handshake =
looks_like_legacy_relay_handshake(id_start_or_offset, id_max_or_frame_len);
let (id_min, id_max, ready_frame, protocol) = if legacy_handshake {
let id_offset = id_start_or_offset;
let ready_frame =
read_raw_frame_after_len_prefix(reader, range_buf[4..8].try_into().unwrap(), deadline)
.await?;
(
id_offset.saturating_add(1),
id_offset.saturating_add(LEGACY_RELAY_ID_RANGE_STEP),
ready_frame,
AgentProtocol::LegacyV1,
)
} else if id_start_or_offset >= id_max_or_frame_len {
return Err(AgentClientError::Handshake(format!(
"invalid relay id range: start={id_start_or_offset}, end={id_max_or_frame_len}"
)));
} else {
let ready_frame = tokio::time::timeout_at(deadline, reader.read_frame_handshake())
.await
.map_err(|_| {
AgentClientError::Handshake(
"read ready frame: timed out before relay sent frame".into(),
)
})?
.map_err(|e| AgentClientError::Handshake(format!("read ready frame: {e}")))?;
(
id_start_or_offset,
id_max_or_frame_len,
ready_frame,
AgentProtocol::Current,
)
};
ensure_usable_id_range(id_min, id_max)?;
let ready_msg = codec::raw_frame_to_message(ready_frame.clone())
.map_err(|e| AgentClientError::Handshake(format!("decode ready frame: {e}")))?;
if ready_msg.t != MessageType::Ready {
return Err(AgentClientError::Handshake(format!(
"expected core.ready frame, got {}",
ready_msg.t.as_str()
)));
}
let ready: Ready = ready_msg
.payload()
.map_err(|e| AgentClientError::Handshake(format!("decode ready payload: {e}")))?;
let negotiated_version = protocol.version().min(ready_msg.v);
Ok(AgentHandshake {
id_min,
id_max,
protocol,
negotiated_version,
ready_body: ready_frame.body,
ready,
})
}
fn first_request_id(id_min: u32) -> u32 {
id_min.max(1)
}
#[cfg(feature = "stream")]
fn ensure_usable_id_range(id_min: u32, id_max: u32) -> AgentClientResult<()> {
if usable_id_count(id_min, id_max) == 0 {
return Err(AgentClientError::Handshake(format!(
"relay id range contains no usable nonzero ids: start={id_min}, end={id_max}"
)));
}
Ok(())
}
fn usable_id_count(id_min: u32, id_max: u32) -> u32 {
id_max.saturating_sub(first_request_id(id_min))
}
#[cfg(feature = "stream")]
fn looks_like_legacy_relay_handshake(id_min: u32, id_max: u32) -> bool {
id_max >= FRAME_HEADER_SIZE as u32
&& id_max <= MAX_FRAME_SIZE
&& (id_min == 0 || id_min >= id_max)
}
#[cfg(feature = "stream")]
async fn read_raw_frame_after_len_prefix<R>(
reader: &mut R,
len_buf: [u8; 4],
deadline: Instant,
) -> AgentClientResult<RawFrame>
where
R: HandshakeReader + ?Sized,
{
let frame_len = u32::from_be_bytes(len_buf);
if frame_len > MAX_FRAME_SIZE {
return Err(AgentClientError::Handshake(format!(
"legacy ready frame too large: {frame_len} bytes (max {MAX_FRAME_SIZE})"
)));
}
if frame_len < FRAME_HEADER_SIZE as u32 {
return Err(AgentClientError::Handshake(format!(
"legacy ready frame too short: {frame_len} bytes"
)));
}
let mut data = vec![0u8; frame_len as usize];
tokio::time::timeout_at(deadline, reader.read_exact_handshake(&mut data))
.await
.map_err(|_| {
AgentClientError::Handshake(
"read legacy ready frame: timed out before relay sent frame".into(),
)
})?
.map_err(|e| AgentClientError::Handshake(format!("read legacy ready frame: {e}")))?;
let id = u32::from_be_bytes(data[0..4].try_into().unwrap());
let flags = data[4];
let body = data[FRAME_HEADER_SIZE..].to_vec();
Ok(RawFrame { id, flags, body })
}
#[cfg(feature = "stream")]
impl<R> HandshakeReader for R
where
R: tokio::io::AsyncRead + Unpin + Send,
{
fn read_exact_handshake<'a>(
&'a mut self,
out: &'a mut [u8],
) -> Pin<Box<dyn Future<Output = AgentClientResult<()>> + Send + 'a>> {
Box::pin(async move {
tokio::io::AsyncReadExt::read_exact(self, out)
.await
.map(|_| ())
.map_err(|e| AgentClientError::Handshake(e.to_string()))
})
}
fn read_frame_handshake<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = AgentClientResult<RawFrame>> + Send + 'a>> {
Box::pin(async move {
codec::read_raw_frame(self)
.await
.map_err(AgentClientError::Protocol)
})
}
}
#[cfg(all(feature = "stream", feature = "uds", unix))]
async fn stream_writer_loop<W>(
mut writer: W,
mut rx: mpsc::Receiver<WriterCommand>,
mut local_release_rx: mpsc::UnboundedReceiver<LocalBulkRelease>,
local_outbound: Option<SharedArenaProducer>,
connection_shutdown_tx: watch::Sender<bool>,
mut connection_shutdown_rx: watch::Receiver<bool>,
) where
W: tokio::io::AsyncWrite + Unpin,
{
loop {
tokio::select! {
biased;
changed = connection_shutdown_rx.changed() => {
if changed.is_err() || *connection_shutdown_rx.borrow() {
break;
}
}
release = local_release_rx.recv() => {
let Some(release) = release else { break; };
let result = async {
let wire = encode_local_bulk_release(release)
.map_err(|error| AgentClientError::LocalTransport(error.to_string()))?;
tokio::io::AsyncWriteExt::write_all(&mut writer, &wire).await?;
tokio::io::AsyncWriteExt::flush(&mut writer).await?;
AgentClientResult::Ok(())
}.await;
if let Err(error) = result {
tracing::debug!("agent client: local release writer error: {error}");
break;
}
}
command = rx.recv() => {
let Some(mut command) = command else { break; };
let result = write_writer_command(&mut writer, &mut command).await;
if let Err(error) = result {
tracing::debug!("agent client: stream writer error: {error}");
let _ = command.ack.send(Err(error));
break;
}
let _ = command.ack.send(Ok(()));
}
}
}
if let Some(producer) = local_outbound.as_ref() {
producer.close();
}
let _ = connection_shutdown_tx.send(true);
}
#[cfg(all(feature = "stream", not(all(feature = "uds", unix))))]
async fn stream_writer_loop<W>(mut writer: W, mut rx: mpsc::Receiver<WriterCommand>)
where
W: tokio::io::AsyncWrite + Unpin,
{
while let Some(mut command) = rx.recv().await {
let result = write_writer_command(&mut writer, &mut command).await;
if let Err(error) = result {
tracing::debug!("agent client: stream writer error: {error}");
let _ = command.ack.send(Err(error));
break;
}
let _ = command.ack.send(Ok(()));
}
}
#[cfg(feature = "stream")]
async fn write_writer_command<W>(
writer: &mut W,
command: &mut WriterCommand,
) -> AgentClientResult<()>
where
W: tokio::io::AsyncWrite + Unpin,
{
match &mut command.frame {
WriterFrame::Control(frame) => codec::write_raw_frame(writer, frame).await?,
WriterFrame::Bulk(record) => codec::write_bulk_record(writer, record).await?,
#[cfg(all(feature = "uds", unix))]
WriterFrame::LocalBulk(prepared) => {
let wire = encode_local_bulk_ref(prepared.descriptor())
.map_err(|error| AgentClientError::LocalTransport(error.to_string()))?;
tokio::io::AsyncWriteExt::write_all(writer, &wire).await?;
tokio::io::AsyncWriteExt::flush(writer).await?;
prepared.commit();
}
}
Ok(())
}
#[cfg(all(feature = "stream", feature = "uds", unix))]
async fn reader_loop<R>(
mut reader: R,
pending: Arc<Mutex<HashMap<u32, CorrelationRoute>>>,
local_inbound: Option<SharedArenaConsumer>,
local_outbound: Option<SharedArenaProducer>,
local_release_tx: mpsc::UnboundedSender<LocalBulkRelease>,
connection_shutdown_tx: watch::Sender<bool>,
mut connection_shutdown_rx: watch::Receiver<bool>,
) where
R: tokio::io::AsyncRead + Unpin,
{
loop {
let frame = match tokio::select! {
changed = connection_shutdown_rx.changed() => {
if changed.is_err() || *connection_shutdown_rx.borrow() {
break;
}
continue;
}
frame = codec::read_raw_frame(&mut reader) => frame,
} {
Ok(frame) => frame,
Err(e) => {
tracing::debug!("agent client: reader EOF or error: {e}");
break;
}
};
if frame.id == 0 && frame.flags == 0 {
let local = match decode_local_body(&frame.body) {
Ok(local) => local,
Err(error) => {
tracing::debug!("agent client: malformed local frame: {error}");
break;
}
};
match local {
LocalShmFrame::BulkRef(descriptor) => {
let Some(consumer) = local_inbound.as_ref() else {
tracing::debug!(
"agent client: local bulk reference without negotiated arena"
);
break;
};
let record = match consumer.receive(descriptor, local_release_tx.clone()) {
Ok(record) => record,
Err(error) => {
tracing::debug!("agent client: rejected local bulk reference: {error}");
break;
}
};
dispatch_frame(InboundFrame::Bulk(record), &pending).await;
}
LocalShmFrame::BulkRelease(release) => {
let Some(producer) = local_outbound.as_ref() else {
tracing::debug!("agent client: local release without negotiated arena");
break;
};
if let Err(error) = producer.release(release) {
tracing::debug!("agent client: rejected local bulk release: {error}");
break;
}
}
LocalShmFrame::UpgradeRequest => {
tracing::debug!("agent client: relay sent an invalid upgrade request");
break;
}
}
continue;
}
dispatch_frame(InboundFrame::Raw(frame), &pending).await;
}
if let Some(producer) = local_outbound.as_ref() {
producer.close();
}
let _ = connection_shutdown_tx.send(true);
let mut map = pending.lock().await;
map.clear();
}
#[cfg(all(feature = "stream", not(all(feature = "uds", unix))))]
async fn reader_loop<R>(mut reader: R, pending: Arc<Mutex<HashMap<u32, CorrelationRoute>>>)
where
R: tokio::io::AsyncRead + Unpin,
{
loop {
let frame = match codec::read_raw_frame(&mut reader).await {
Ok(frame) => frame,
Err(error) => {
tracing::debug!("agent client: reader EOF or error: {error}");
break;
}
};
dispatch_frame(InboundFrame::Raw(frame), &pending).await;
}
pending.lock().await.clear();
}
#[cfg(feature = "stream")]
async fn dispatch_frame(frame: InboundFrame, pending: &Arc<Mutex<HashMap<u32, CorrelationRoute>>>) {
let id = frame.id();
let flags = frame.flags();
let is_terminal = (flags & FLAG_TERMINAL) != 0;
let tx = {
let mut map = pending.lock().await;
let Some(route) = map.get(&id) else {
tracing::trace!("agent client: no pending handler for id={id}");
return;
};
if route.state == CorrelationState::Cancelling && flags == FLAG_BULK {
return;
}
let tx = route.tx.clone();
if is_terminal {
map.remove(&id);
}
tx
};
if tx.send(frame).await.is_err() {
pending.lock().await.remove(&id);
}
}
async fn decode_stream_task(mut raw_rx: mpsc::Receiver<InboundFrame>, tx: mpsc::Sender<Message>) {
while let Some(frame) = raw_rx.recv().await {
#[cfg(all(feature = "uds", unix))]
let frame = match frame {
InboundFrame::Raw(frame) => frame,
InboundFrame::Bulk(_) => {
tracing::warn!("agent client: raw bulk record reached a control-only stream");
break;
}
};
#[cfg(not(all(feature = "uds", unix)))]
let InboundFrame::Raw(frame) = frame;
if frame.flags & FLAG_BULK != 0 {
tracing::warn!("agent client: raw bulk record reached a control-only stream");
break;
}
match codec::raw_frame_to_message(frame) {
Ok(msg) => {
if tx.send(msg).await.is_err() {
break;
}
}
Err(e) => {
tracing::warn!("agent client: failed to decode frame in stream: {e}");
}
}
}
}
async fn decode_frame_stream_task(
mut raw_rx: mpsc::Receiver<InboundFrame>,
tx: mpsc::Sender<AgentFrame>,
) {
while let Some(frame) = raw_rx.recv().await {
let decoded = match frame {
#[cfg(all(feature = "uds", unix))]
InboundFrame::Bulk(record) => Ok(AgentFrame::Bulk(record)),
InboundFrame::Raw(frame) if frame.flags & FLAG_BULK != 0 => {
codec::raw_frame_to_bulk(frame, MAX_BULK_RECORD_PAYLOAD).map(AgentFrame::Bulk)
}
InboundFrame::Raw(frame) => codec::raw_frame_to_message(frame).map(AgentFrame::Control),
};
match decoded {
Ok(frame) => {
if tx.send(frame).await.is_err() {
break;
}
}
Err(error) => {
tracing::warn!("agent client: failed to decode frame in bulk stream: {error}");
break;
}
}
}
}
async fn materialize_raw_stream_task(
mut inbound_rx: mpsc::Receiver<InboundFrame>,
tx: mpsc::Sender<RawFrame>,
) {
while let Some(frame) = inbound_rx.recv().await {
let Ok(frame) = frame.into_raw_frame() else {
break;
};
if tx.send(frame).await.is_err() {
break;
}
}
}
fn encode_message_body<T: Serialize>(
version: u8,
t: MessageType,
payload: &T,
) -> AgentClientResult<Vec<u8>> {
let mut msg = Message::with_payload(t, 0, payload)?;
msg.v = version;
let mut body = Vec::new();
ciborium::into_writer(&msg, &mut body).map_err(microsandbox_protocol::ProtocolError::from)?;
Ok(body)
}
impl Drop for AgentClient {
fn drop(&mut self) {
self.reader_handle.abort();
self.writer_handle.abort();
}
}
#[cfg(test)]
mod tests {
#[cfg(all(feature = "uds", unix))]
use crate::local_shm::{LocalShmServer, send_local_shm_upgrade};
#[cfg(all(feature = "uds", unix))]
use bytes::Bytes;
#[cfg(all(feature = "uds", unix))]
use microsandbox_protocol::core::Ready;
#[cfg(all(feature = "uds", unix))]
use microsandbox_protocol::exec::ExecRequest;
#[cfg(all(feature = "uds", unix))]
use microsandbox_protocol::message::PROTOCOL_VERSION;
#[cfg(all(feature = "uds", unix))]
use tokio::io::AsyncWriteExt;
#[cfg(all(feature = "uds", unix))]
use tokio::net::UnixListener;
#[cfg(all(feature = "uds", unix))]
use tokio::sync::oneshot;
use super::*;
#[cfg(all(feature = "uds", unix))]
#[tokio::test]
async fn connect_decodes_ready_payload() {
let temp = tempfile::tempdir().unwrap();
let sock_path = temp.path().join("agent.sock");
let listener = UnixListener::bind(&sock_path).unwrap();
let ready = Ready {
boot_time_ns: 11,
init_time_ns: 22,
ready_time_ns: 33,
agent_version: "9.9.9".to_string(),
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&1u32.to_be_bytes()).await.unwrap();
socket.write_all(&8u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut socket, &ready_msg).await.unwrap();
});
let client =
AgentClient::connect_with_deadline(&sock_path, Instant::now() + Duration::from_secs(1))
.await
.unwrap();
assert_eq!(client.protocol(), AgentProtocol::Current);
assert_eq!(client.negotiated_version(), PROTOCOL_VERSION);
assert!(client.supports(MessageType::FsRequest));
assert_eq!(client.agent_version(), "9.9.9");
let decoded = client.ready().unwrap();
assert_eq!(decoded.boot_time_ns, ready.boot_time_ns);
assert_eq!(decoded.init_time_ns, ready.init_time_ns);
assert_eq!(decoded.ready_time_ns, ready.ready_time_ns);
let raw_msg: Message = ciborium::from_reader(client.ready_bytes()).unwrap();
assert_eq!(raw_msg.t, MessageType::Ready);
}
#[cfg(all(feature = "uds", unix))]
#[tokio::test]
async fn connect_selects_shared_arena_and_sends_bulk_by_reference() {
let temp = tempfile::tempdir().unwrap();
let sock_path = temp.path().join("agent.sock");
let listener = UnixListener::bind(&sock_path).unwrap();
let ready = Ready {
local_transport: Some(
microsandbox_protocol::transport::LocalTransportReady::shared_arena_v1(),
),
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
let expected = Bytes::from_static(b"shared payload, socket descriptor");
let expected_server = expected.clone();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&1u32.to_be_bytes()).await.unwrap();
socket.write_all(&1024u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut socket, &ready_msg).await.unwrap();
let upgrade = codec::read_raw_frame(&mut socket).await.unwrap();
assert_eq!(upgrade.id, 0);
assert_eq!(upgrade.flags, 0);
assert_eq!(
decode_local_body(&upgrade.body).unwrap(),
LocalShmFrame::UpgradeRequest
);
let arenas = LocalShmServer::create().unwrap();
send_local_shm_upgrade(&socket, Some(arenas.client_fds()))
.await
.unwrap();
let descriptor = codec::read_raw_frame(&mut socket).await.unwrap();
let LocalShmFrame::BulkRef(descriptor) = decode_local_body(&descriptor.body).unwrap()
else {
panic!("client sent bulk bytes in-band after selecting the shared arena");
};
let (release_tx, _release_rx) = mpsc::unbounded_channel();
let record = arenas.inbound.receive(descriptor, release_tx).unwrap();
assert_eq!(record.payload, expected_server);
});
let client =
AgentClient::connect_with_deadline(&sock_path, Instant::now() + Duration::from_secs(1))
.await
.unwrap();
client
.send_bulk(BulkRecord {
id: 1,
kind: microsandbox_protocol::bulk::BulkKind::Filesystem,
flow: microsandbox_protocol::bulk::BulkFlow::HostToGuest,
offset: 0,
payload: expected,
})
.await
.unwrap();
server.await.unwrap();
}
#[cfg(all(feature = "named-pipe", windows))]
#[tokio::test]
async fn connect_decodes_ready_payload_from_named_pipe() {
use microsandbox_protocol::core::Ready;
use microsandbox_protocol::message::PROTOCOL_VERSION;
use tokio::io::AsyncWriteExt;
use tokio::net::windows::named_pipe::{PipeMode, ServerOptions};
let pipe_path = unique_named_pipe("ready");
let server = ServerOptions::new()
.first_pipe_instance(true)
.pipe_mode(PipeMode::Byte)
.create(&pipe_path)
.unwrap();
let ready = Ready {
boot_time_ns: 11,
init_time_ns: 22,
ready_time_ns: 33,
agent_version: "named-pipe-test".to_string(),
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
tokio::spawn(async move {
let mut server = server;
server.connect().await.unwrap();
server.write_all(&1u32.to_be_bytes()).await.unwrap();
server.write_all(&8u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut server, &ready_msg).await.unwrap();
});
let client = AgentClient::connect_with_deadline(
std::path::Path::new(&pipe_path),
Instant::now() + Duration::from_secs(1),
)
.await
.unwrap();
assert_eq!(client.protocol(), AgentProtocol::Current);
assert_eq!(client.negotiated_version(), PROTOCOL_VERSION);
assert_eq!(client.agent_version(), "named-pipe-test");
let decoded = client.ready().unwrap();
assert_eq!(decoded.boot_time_ns, ready.boot_time_ns);
assert_eq!(decoded.init_time_ns, ready.init_time_ns);
assert_eq!(decoded.ready_time_ns, ready.ready_time_ns);
}
#[cfg(all(feature = "uds", unix))]
#[tokio::test]
async fn connect_negotiates_down_to_older_guest_generation() {
let temp = tempfile::tempdir().unwrap();
let sock_path = temp.path().join("agent.sock");
let listener = UnixListener::bind(&sock_path).unwrap();
let ready = Ready {
boot_time_ns: 1,
init_time_ns: 2,
ready_time_ns: 3,
..Default::default()
};
let mut ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
ready_msg.v = 1;
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&1u32.to_be_bytes()).await.unwrap();
socket
.write_all(µsandbox_protocol::AGENT_RELAY_ID_RANGE_STEP.to_be_bytes())
.await
.unwrap();
codec::write_message(&mut socket, &ready_msg).await.unwrap();
});
let client =
AgentClient::connect_with_deadline(&sock_path, Instant::now() + Duration::from_secs(1))
.await
.unwrap();
assert_eq!(client.protocol(), AgentProtocol::Current);
assert_eq!(client.negotiated_version(), 1);
assert!(client.supports(MessageType::ExecRequest));
assert!(!client.supports(MessageType::FsRequest));
}
#[cfg(all(feature = "uds", unix))]
#[tokio::test]
async fn connect_accepts_legacy_relay_handshake() {
assert_accepts_legacy_relay_handshake(0).await;
assert_accepts_legacy_relay_handshake(268_435_455).await;
}
#[cfg(all(feature = "uds", unix))]
#[tokio::test]
async fn legacy_relay_requests_use_v1_and_legacy_id_range() {
let temp = tempfile::tempdir().unwrap();
let sock_path = temp.path().join("agent.sock");
let listener = UnixListener::bind(&sock_path).unwrap();
let ready = Ready {
boot_time_ns: 11,
init_time_ns: 22,
ready_time_ns: 33,
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
let id_offset = 268_435_455u32;
let (frame_tx, frame_rx) = oneshot::channel();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&id_offset.to_be_bytes()).await.unwrap();
codec::write_message(&mut socket, &ready_msg).await.unwrap();
let frame = codec::read_raw_frame(&mut socket).await.unwrap();
frame_tx.send(frame).unwrap();
});
let client =
AgentClient::connect_with_deadline(&sock_path, Instant::now() + Duration::from_secs(1))
.await
.unwrap();
let request = ExecRequest {
cmd: "/bin/true".into(),
args: Vec::new(),
env: Vec::new(),
cwd: None,
user: None,
tty: false,
rows: 24,
cols: 80,
rlimits: Vec::new(),
};
let (id, _rx) = client
.stream(MessageType::ExecRequest, &request)
.await
.unwrap();
let frame = frame_rx.await.unwrap();
let message = codec::raw_frame_to_message(frame).unwrap();
assert_eq!(id, id_offset + 1);
assert_eq!(message.id, id_offset + 1);
assert_eq!(message.v, LEGACY_PROTOCOL_VERSION);
assert_eq!(message.t, MessageType::ExecRequest);
}
#[test]
fn version_compat_across_generations() {
use MessageType::{BulkAccepted, ExecRequest, FsRequest};
let cases = [
(ExecRequest, 1, true),
(ExecRequest, 2, true),
(ExecRequest, 3, true),
(FsRequest, 1, false),
(FsRequest, 2, true),
(FsRequest, 3, true),
(BulkAccepted, 7, false),
(BulkAccepted, 8, true),
(MessageType::WorkloadFreeze, 8, false),
(MessageType::WorkloadThaw, 8, false),
(MessageType::WorkloadFreeze, 9, true),
(MessageType::WorkloadThaw, 9, true),
];
for (t, generation, allowed) in cases {
assert_eq!(
AgentClient::ensure_version_compat_for(t, generation).is_ok(),
allowed,
"{t:?} at generation {generation}"
);
}
}
#[test]
fn version_compat_rejection_is_typed() {
let err =
AgentClient::ensure_version_compat_for(MessageType::FsRequest, LEGACY_PROTOCOL_VERSION)
.unwrap_err();
assert!(matches!(
err,
AgentClientError::UnsupportedOperation {
needs: 2,
peer: 1,
..
}
));
}
#[cfg(all(feature = "uds", unix))]
#[tokio::test]
async fn connect_preserves_current_peer_protocol_version() {
let temp = tempfile::tempdir().unwrap();
let sock_path = temp.path().join("agent.sock");
let listener = UnixListener::bind(&sock_path).unwrap();
let ready = Ready {
boot_time_ns: 11,
init_time_ns: 22,
ready_time_ns: 33,
..Default::default()
};
let mut ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
ready_msg.v = 2;
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&1u32.to_be_bytes()).await.unwrap();
socket
.write_all(µsandbox_protocol::AGENT_RELAY_ID_RANGE_STEP.to_be_bytes())
.await
.unwrap();
codec::write_message(&mut socket, &ready_msg).await.unwrap();
});
let client =
AgentClient::connect_with_deadline(&sock_path, Instant::now() + Duration::from_secs(1))
.await
.unwrap();
assert_eq!(client.protocol(), AgentProtocol::Current);
assert_eq!(client.negotiated_version(), 2);
assert!(!client.supports(MessageType::TcpConnect));
}
#[cfg(all(feature = "uds", unix))]
async fn assert_accepts_legacy_relay_handshake(id_offset: u32) {
let temp = tempfile::tempdir().unwrap();
let sock_path = temp.path().join("agent.sock");
let listener = UnixListener::bind(&sock_path).unwrap();
let ready = Ready {
boot_time_ns: 11,
init_time_ns: 22,
ready_time_ns: 33,
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
socket.write_all(&id_offset.to_be_bytes()).await.unwrap();
codec::write_message(&mut socket, &ready_msg).await.unwrap();
});
let client =
AgentClient::connect_with_deadline(&sock_path, Instant::now() + Duration::from_secs(1))
.await
.unwrap();
assert_eq!(client.protocol(), AgentProtocol::LegacyV1);
assert_eq!(client.negotiated_version(), LEGACY_PROTOCOL_VERSION);
let decoded = client.ready().unwrap();
assert_eq!(decoded.boot_time_ns, ready.boot_time_ns);
assert_eq!(decoded.init_time_ns, ready.init_time_ns);
assert_eq!(decoded.ready_time_ns, ready.ready_time_ns);
}
#[cfg(all(feature = "named-pipe", windows))]
fn unique_named_pipe(name: &str) -> String {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
format!(
r"\\.\pipe\msb-agent-client-{name}-{}-{nanos}",
std::process::id()
)
}
#[cfg(feature = "stream")]
#[tokio::test]
async fn connect_stream_handshakes_and_streams_exec() {
use microsandbox_protocol::exec::{ExecExited, ExecRequest, ExecStdout};
use tokio::io::AsyncWriteExt;
let (client_io, mut server_io) = tokio::io::duplex(64 * 1024);
let ready = Ready {
boot_time_ns: 11,
init_time_ns: 22,
ready_time_ns: 33,
agent_version: "stream-test".to_string(),
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
tokio::spawn(async move {
server_io.write_all(&1u32.to_be_bytes()).await.unwrap();
server_io.write_all(&1024u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut server_io, &ready_msg)
.await
.unwrap();
let request = codec::read_raw_frame(&mut server_io).await.unwrap();
let stdout = Message::with_payload(
MessageType::ExecStdout,
request.id,
&ExecStdout {
data: b"hi".to_vec(),
},
)
.unwrap();
codec::write_message(&mut server_io, &stdout).await.unwrap();
let exited =
Message::with_payload(MessageType::ExecExited, request.id, &ExecExited { code: 0 })
.unwrap();
codec::write_message(&mut server_io, &exited).await.unwrap();
});
let client = AgentClient::connect_stream_with_deadline(
client_io,
Instant::now() + Duration::from_secs(1),
)
.await
.unwrap();
assert_eq!(client.protocol(), AgentProtocol::Current);
assert_eq!(client.agent_version(), "stream-test");
assert!(client.supports(MessageType::ExecRequest));
let request = ExecRequest {
cmd: "echo".into(),
args: vec!["hi".into()],
env: Vec::new(),
cwd: None,
user: None,
tty: false,
rows: 24,
cols: 80,
rlimits: Vec::new(),
};
let (_id, mut rx) = client
.stream(MessageType::ExecRequest, &request)
.await
.unwrap();
let first = rx.recv().await.unwrap();
assert_eq!(first.t, MessageType::ExecStdout);
let out: ExecStdout = first.payload().unwrap();
assert_eq!(out.data, b"hi");
let second = rx.recv().await.unwrap();
assert_eq!(second.t, MessageType::ExecExited);
let exit: ExecExited = second.payload().unwrap();
assert_eq!(exit.code, 0);
}
#[cfg(feature = "stream")]
#[tokio::test]
async fn correlation_ids_are_single_use_until_reconnect() {
use microsandbox_protocol::core::{Ping, Pong};
use tokio::io::AsyncWriteExt;
let (client_io, mut server_io) = tokio::io::duplex(64 * 1024);
let ready_msg = Message::with_payload(MessageType::Ready, 0, &Ready::default()).unwrap();
let server = tokio::spawn(async move {
server_io.write_all(&1u32.to_be_bytes()).await.unwrap();
server_io.write_all(&3u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut server_io, &ready_msg)
.await
.unwrap();
for expected_id in 1..3 {
let request = codec::read_raw_frame(&mut server_io).await.unwrap();
assert_eq!(request.id, expected_id);
let response =
Message::with_payload(MessageType::Pong, request.id, &Pong {}).unwrap();
codec::write_message(&mut server_io, &response)
.await
.unwrap();
}
});
let client = AgentClient::connect_stream(client_io).await.unwrap();
client.request(MessageType::Ping, &Ping {}).await.unwrap();
client.request(MessageType::Ping, &Ping {}).await.unwrap();
assert!(matches!(
client.request(MessageType::Ping, &Ping {}).await,
Err(AgentClientError::IdRangeExhausted)
));
server.await.unwrap();
}
#[cfg(feature = "stream")]
#[tokio::test]
async fn connect_stream_carries_bidirectional_raw_bulk_records() {
use microsandbox_protocol::bulk::{
BULK_FLOW_MASK_GUEST_TO_HOST, BULK_FLOW_MASK_HOST_TO_GUEST, BulkAccepted, BulkFinish,
BulkFlow, BulkKind, BulkOffer, DEFAULT_BULK_RECORD_PAYLOAD, DEFAULT_BULK_WINDOW,
};
use microsandbox_protocol::tcp::{TcpClosed, TcpConnect, TcpConnected};
use tokio::io::AsyncWriteExt;
let (client_io, mut server_io) = tokio::io::duplex(1024 * 1024);
let ready = Ready {
agent_version: "bulk-stream-test".to_string(),
..Default::default()
};
let ready_msg = Message::with_payload(MessageType::Ready, 0, &ready).unwrap();
let server = tokio::spawn(async move {
server_io.write_all(&1u32.to_be_bytes()).await.unwrap();
server_io.write_all(&1024u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut server_io, &ready_msg)
.await
.unwrap();
let opening = codec::read_raw_frame(&mut server_io).await.unwrap();
let opening_id = opening.id;
let opening = codec::raw_frame_to_message(opening).unwrap();
assert_eq!(opening.t, MessageType::TcpConnect);
let connected =
Message::with_payload(MessageType::TcpConnected, opening_id, &TcpConnected {})
.unwrap();
codec::write_message(&mut server_io, &connected)
.await
.unwrap();
let accepted = Message::with_payload(
MessageType::BulkAccepted,
opening_id,
&BulkAccepted {
kind: BulkKind::Tcp,
flows: BULK_FLOW_MASK_HOST_TO_GUEST | BULK_FLOW_MASK_GUEST_TO_HOST,
format: 1,
max_record_payload: DEFAULT_BULK_RECORD_PAYLOAD,
host_to_guest_credit_limit: DEFAULT_BULK_WINDOW,
guest_to_host_credit_limit: DEFAULT_BULK_WINDOW,
},
)
.unwrap();
codec::write_message(&mut server_io, &accepted)
.await
.unwrap();
let first = codec::read_raw_frame(&mut server_io).await.unwrap();
let first = codec::raw_frame_to_bulk(first, DEFAULT_BULK_RECORD_PAYLOAD).unwrap();
let second = codec::read_raw_frame(&mut server_io).await.unwrap();
let second = codec::raw_frame_to_bulk(second, DEFAULT_BULK_RECORD_PAYLOAD).unwrap();
assert_eq!(first.id, opening_id);
assert_eq!(first.flow, BulkFlow::HostToGuest);
assert_eq!(first.offset, 0);
assert_eq!(first.payload.as_ref(), b"host-to-");
assert_eq!(second.id, opening_id);
assert_eq!(second.flow, BulkFlow::HostToGuest);
assert_eq!(second.offset, 8);
assert_eq!(second.payload.as_ref(), b"guest");
codec::write_bulk_record(
&mut server_io,
&BulkRecord {
id: opening_id,
kind: BulkKind::Tcp,
flow: BulkFlow::GuestToHost,
offset: 0,
payload: b"guest-to-host".as_slice().into(),
},
)
.await
.unwrap();
let finish = Message::with_payload(
MessageType::BulkFinish,
opening_id,
&BulkFinish {
kind: BulkKind::Tcp,
flow: BulkFlow::GuestToHost,
final_offset: 13,
},
)
.unwrap();
codec::write_message(&mut server_io, &finish).await.unwrap();
let closed =
Message::with_payload(MessageType::TcpClosed, opening_id, &TcpClosed {}).unwrap();
codec::write_message(&mut server_io, &closed).await.unwrap();
});
let client = AgentClient::connect_stream_with_deadline(
client_io,
Instant::now() + Duration::from_secs(1),
)
.await
.unwrap();
let offer = BulkOffer::tcp();
let (id, mut rx) = client
.stream_frames(
MessageType::TcpConnect,
&TcpConnect {
host: "example.test".into(),
port: 80,
bulk: Some(offer),
},
)
.await
.unwrap();
assert!(
matches!(rx.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::TcpConnected)
);
assert!(
matches!(rx.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::BulkAccepted)
);
client
.send_bulk(BulkRecord {
id,
kind: BulkKind::Tcp,
flow: BulkFlow::HostToGuest,
offset: 0,
payload: b"host-to-".as_slice().into(),
})
.await
.unwrap();
client
.send_bulk(BulkRecord {
id,
kind: BulkKind::Tcp,
flow: BulkFlow::HostToGuest,
offset: 8,
payload: b"guest".as_slice().into(),
})
.await
.unwrap();
let Some(AgentFrame::Bulk(record)) = rx.recv().await else {
panic!("expected raw bulk record");
};
assert_eq!(record.payload.as_ref(), b"guest-to-host");
assert!(
matches!(rx.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::BulkFinish)
);
assert!(
matches!(rx.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::TcpClosed)
);
server.await.unwrap();
}
#[cfg(feature = "stream")]
#[tokio::test]
async fn bulk_cancel_discards_late_raw_but_retains_terminal_route() {
use microsandbox_protocol::bulk::{
BULK_FLOW_MASK_GUEST_TO_HOST, BulkAccepted, BulkCancelReason, BulkFlow, BulkKind,
BulkOffer, DEFAULT_BULK_RECORD_PAYLOAD, DEFAULT_BULK_WINDOW,
};
use microsandbox_protocol::tcp::{TcpClosed, TcpConnect, TcpConnected};
use tokio::io::AsyncWriteExt;
let (client_io, mut server_io) = tokio::io::duplex(1024 * 1024);
let (late_sent, late_observed) = tokio::sync::oneshot::channel();
let (send_terminal, terminal_allowed) = tokio::sync::oneshot::channel();
let ready_msg = Message::with_payload(MessageType::Ready, 0, &Ready::default()).unwrap();
let server = tokio::spawn(async move {
server_io.write_all(&1u32.to_be_bytes()).await.unwrap();
server_io.write_all(&1024u32.to_be_bytes()).await.unwrap();
codec::write_message(&mut server_io, &ready_msg)
.await
.unwrap();
let opening = codec::read_raw_frame(&mut server_io).await.unwrap();
let id = opening.id;
let connected =
Message::with_payload(MessageType::TcpConnected, id, &TcpConnected {}).unwrap();
codec::write_message(&mut server_io, &connected)
.await
.unwrap();
let accepted = Message::with_payload(
MessageType::BulkAccepted,
id,
&BulkAccepted {
kind: BulkKind::Tcp,
flows: BULK_FLOW_MASK_GUEST_TO_HOST,
format: 1,
max_record_payload: DEFAULT_BULK_RECORD_PAYLOAD,
host_to_guest_credit_limit: 0,
guest_to_host_credit_limit: DEFAULT_BULK_WINDOW,
},
)
.unwrap();
codec::write_message(&mut server_io, &accepted)
.await
.unwrap();
let cancel = codec::read_raw_frame(&mut server_io).await.unwrap();
assert_eq!(cancel.id, id);
assert_eq!(
codec::raw_frame_to_message(cancel).unwrap().t,
MessageType::BulkCancel
);
codec::write_bulk_record(
&mut server_io,
&BulkRecord {
id,
kind: BulkKind::Tcp,
flow: BulkFlow::GuestToHost,
offset: 0,
payload: b"late".as_slice().into(),
},
)
.await
.unwrap();
let _ = late_sent.send(());
let _ = terminal_allowed.await;
let closed = Message::with_payload(MessageType::TcpClosed, id, &TcpClosed {}).unwrap();
codec::write_message(&mut server_io, &closed).await.unwrap();
});
let client = AgentClient::connect_stream(client_io).await.unwrap();
let (id, mut frames) = client
.stream_frames(
MessageType::TcpConnect,
&TcpConnect {
host: "example.test".into(),
port: 80,
bulk: Some(BulkOffer::tcp()),
},
)
.await
.unwrap();
assert!(
matches!(frames.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::TcpConnected)
);
assert!(
matches!(frames.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::BulkAccepted)
);
client
.cancel_bulk(
id,
&BulkCancel {
kind: BulkKind::Tcp,
reason: BulkCancelReason::CallerCancelled,
message: "test cancellation".into(),
},
)
.await
.unwrap();
late_observed.await.unwrap();
assert!(client.pending.lock().await.contains_key(&id));
let _ = send_terminal.send(());
assert!(
matches!(frames.recv().await, Some(AgentFrame::Control(message)) if message.t == MessageType::TcpClosed)
);
assert!(!client.pending.lock().await.contains_key(&id));
server.await.unwrap();
}
}