Skip to main content

volo_grpc/codec/
encode.rs

1use 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                    // send each node in linked bytes as a separate frame
74                    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        // frame
143        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        // frame
168        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        // frame
202        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        // frame
236        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}