use std::future::Future;
use std::io;
use std::mem;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncSeek, ReadBuf};
use crate::error::Error;
use crate::io::file_in_stream::GoosefsFileInStream;
type OwnedReadFut =
Pin<Box<dyn Future<Output = (Box<GoosefsFileInStream>, io::Result<Vec<u8>>)> + Send>>;
type OwnedSeekFut =
Pin<Box<dyn Future<Output = (Box<GoosefsFileInStream>, io::Result<i64>)> + Send>>;
enum State {
Idle(Box<GoosefsFileInStream>),
Reading(OwnedReadFut),
Seeking(OwnedSeekFut),
Empty,
}
pub struct GoosefsAsyncReader {
state: State,
}
impl GoosefsAsyncReader {
pub fn new(stream: GoosefsFileInStream) -> Self {
Self {
state: State::Idle(Box::new(stream)),
}
}
#[allow(clippy::result_large_err)] pub fn into_inner(self) -> Result<GoosefsFileInStream, Self> {
match self.state {
State::Idle(s) => Ok(*s),
_ => Err(self),
}
}
pub fn get_ref(&self) -> Option<&GoosefsFileInStream> {
if let State::Idle(s) = &self.state {
Some(s)
} else {
None
}
}
pub fn get_mut(&mut self) -> Option<&mut GoosefsFileInStream> {
if let State::Idle(s) = &mut self.state {
Some(s)
} else {
None
}
}
}
fn take_idle(state: &mut State) -> Box<GoosefsFileInStream> {
match mem::replace(state, State::Empty) {
State::Idle(s) => s,
_ => unreachable!("take_idle called on non-Idle state"),
}
}
fn sdk_err_to_io(err: Error) -> io::Error {
io::Error::other(err.to_string())
}
impl AsyncRead for GoosefsAsyncReader {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
loop {
match &mut this.state {
State::Idle(_) => {
if buf.remaining() == 0 {
return Poll::Ready(Ok(()));
}
let cap = buf.remaining();
let mut stream = take_idle(&mut this.state);
let fut: OwnedReadFut = Box::pin(async move {
let mut tmp = vec![0u8; cap];
let result = stream.read(&mut tmp).await.map_err(sdk_err_to_io);
let bytes = match result {
Ok(n) => {
tmp.truncate(n);
Ok(tmp)
}
Err(e) => Err(e),
};
(stream, bytes)
});
this.state = State::Reading(fut);
}
State::Reading(fut) => match fut.as_mut().poll(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready((stream, result)) => {
this.state = State::Idle(stream);
return match result {
Ok(data) => {
let take = data.len().min(buf.remaining());
buf.put_slice(&data[..take]);
Poll::Ready(Ok(()))
}
Err(e) => Poll::Ready(Err(e)),
};
}
},
State::Seeking(_) => {
return Poll::Ready(Err(io::Error::other(
"cannot read while a seek is in flight",
)));
}
State::Empty => unreachable!("Empty state observed in poll_read"),
}
}
}
}
impl AsyncSeek for GoosefsAsyncReader {
fn start_seek(self: Pin<&mut Self>, position: io::SeekFrom) -> io::Result<()> {
let this = self.get_mut();
match &this.state {
State::Reading(_) => Err(io::Error::other(
"cannot start a seek while a read is in flight",
)),
State::Seeking(_) => Err(io::Error::other("another seek is already in flight")),
State::Empty => unreachable!("Empty state observed in start_seek"),
State::Idle(_) => {
let stream = take_idle(&mut this.state);
let fut: OwnedSeekFut = Box::pin(async move {
let (s, result) = (*stream).seek_owned(position).await;
(Box::new(s), result)
});
this.state = State::Seeking(fut);
Ok(())
}
}
}
fn poll_complete(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<u64>> {
let this = self.get_mut();
match &mut this.state {
State::Idle(s) => {
let pos = s.pos().max(0) as u64;
Poll::Ready(Ok(pos))
}
State::Seeking(fut) => match fut.as_mut().poll(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready((stream, result)) => {
this.state = State::Idle(stream);
match result {
Ok(pos) => Poll::Ready(Ok(pos.max(0) as u64)),
Err(e) => Poll::Ready(Err(e)),
}
}
},
State::Reading(_) => Poll::Ready(Err(io::Error::other(
"cannot complete a seek while a read is in flight",
))),
State::Empty => unreachable!("Empty state observed in poll_complete"),
}
}
}
impl GoosefsFileInStream {
pub(crate) async fn seek_owned(mut self, from: io::SeekFrom) -> (Self, io::Result<i64>) {
let result = self.seek_from(from).await.map_err(sdk_err_to_io);
(self, result)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sdk_err_to_io_uses_other_kind() {
let sdk = Error::Internal {
message: "boom".to_string(),
source: None,
};
let mapped = sdk_err_to_io(sdk);
assert_eq!(mapped.kind(), io::ErrorKind::Other);
assert!(mapped.to_string().contains("boom"));
}
#[test]
fn test_state_empty_replacement_round_trip() {
let mut s = State::Empty;
let replaced = mem::replace(&mut s, State::Empty);
assert!(matches!(replaced, State::Empty));
assert!(matches!(s, State::Empty));
}
#[test]
fn test_send_and_unpin() {
fn assert_send<T: Send>() {}
fn assert_unpin<T: Unpin>() {}
assert_send::<GoosefsAsyncReader>();
assert_unpin::<GoosefsAsyncReader>();
}
}