volo_grpc/codec/
encode.rs1use bytes::{BufMut, Bytes};
2use futures::{Stream, StreamExt};
3use http_body::Frame;
4use linkedbytes::Node;
5use pilota::{LinkedBytes, pb::Message};
6
7use super::{DefaultEncoder, PREFIX_LEN};
8use crate::{
9 BoxStream, Status,
10 codec::{
11 BUFFER_SIZE, Encoder,
12 compression::{CompressionEncoding, compress},
13 },
14};
15
16pub fn encode<T, S>(
17 source: S,
18 compression_encoding: Option<CompressionEncoding>,
19) -> BoxStream<'static, Result<Frame<Bytes>, Status>>
20where
21 S: Stream<Item = Result<T, Status>> + Send + 'static,
22 T: Message + 'static,
23{
24 Box::pin(async_stream::stream! {
25 futures_util::pin_mut!(source);
26
27 loop {
28 match source.next().await {
29 Some(Ok(item)) => {
30 let mut buf = LinkedBytes::with_capacity(BUFFER_SIZE);
31 let mut compressed_buf = if compression_encoding.is_some() {
32 LinkedBytes::with_capacity(BUFFER_SIZE)
33 } else {
34 LinkedBytes::new()
35 };
36
37 buf.reserve(PREFIX_LEN);
38 unsafe {
39 buf.advance_mut(PREFIX_LEN);
40 }
41
42 let mut encoder=DefaultEncoder::default();
43
44 if let Some(config)=compression_encoding{
45 encoder.encode(item, &mut compressed_buf)
46 .map_err(|err| Status::internal(format!("Error encoding: {err}")))?;
47 compress(config,&mut compressed_buf.concat(), buf.bytes_mut())
48 .map_err(|err| Status::internal(format!("Error compressing: {err}")))?;
49 } else {
50 encoder.encode(item, &mut buf)
51 .map_err(|err| Status::internal(format!("Error encoding: {err}")))?;
52 }
53
54 let len = buf.len() - PREFIX_LEN;
55 assert!(len <= u32::MAX as usize);
56 {
57 if let Some(node) = buf.get_list_mut(0) {
58 match node {
59 linkedbytes::Node::BytesMut(bytes_mut) => {
60 let mut dest = &mut bytes_mut[..PREFIX_LEN];
61 dest.put_u8(compression_encoding.is_some() as u8);
62 dest.put_u32(len as u32);
63 }
64 _ => unreachable!("reserve_node_idx is not a bytesmut"),
65 };
66 } else {
67 let mut dest = &mut buf.bytes_mut()[..PREFIX_LEN];
68 dest.put_u8(compression_encoding.is_some() as u8);
69 dest.put_u32(len as u32);
70 }
71 }
72
73 for node in buf.into_iter_list() {
75 let bytes = match node {
76 Node::Bytes(bytes) => bytes,
77 Node::BytesMut(bytesmut) => bytesmut.freeze(),
78 Node::FastStr(faststr) => faststr.into_bytes(),
79 };
80 if !bytes.is_empty() {
81 yield Ok(Frame::data(bytes));
82 }
83 }
84 },
85 Some(Err(status)) => yield Err(status),
86 None => break,
87 }
88 }
89 })
90}
91
92pub mod tests {
93
94 #[derive(Debug, Default, Clone, PartialEq)]
95 pub struct EchoRequest {
96 pub message: ::pilota::FastStr,
97 }
98 impl pilota::pb::Message for EchoRequest {
99 #[inline]
100 fn encoded_len(&self, ctx: &mut pilota::pb::EncodeLengthContext) -> usize {
101 pilota::pb::encoding::faststr::encoded_len(ctx, 1, &self.message)
102 }
103
104 #[allow(unused_variables)]
105 fn encode_raw(&self, buf: &mut pilota::LinkedBytes) {
106 pilota::pb::encoding::faststr::encode(1, &self.message, buf);
107 }
108
109 #[allow(unused_variables)]
110 fn merge_field(
111 &mut self,
112 tag: u32,
113 wire_type: pilota::pb::encoding::WireType,
114 buf: &mut pilota::Bytes,
115 ctx: &mut pilota::pb::encoding::DecodeContext,
116 _is_root: bool,
117 ) -> core::result::Result<(), pilota::pb::DecodeError> {
118 const STRUCT_NAME: &str = stringify!(EchoRequest);
119
120 match tag {
121 1 => {
122 let mut _inner_pilota_value = &mut self.message;
123 pilota::pb::encoding::faststr::merge(wire_type, _inner_pilota_value, buf, ctx)
124 .map_err(|mut error| {
125 error.push(STRUCT_NAME, stringify!(message));
126 error
127 })
128 }
129 _ => pilota::pb::encoding::skip_field(wire_type, tag, buf, ctx),
130 }
131 }
132 }
133
134 #[tokio::test]
135 async fn test_encode() {
136 use super::*;
137 let source = async_stream::stream! {
138 yield Ok(EchoRequest { message: "Volo".into() });
139 };
140
141 let mut stream = encode(source, None);
142 let frame = stream.next().await.unwrap().unwrap();
144 assert!(frame.is_data());
145 let data = frame.data_ref().unwrap();
146 assert_eq!(&data[..PREFIX_LEN], b"\x00\x00\x00\x00\x06");
147 assert_eq!(&data[PREFIX_LEN..], b"\x0a\x04Volo");
148
149 assert!(stream.next().await.is_none());
150 }
151
152 #[cfg(feature = "gzip")]
153 #[tokio::test]
154 async fn test_encode_gzip() {
155 use bytes::BytesMut;
156
157 use super::*;
158 use crate::codec::compression::{GzipConfig, decompress};
159
160 let source = async_stream::stream! {
161 yield Ok(EchoRequest { message: "Volo".into() });
162 };
163
164 let compression_encoding = Some(CompressionEncoding::Gzip(Some(GzipConfig::default())));
165 let mut stream = encode(source, compression_encoding);
166
167 let frame = stream.next().await.unwrap().unwrap();
169 assert!(frame.is_data());
170 let data = frame.data_ref().unwrap();
171 assert_eq!(&data[..PREFIX_LEN], b"\x01\x00\x00\x00\x1a");
172
173 let mut compressed_data = BytesMut::from(&data[PREFIX_LEN..]);
174 let mut uncompressed_data_mut = BytesMut::new();
175 decompress(
176 compression_encoding.unwrap(),
177 &mut compressed_data,
178 &mut uncompressed_data_mut,
179 )
180 .unwrap();
181 assert_eq!(&uncompressed_data_mut[..], b"\x0a\x04Volo");
182
183 assert!(stream.next().await.is_none());
184 }
185
186 #[cfg(feature = "zlib")]
187 #[tokio::test]
188 async fn test_encode_zlib() {
189 use bytes::BytesMut;
190
191 use super::*;
192 use crate::codec::compression::{ZlibConfig, decompress};
193
194 let source = async_stream::stream! {
195 yield Ok(EchoRequest { message: "Volo".into() });
196 };
197
198 let compression_encoding = Some(CompressionEncoding::Zlib(Some(ZlibConfig::default())));
199 let mut stream = encode(source, compression_encoding);
200
201 let frame = stream.next().await.unwrap().unwrap();
203 assert!(frame.is_data());
204 let data = frame.data_ref().unwrap();
205 assert_eq!(&data[..PREFIX_LEN], b"\x01\x00\x00\x00\x0e");
206
207 let mut compressed_data = BytesMut::from(&data[PREFIX_LEN..]);
208 let mut uncompressed_data_mut = BytesMut::new();
209 decompress(
210 compression_encoding.unwrap(),
211 &mut compressed_data,
212 &mut uncompressed_data_mut,
213 )
214 .unwrap();
215 assert_eq!(&uncompressed_data_mut[..], b"\x0a\x04Volo");
216
217 assert!(stream.next().await.is_none());
218 }
219
220 #[cfg(feature = "zstd")]
221 #[tokio::test]
222 async fn test_encode_zstd() {
223 use bytes::BytesMut;
224
225 use super::*;
226 use crate::codec::compression::{ZstdConfig, decompress};
227
228 let source = async_stream::stream! {
229 yield Ok(EchoRequest { message: "Volo".into() });
230 };
231
232 let compression_encoding = Some(CompressionEncoding::Zstd(Some(ZstdConfig::default())));
233 let mut stream = encode(source, compression_encoding);
234
235 let frame = stream.next().await.unwrap().unwrap();
237 assert!(frame.is_data());
238 let data = frame.data_ref().unwrap();
239 assert_eq!(&data[..PREFIX_LEN], b"\x01\x00\x00\x00\x0f");
240
241 let mut compressed_data = BytesMut::from(&data[PREFIX_LEN..]);
242 let mut uncompressed_data_mut = BytesMut::new();
243 decompress(
244 compression_encoding.unwrap(),
245 &mut compressed_data,
246 &mut uncompressed_data_mut,
247 )
248 .unwrap();
249 assert_eq!(&uncompressed_data_mut[..], b"\x0a\x04Volo");
250
251 assert!(stream.next().await.is_none());
252 }
253}