use std::{future::Future, pin::Pin, task::Poll};
use wasm_bindgen_futures::JsFuture;
#[derive(Debug, Default)]
pub enum Op {
#[default]
Idle,
ReadPending(JsFuture),
ConsumingReadBuffer {
read_buffer: js_sys::Uint8Array,
already_read: usize,
},
}
#[derive(Debug)]
pub enum Mode {
Byob {
reader: web_sys::ReadableStreamByobReader,
internal_buf: Option<js_sys::ArrayBuffer>,
},
Default {
reader: web_sys::ReadableStreamDefaultReader,
},
}
impl Mode {
pub fn release_lock(&self) {
match self {
Self::Byob { reader, .. } => reader.release_lock(),
Self::Default { reader } => reader.release_lock(),
}
}
pub fn cancel_with_reason(&self, reason: &wasm_bindgen::JsValue) -> js_sys::Promise {
match self {
Self::Byob { reader, .. } => reader.cancel_with_reason(reason),
Self::Default { reader } => reader.cancel_with_reason(reason),
}
}
pub fn closed(&self) -> js_sys::Promise {
match self {
Self::Byob { reader, .. } => reader.closed(),
Self::Default { reader } => reader.closed(),
}
}
fn start_read(&mut self, requested_size: u32) -> JsFuture {
match self {
Self::Byob {
reader,
internal_buf,
} => {
let internal_buf = internal_buf
.take()
.filter(|internal_buf| {
let actual_size = internal_buf.byte_length();
debug_assert!(actual_size > 0);
actual_size >= requested_size
})
.unwrap_or_else(|| js_sys::ArrayBuffer::new(requested_size));
let internal_buf_view = js_sys::Uint8Array::new_with_byte_offset_and_length(
&internal_buf,
0,
requested_size,
);
JsFuture::from(reader.read_with_array_buffer_view(&internal_buf_view))
}
Self::Default { reader } => JsFuture::from(reader.read()),
}
}
fn reclaim_read_buffer(&mut self, read_buffer: &js_sys::Uint8Array) {
if let Self::Byob { internal_buf, .. } = self {
*internal_buf = Some(read_buffer.buffer());
}
}
}
fn parse_read_result(
mode: &Mode,
result: Result<wasm_bindgen::JsValue, wasm_bindgen::JsValue>,
) -> Result<Option<js_sys::Uint8Array>, ReadError> {
let result = result.map_err(ReadError::Read)?;
let result: crate::sys::ReadableStreamReaderValue = result.into();
match result.value() {
Some(read_buffer) => Ok(Some(read_buffer)),
None => match mode {
Mode::Byob { .. } => Err(ReadError::ByobReadConsumedBuffer),
Mode::Default { .. } => Ok(None),
},
}
}
fn consume_read_buffer(
mode: &mut Mode,
op: &mut Op,
read_buffer: js_sys::Uint8Array,
already_read: usize,
dest: &mut [u8],
) -> usize {
let read_buffer_size = read_buffer.byte_length() as usize;
let remaining_size = read_buffer_size - already_read;
let copy_size = remaining_size.min(dest.len());
if already_read == 0 && copy_size == read_buffer_size {
read_buffer.copy_to(&mut dest[..copy_size]);
} else {
let source_view =
read_buffer.subarray(already_read as u32, (already_read + copy_size) as u32);
source_view.copy_to(&mut dest[..copy_size]);
}
if already_read + copy_size < read_buffer_size {
*op = Op::ConsumingReadBuffer {
read_buffer,
already_read: already_read + copy_size,
};
} else {
mode.reclaim_read_buffer(&read_buffer);
}
copy_size
}
impl From<web_sys::ReadableStreamByobReader> for Mode {
fn from(reader: web_sys::ReadableStreamByobReader) -> Self {
Self::Byob {
reader,
internal_buf: None,
}
}
}
impl From<web_sys::ReadableStreamDefaultReader> for Mode {
fn from(reader: web_sys::ReadableStreamDefaultReader) -> Self {
Self::Default { reader }
}
}
#[derive(Debug)]
pub enum ReadError {
Read(wasm_bindgen::JsValue),
ByobReadConsumedBuffer,
}
impl From<ReadError> for std::io::Error {
fn from(err: ReadError) -> Self {
match err {
ReadError::Read(err) => super::js_value_to_io_error(err),
ReadError::ByobReadConsumedBuffer => {
std::io::Error::other("BYOB read consumed the buffer and did not provide a new one")
}
}
}
}
#[derive(Debug)]
pub struct Reader {
pub inner: Mode,
pub op: Op,
}
impl Reader {
pub fn new(inner: impl Into<Mode>) -> Self {
Self {
inner: inner.into(),
op: Op::default(),
}
}
pub fn with_buf(
inner: web_sys::ReadableStreamByobReader,
internal_buf: js_sys::ArrayBuffer,
) -> Self {
Self {
inner: Mode::Byob {
reader: inner,
internal_buf: Some(internal_buf),
},
op: Op::default(),
}
}
pub async fn read_into(&mut self, buf: &mut [u8]) -> Result<usize, ReadError> {
let (read_buffer, already_read) = match std::mem::take(&mut self.op) {
Op::ConsumingReadBuffer {
read_buffer,
already_read,
} => (read_buffer, already_read),
op => {
let fut = match op {
Op::ReadPending(fut) => fut,
_ => {
let requested_size = buf.len().try_into().unwrap();
self.inner.start_read(requested_size)
}
};
match parse_read_result(&self.inner, fut.await)? {
Some(read_buffer) => (read_buffer, 0),
None => return Ok(0),
}
}
};
Ok(consume_read_buffer(
&mut self.inner,
&mut self.op,
read_buffer,
already_read,
buf,
))
}
}
impl tokio::io::AsyncRead for Reader {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let this = self.get_mut();
if buf.remaining() == 0 {
return Poll::Ready(Ok(()));
}
match std::mem::take(&mut this.op) {
Op::ReadPending(mut fut) => {
let result = match Pin::new(&mut fut).poll(cx) {
Poll::Pending => {
this.op = Op::ReadPending(fut);
return Poll::Pending;
}
Poll::Ready(result) => result,
};
let Some(read_buffer) = parse_read_result(&this.inner, result)? else {
return Poll::Ready(Ok(()));
};
this.op = Op::ConsumingReadBuffer {
read_buffer,
already_read: 0,
};
Pin::new(this).poll_read(cx, buf)
}
Op::ConsumingReadBuffer {
read_buffer,
already_read,
} => {
let remaining_size = read_buffer.byte_length() as usize - already_read;
let copy_size = remaining_size.min(buf.remaining());
let write_slice = buf.initialize_unfilled_to(copy_size);
let copied = consume_read_buffer(
&mut this.inner,
&mut this.op,
read_buffer,
already_read,
write_slice,
);
buf.advance(copied);
Poll::Ready(Ok(()))
}
Op::Idle => {
let requested_size = buf.remaining().try_into().unwrap();
let fut = this.inner.start_read(requested_size);
this.op = Op::ReadPending(fut);
Pin::new(this).poll_read(cx, buf)
}
}
}
}