use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use arcbox_connect::v1::SandboxStreamWindow;
use arcbox_constants::wire::{MessageType, SANDBOX_STREAM_WINDOW};
use arcbox_transport::vsock::{Credit, VsockReceiver, VsockSender};
use buffa::Message as _;
use bytes::Bytes;
use tokio::sync::mpsc;
use tokio_stream::Stream;
use super::{AgentClient, wire};
use crate::error::{EngineError, Result};
pub(super) struct StreamKind<T> {
pub(super) frame: MessageType,
pub(super) end: Option<MessageType>,
pub(super) decode: fn(&[u8]) -> Result<T>,
pub(super) is_last: fn(&T) -> bool,
}
pub struct SandboxStream<T> {
frames: mpsc::UnboundedReceiver<(usize, Result<T>)>,
credit: Arc<Credit>,
unreturned: usize,
out: mpsc::UnboundedSender<Bytes>,
}
impl<T> SandboxStream<T> {
pub async fn recv(&mut self) -> Option<Result<T>> {
std::future::poll_fn(|cx| self.poll_recv(cx)).await
}
fn poll_recv(&mut self, cx: &mut Context<'_>) -> Poll<Option<Result<T>>> {
let (cost, item) = match self.frames.poll_recv(cx) {
Poll::Ready(Some(next)) => next,
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
};
self.consumed(cost);
Poll::Ready(Some(item))
}
fn consumed(&mut self, cost: usize) {
self.unreturned += cost;
if self.unreturned < SANDBOX_STREAM_WINDOW as usize / 2 {
return;
}
let bytes = std::mem::take(&mut self.unreturned);
if self.credit.grant(bytes).is_err() {
return;
}
let grant = SandboxStreamWindow {
bytes: bytes as u32,
..Default::default()
};
let frame =
wire::build_message(MessageType::SandboxStreamWindow, "", &grant.encode_to_vec());
let _ = self.out.send(frame);
}
}
impl<T> Stream for SandboxStream<T> {
type Item = Result<T>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.poll_recv(cx)
}
}
impl AgentClient {
pub(super) async fn open_sandbox_stream<T: Send + 'static>(
mut self,
request: MessageType,
payload: &[u8],
kind: StreamKind<T>,
) -> Result<SandboxStream<T>> {
if !self.connected {
self.connect().await?;
}
let buf = wire::build_message(request, "", payload);
self.transport
.async_send(buf)
.await
.map_err(|source| EngineError::Transport {
context: "failed to send sandbox stream request",
source,
})?;
let (sender, mut receiver) =
self.transport
.into_split()
.map_err(|source| EngineError::Transport {
context: "failed to split sandbox stream transport",
source,
})?;
let credit = Arc::new(Credit::new(SANDBOX_STREAM_WINDOW as usize));
let (out, out_rx) = mpsc::unbounded_channel();
let writer = tokio::spawn(write_frames(sender, out_rx));
let (frames_tx, frames) = mpsc::unbounded_channel();
tokio::spawn({
let credit = Arc::clone(&credit);
async move {
tokio::select! {
() = relay(&mut receiver, &frames_tx, &credit, &kind) => {}
() = frames_tx.closed() => {}
}
writer.abort();
}
});
Ok(SandboxStream {
frames,
credit,
unreturned: 0,
out,
})
}
}
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 relay<T>(
receiver: &mut VsockReceiver,
out: &mpsc::UnboundedSender<(usize, Result<T>)>,
credit: &Credit,
kind: &StreamKind<T>,
) {
loop {
let (cost, item) = match next_frame(receiver, credit, kind).await {
Ok(None) => return,
Ok(Some((cost, item))) => (cost, Ok(item)),
Err(e) => (0, Err(e)),
};
let last = item.as_ref().map_or(true, kind.is_last);
if out.send((cost, item)).is_err() || last {
return;
}
}
}
async fn next_frame<T>(
receiver: &mut VsockReceiver,
credit: &Credit,
kind: &StreamKind<T>,
) -> Result<Option<(usize, T)>> {
let raw = receiver
.recv()
.await
.map_err(|source| EngineError::Transport {
context: "failed to receive sandbox stream frame",
source,
})?;
let (resp_type, _, payload) = wire::parse_response(&raw)?;
let cost = payload.len();
credit.take(cost).map_err(|_| {
EngineError::Machine("guest agent overran the sandbox stream window".into())
})?;
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 kind.end.is_some_and(|end| resp_type == end as u32) {
return Ok(None);
}
AgentClient::expect_response_type(resp_type, kind.frame)?;
Ok(Some((cost, (kind.decode)(&payload)?)))
}
pub(super) fn decode_error(e: impl std::fmt::Display) -> EngineError {
EngineError::Machine(format!("decode error: {e}"))
}