use crate::read::{AsyncReadState, AsyncReadTyped, ChecksumReadState};
use crate::write::{AsyncWriteState, AsyncWriteTyped, MessageFeatures};
use crate::{ChecksumEnabled, Error, PROTOCOL_VERSION};
use futures_core::Stream;
use futures_io::{AsyncRead, AsyncWrite};
use futures_util::{Sink, SinkExt};
use serde::de::DeserializeOwned;
use serde::Serialize;
use std::collections::VecDeque;
use std::pin::Pin;
use std::task::{Context, Poll};
#[derive(Debug)]
pub struct DuplexStreamTyped<
RW: AsyncRead + AsyncWrite + Unpin,
T: Serialize + DeserializeOwned + Unpin,
> {
rw: Option<RW>,
read_state: AsyncReadState,
read_buffer: Vec<u8>,
write_state: AsyncWriteState,
write_buffer: Vec<u8>,
primed_values: VecDeque<T>,
checksum_read_state: ChecksumReadState,
message_features: MessageFeatures,
}
impl<RW: AsyncRead + AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin>
DuplexStreamTyped<RW, T>
{
pub fn new_with_limit(rw: RW, size_limit: u64, checksum_enabled: ChecksumEnabled) -> Self {
Self {
rw: Some(rw),
read_state: AsyncReadState::ReadingVersion {
version_in_progress: [0; 8],
version_in_progress_assigned: 0,
},
read_buffer: Vec::new(),
write_state: AsyncWriteState::WritingVersion {
version: PROTOCOL_VERSION.to_le_bytes(),
len_sent: 0,
},
write_buffer: Vec::new(),
primed_values: VecDeque::new(),
checksum_read_state: checksum_enabled.into(),
message_features: MessageFeatures {
size_limit,
checksum_enabled: checksum_enabled.into(),
},
}
}
pub fn new(rw: RW, checksum_enabled: ChecksumEnabled) -> Self {
Self::new_with_limit(rw, 1024_u64.pow(2), checksum_enabled)
}
pub fn inner(&self) -> &RW {
self.rw.as_ref().expect("infallible")
}
pub fn into_inner(mut self) -> RW {
self.rw.take().expect("infallible")
}
pub fn optimize_memory_usage(&mut self) {
match self.read_state {
AsyncReadState::ReadingItem { .. } => self.read_buffer.shrink_to_fit(),
_ => {
self.read_buffer = Vec::new();
}
}
match self.write_state {
AsyncWriteState::WritingValue { .. } => self.write_buffer.shrink_to_fit(),
_ => {
self.write_buffer = Vec::new();
}
}
}
pub fn current_memory_usage(&self) -> usize {
self.write_buffer.capacity() + self.read_buffer.capacity()
}
pub fn checksum_send_enabled(&self) -> bool {
self.message_features.checksum_enabled
}
pub fn checksum_receive_enabled(&self) -> bool {
self.checksum_read_state == ChecksumReadState::Yes
}
}
impl<RW: AsyncRead + AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> Stream
for DuplexStreamTyped<RW, T>
{
type Item = Result<T, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let Self {
ref mut rw,
ref mut read_state,
ref mut read_buffer,
ref message_features,
ref mut checksum_read_state,
..
} = *self.as_mut();
AsyncReadTyped::poll_next_impl(
read_state,
rw.as_mut().expect("infallible"),
read_buffer,
message_features.size_limit,
checksum_read_state,
cx,
)
}
}
impl<RW: AsyncRead + AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> Sink<T>
for DuplexStreamTyped<RW, T>
{
type Error = Error;
fn poll_ready(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(mut self: Pin<&mut Self>, item: T) -> Result<(), Self::Error> {
self.primed_values.push_front(item);
Ok(())
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
let Self {
ref mut rw,
ref mut write_state,
ref mut write_buffer,
ref mut primed_values,
ref message_features,
..
} = *self.as_mut();
let rw = rw.as_mut().expect("infallible");
match futures_core::ready!(AsyncWriteTyped::maybe_send(
rw,
write_state,
write_buffer,
primed_values,
*message_features,
cx,
false,
))? {
Some(()) => {
Pin::new(rw).poll_flush(cx).map(|r| r.map_err(Error::Io))
}
None => Poll::Ready(Ok(())),
}
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
let Self {
ref mut rw,
ref mut write_state,
ref mut write_buffer,
ref mut primed_values,
ref message_features,
..
} = *self.as_mut();
let rw = rw.as_mut().expect("infallible");
match futures_core::ready!(AsyncWriteTyped::maybe_send(
rw,
write_state,
write_buffer,
primed_values,
*message_features,
cx,
true,
))? {
Some(()) => {
Pin::new(rw).poll_close(cx).map(|r| r.map_err(Error::Io))
}
None => Poll::Ready(Ok(())),
}
}
}
impl<RW: AsyncRead + AsyncWrite + Unpin, T: Serialize + Unpin + DeserializeOwned> Drop
for DuplexStreamTyped<RW, T>
{
fn drop(&mut self) {
if self.rw.is_some() {
let _ = futures_executor::block_on(SinkExt::close(self));
}
}
}