1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
use bytes::{Buf, BufMut, BytesMut};
use prost::Message;
#[cfg(feature = "smol-backend")]
use smol::io::{AsyncReadExt, AsyncWriteExt};
#[cfg(feature = "smol-backend")]
use smol::prelude::{AsyncRead, AsyncWrite};
use tm_protos::abci::{Request, Response};
#[cfg(feature = "tokio-backend")]
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))
}
}