use bytes::{Buf, BufMut, BytesMut};
use prost::Message;
use tm_protos::abci::{Request, Response};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use crate::error::Error;
pub const MAX_VARINT_LENGTH: usize = 16;
pub struct ICodec<R> {
stream: R,
read_buf: BytesMut,
read_window: Vec<u8>,
}
impl<R> ICodec<R> {
pub fn new(stream: R, read_buf_size: usize) -> Self {
Self {
stream,
read_buf: BytesMut::new(),
read_window: vec![0_u8; read_buf_size],
}
}
}
impl<R> ICodec<R>
where
R: AsyncRead + Unpin,
{
pub async fn next(&mut self) -> Option<Result<Request, Error>> {
loop {
match decode_length_delimited::<Request>(&mut self.read_buf) {
Ok(Some(incoming)) => return Some(Ok(incoming)),
Err(e) => return Some(Err(e)),
_ => (), }
let bytes_read = match self.stream.read(self.read_window.as_mut()).await {
Ok(br) => br,
Err(e) => return Some(Err(Error::StdIoError(e))),
};
if bytes_read == 0 {
return None;
}
self.read_buf
.extend_from_slice(&self.read_window[..bytes_read]);
}
}
}
pub struct OCodec<W> {
stream: W,
write_buf: BytesMut,
}
impl<W> OCodec<W> {
pub fn new(stream: W) -> Self {
Self {
stream,
write_buf: BytesMut::default(),
}
}
}
impl<W> OCodec<W>
where
W: AsyncWrite + Unpin,
{
pub async fn send(&mut self, message: Response) -> Result<(), Error> {
encode_length_delimited(message, &mut self.write_buf)?;
while !self.write_buf.is_empty() {
let bytes_written = self
.stream
.write(self.write_buf.as_ref())
.await
.map_err(Error::StdIoError)?;
if bytes_written == 0 {
return Err(Error::StdIoError(std::io::Error::new(
std::io::ErrorKind::WriteZero,
"failed to write to underlying stream",
)));
}
self.write_buf.advance(bytes_written);
}
self.stream.flush().await.map_err(Error::StdIoError)?;
Ok(())
}
}
pub fn encode_length_delimited<M, B>(message: M, mut dst: &mut B) -> Result<(), Error>
where
M: Message,
B: BufMut,
{
let mut buf = BytesMut::new();
message.encode(&mut buf).map_err(Error::ProstEncodeError)?;
let buf = buf.freeze();
prost::encoding::encode_varint(buf.len() as u64, &mut dst);
dst.put(buf);
Ok(())
}
pub fn decode_length_delimited<M>(src: &mut BytesMut) -> Result<Option<M>, Error>
where
M: Message + Default,
{
let src_len = src.len();
let mut tmp = src.clone().freeze();
let encoded_len = match prost::encoding::decode_varint(&mut tmp) {
Ok(len) => len,
Err(_) if src_len <= MAX_VARINT_LENGTH => return Ok(None),
Err(e) => return Err(Error::ProstDecodeError(e)),
};
let remaining = tmp.remaining() as u64;
if remaining < encoded_len {
Ok(None)
} else {
let delim_len = src_len - tmp.remaining();
src.advance(delim_len + (encoded_len as usize));
let mut result_bytes = BytesMut::from(tmp.split_to(encoded_len as usize).as_ref());
let res = M::decode(&mut result_bytes).map_err(Error::ProstDecodeError)?;
Ok(Some(res))
}
}