use crate::{
ChecksumEnabled, Error, CHECKSUM_DISABLED, CHECKSUM_ENABLED, PROTOCOL_VERSION, U16_MARKER,
U32_MARKER, U64_MARKER, ZST_MARKER,
};
use bincode::Options;
use futures_io::AsyncWrite;
use futures_util::{Sink, SinkExt};
use serde::de::DeserializeOwned;
use serde::Serialize;
use siphasher::sip::SipHasher;
use std::collections::VecDeque;
use std::hash::Hasher;
use std::mem::size_of;
use std::pin::Pin;
use std::task::{Context, Poll};
#[derive(Debug)]
pub struct AsyncWriteTyped<W: AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> {
raw: Option<W>,
write_buffer: Vec<u8>,
state: AsyncWriteState,
primed_values: VecDeque<T>,
message_features: MessageFeatures,
}
#[derive(Debug)]
pub(crate) enum AsyncWriteState {
WritingVersion { version: [u8; 8], len_sent: usize },
WritingChecksumEnabled,
Idle,
WritingValue { bytes_sent: usize },
Closing,
Closed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct MessageFeatures {
pub size_limit: u64,
pub checksum_enabled: bool,
}
impl<W: AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> Sink<T>
for AsyncWriteTyped<W, 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 raw,
ref mut write_buffer,
ref mut state,
ref mut primed_values,
ref message_features,
} = *self.as_mut();
match futures_core::ready!(Self::maybe_send(
raw.as_mut().expect("infallible"),
state,
write_buffer,
primed_values,
*message_features,
cx,
false,
))? {
Some(()) => {
Pin::new(raw.as_mut().expect("infallible"))
.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 raw,
ref mut state,
ref mut write_buffer,
ref mut primed_values,
ref message_features,
} = *self.as_mut();
match futures_core::ready!(Self::maybe_send(
raw.as_mut().expect("infallible"),
state,
write_buffer,
primed_values,
*message_features,
cx,
true,
))? {
Some(()) => {
Pin::new(raw.as_mut().expect("infallible"))
.poll_close(cx)
.map(|r| r.map_err(Error::Io))
}
None => Poll::Ready(Ok(())),
}
}
}
impl<W: AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> AsyncWriteTyped<W, T> {
pub(crate) fn maybe_send(
raw: &mut W,
state: &mut AsyncWriteState,
write_buffer: &mut Vec<u8>,
primed_values: &mut VecDeque<T>,
message_features: MessageFeatures,
cx: &mut Context<'_>,
closing: bool,
) -> Poll<Result<Option<()>, Error>> {
let MessageFeatures {
checksum_enabled,
size_limit,
} = message_features;
loop {
return match state {
AsyncWriteState::WritingVersion { version, len_sent } => {
while *len_sent < size_of::<u64>() {
let len = futures_core::ready!(
Pin::new(&mut *raw).poll_write(cx, &version[(*len_sent)..])
)?;
*len_sent += len;
}
*state = AsyncWriteState::WritingChecksumEnabled;
continue;
}
AsyncWriteState::WritingChecksumEnabled => {
let to_send = if checksum_enabled {
CHECKSUM_ENABLED
} else {
CHECKSUM_DISABLED
};
if futures_core::ready!(Pin::new(&mut *raw).poll_write(cx, &[to_send]))? == 1 {
*state = AsyncWriteState::Idle;
}
continue;
}
AsyncWriteState::Idle => {
if let Some(item) = primed_values.pop_back() {
write_buffer.clear();
let length = crate::bincode_options(size_limit)
.serialized_size(&item)
.map_err(Error::Bincode)?;
if length > size_limit {
return Poll::Ready(Err(Error::SentMessageTooLarge));
}
if length == 0 {
write_buffer.push(ZST_MARKER);
} else if length < U16_MARKER as u64 {
write_buffer.extend((length as u8).to_le_bytes());
} else if length < 2_u64.pow(16) {
write_buffer.push(U16_MARKER);
write_buffer.extend((length as u16).to_le_bytes());
} else if length < 2_u64.pow(32) {
write_buffer.push(U32_MARKER);
write_buffer.extend((length as u32).to_le_bytes());
} else {
write_buffer.push(U64_MARKER);
write_buffer.extend(length.to_le_bytes());
}
let length_length = write_buffer.len();
crate::bincode_options(size_limit)
.serialize_into(&mut *write_buffer, &item)
.map_err(Error::Bincode)?;
if checksum_enabled {
let mut hasher = SipHasher::new();
hasher.write(&write_buffer[length_length..]);
let checksum = hasher.finish();
write_buffer.extend(checksum.to_le_bytes());
}
*state = AsyncWriteState::WritingValue { bytes_sent: 0 };
continue;
} else if closing {
*state = AsyncWriteState::Closing;
continue;
} else {
Poll::Ready(Ok(Some(())))
}
}
AsyncWriteState::WritingValue { bytes_sent } => {
while *bytes_sent < write_buffer.len() {
let len = futures_core::ready!(
Pin::new(&mut *raw).poll_write(cx, &write_buffer[*bytes_sent..])
)?;
*bytes_sent += len;
}
*state = AsyncWriteState::Idle;
if primed_values.is_empty() {
return Poll::Ready(Ok(Some(())));
}
continue;
}
AsyncWriteState::Closing => {
let len = futures_core::ready!(Pin::new(&mut *raw).poll_write(cx, &[0]))?;
if len == 1 {
*state = AsyncWriteState::Closed;
Poll::Ready(Ok(Some(())))
} else {
continue;
}
}
AsyncWriteState::Closed => Poll::Ready(Ok(None)),
};
}
}
}
impl<W: AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> AsyncWriteTyped<W, T> {
pub fn new_with_limit(raw: W, size_limit: u64, checksum_enabled: ChecksumEnabled) -> Self {
Self {
raw: Some(raw),
write_buffer: Vec::new(),
state: AsyncWriteState::WritingVersion {
version: PROTOCOL_VERSION.to_le_bytes(),
len_sent: 0,
},
message_features: MessageFeatures {
size_limit,
checksum_enabled: checksum_enabled.into(),
},
primed_values: VecDeque::new(),
}
}
pub fn new(raw: W, checksum_enabled: ChecksumEnabled) -> Self {
Self::new_with_limit(raw, 1024u64.pow(2), checksum_enabled)
}
pub fn inner(&self) -> &W {
self.raw.as_ref().expect("infallible")
}
pub fn into_inner(mut self) -> W {
self.raw.take().expect("infallible")
}
pub fn optimize_memory_usage(&mut self) {
match self.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()
}
pub fn checksum_enabled(&self) -> bool {
self.message_features.checksum_enabled
}
}
impl<W: AsyncWrite + Unpin, T: Serialize + DeserializeOwned + Unpin> Drop
for AsyncWriteTyped<W, T>
{
fn drop(&mut self) {
if self.raw.is_some() {
let _ = futures_executor::block_on(SinkExt::close(self));
}
}
}