use crate::poll::Pollable;
use alloc::boxed::Box;
use bytes::Bytes;
use wasmtime::Result;
const MAX_BLOCKING_ATTEMPTS: u8 = 10;
#[async_trait::async_trait]
pub trait InputStream: Pollable {
fn read(&mut self, size: usize) -> StreamResult<Bytes>;
async fn blocking_read(&mut self, size: usize) -> StreamResult<Bytes> {
if size == 0 {
self.ready().await;
return self.read(size);
}
let mut i = 0;
loop {
self.ready().await;
let data = self.read(size)?;
if !data.is_empty() {
return Ok(data);
}
if i >= MAX_BLOCKING_ATTEMPTS {
return Err(StreamError::trap("max blocking attempts exceeded"));
}
i += 1;
}
}
fn skip(&mut self, nelem: usize) -> StreamResult<usize> {
let bs = self.read(nelem)?;
Ok(bs.len())
}
async fn blocking_skip(&mut self, nelem: usize) -> StreamResult<usize> {
let bs = self.blocking_read(nelem).await?;
Ok(bs.len())
}
async fn cancel(&mut self) {}
}
pub type Error = wasmtime::Error;
pub type StreamResult<T> = Result<T, StreamError>;
#[derive(Debug)]
pub enum StreamError {
Closed,
LastOperationFailed(wasmtime::Error),
Trap(wasmtime::Error),
}
impl StreamError {
pub fn trap(msg: &str) -> StreamError {
StreamError::Trap(wasmtime::format_err!("{msg}"))
}
}
impl alloc::fmt::Display for StreamError {
fn fmt(&self, f: &mut alloc::fmt::Formatter<'_>) -> alloc::fmt::Result {
match self {
StreamError::Closed => write!(f, "closed"),
StreamError::LastOperationFailed(e) => write!(f, "last operation failed: {e}"),
StreamError::Trap(e) => write!(f, "trap: {e}"),
}
}
}
impl core::error::Error for StreamError {
fn source(&self) -> Option<&(dyn core::error::Error + 'static)> {
match self {
StreamError::Closed => None,
StreamError::LastOperationFailed(e) | StreamError::Trap(e) => e.source(),
}
}
}
impl From<wasmtime::component::ResourceTableError> for StreamError {
fn from(error: wasmtime::component::ResourceTableError) -> Self {
Self::Trap(error.into())
}
}
#[async_trait::async_trait]
pub trait OutputStream: Pollable {
fn write(&mut self, bytes: Bytes) -> StreamResult<()>;
fn flush(&mut self) -> StreamResult<()>;
fn check_write(&mut self) -> StreamResult<usize>;
async fn blocking_write_and_flush(&mut self, mut bytes: Bytes) -> StreamResult<()> {
loop {
let permit = self.write_ready().await?;
let len = bytes.len().min(permit);
let chunk = bytes.split_to(len);
self.write(chunk)?;
if bytes.is_empty() {
break;
}
}
match self.flush() {
Ok(_) => {}
Err(StreamError::Closed) => {}
Err(e) => Err(e)?,
};
match self.write_ready().await {
Ok(_) => {}
Err(StreamError::Closed) => {}
Err(e) => Err(e)?,
};
Ok(())
}
fn write_zeroes(&mut self, nelem: usize) -> StreamResult<()> {
let n = self.check_write()?;
if nelem > n {
return Err(StreamError::trap(
"cannot write more zeroes than `check_write` allows",
));
};
let bs = Bytes::from_iter(core::iter::repeat(0).take(nelem));
self.write(bs)?;
Ok(())
}
async fn write_ready(&mut self) -> StreamResult<usize> {
let mut i = 0;
loop {
self.ready().await;
let n = self.check_write()?;
if n > 0 {
return Ok(n);
}
if i >= MAX_BLOCKING_ATTEMPTS {
return Err(StreamError::trap("max blocking attempts exceeded"));
}
i += 1;
}
}
async fn cancel(&mut self) {}
}
#[async_trait::async_trait]
impl Pollable for Box<dyn OutputStream> {
async fn ready(&mut self) {
(**self).ready().await
}
}
#[async_trait::async_trait]
impl Pollable for Box<dyn InputStream> {
async fn ready(&mut self) {
(**self).ready().await
}
}
pub type DynInputStream = Box<dyn InputStream>;
pub type DynOutputStream = Box<dyn OutputStream>;