flare_core/common/
protobuf_decoder.rs1use bytes::BytesMut;
10use prost::Message;
11use std::io::{self, Cursor};
12
13pub struct ProtobufDecoder<T> {
15 _phantom: std::marker::PhantomData<T>,
17 buffer: BytesMut,
19}
20
21impl<T> ProtobufDecoder<T>
22where
23 T: Message + Default,
24{
25 pub fn new() -> Self {
27 Self {
28 _phantom: std::marker::PhantomData,
29 buffer: BytesMut::new(),
30 }
31 }
32
33 pub fn add_data(&mut self, data: &[u8]) {
35 self.buffer.extend_from_slice(data);
36 }
37
38 pub fn decode_next(&mut self) -> Result<Option<T>, Box<dyn std::error::Error + Send + Sync>> {
44 if let Some((message_len, offset)) = self.read_varint()? {
46 if self.buffer.len() < offset + message_len {
48 return Ok(None);
50 }
51
52 let message_data = self.buffer.split_to(offset + message_len).freeze();
54 let message_bytes = &message_data[offset..];
55
56 let message = T::decode(message_bytes)?;
58 Ok(Some(message))
59 } else {
60 Ok(None)
62 }
63 }
64
65 fn read_varint(&self) -> Result<Option<(usize, usize)>, io::Error> {
68 let mut cursor = Cursor::new(&self.buffer);
69 let mut result: u64 = 0;
70 let mut shift = 0;
71
72 for _i in 0..10 {
73 if cursor.position() as usize >= self.buffer.len() {
75 return Ok(None);
77 }
78
79 let byte = cursor.get_ref()[cursor.position() as usize];
80 cursor.set_position(cursor.position() + 1);
81
82 let value = (byte & 0x7F) as u64;
84 result |= value << shift;
85
86 if (byte & 0x80) == 0 {
88 return Ok(Some((result as usize, (cursor.position()) as usize)));
90 }
91
92 shift += 7;
93 }
94
95 Err(io::Error::new(
96 io::ErrorKind::InvalidData,
97 "Varint too long",
98 ))
99 }
100
101 pub fn has_remaining(&self) -> bool {
103 !self.buffer.is_empty()
104 }
105
106 pub fn remaining_len(&self) -> usize {
108 self.buffer.len()
109 }
110
111 pub fn clear(&mut self) {
113 self.buffer.clear();
114 }
115}
116
117impl<T> Default for ProtobufDecoder<T>
118where
119 T: Message + Default,
120{
121 fn default() -> Self {
122 Self::new()
123 }
124}
125
126pub fn safe_protobuf_decode<T>(data: &[u8]) -> Result<T, Box<dyn std::error::Error + Send + Sync>>
130where
131 T: Message + Default,
132{
133 match T::decode(data) {
134 Ok(message) => Ok(message),
135 Err(e) => Err(Box::new(e)),
136 }
137}
138
139pub fn safe_string_decode(data: &[u8]) -> Result<String, Box<dyn std::error::Error + Send + Sync>> {
143 match std::str::from_utf8(data) {
144 Ok(s) => Ok(s.to_string()),
145 Err(e) => Err(Box::new(e)),
146 }
147}
148
149#[cfg(test)]
150mod tests {
151 use super::*;
152 use prost::Message;
153
154 #[derive(Clone, PartialEq, Message)]
156 struct TestMessage {
157 #[prost(string, tag = "1")]
158 content: String,
159 #[prost(int32, tag = "2")]
160 id: i32,
161 }
162
163 #[test]
164 fn test_protobuf_decoder() {
165 let mut decoder = ProtobufDecoder::<TestMessage>::new();
166
167 let test_msg = TestMessage {
169 content: "Hello, World!".to_string(),
170 id: 42,
171 };
172
173 let encoded_msg = test_msg.encode_to_vec();
175 let mut prefixed_data = Vec::new();
176 prost::encoding::encode_varint(encoded_msg.len() as u64, &mut prefixed_data);
177 prefixed_data.extend_from_slice(&encoded_msg);
178
179 decoder.add_data(&prefixed_data);
181
182 let decoded_msg = decoder.decode_next().unwrap().unwrap();
184
185 assert_eq!(decoded_msg.content, "Hello, World!");
186 assert_eq!(decoded_msg.id, 42);
187 }
188
189 #[test]
190 fn test_protobuf_decoder_partial_data() {
191 let mut decoder = ProtobufDecoder::<TestMessage>::new();
192
193 let test_msg = TestMessage {
195 content: "Hello, Partial Data Test!".to_string(),
196 id: 123,
197 };
198
199 let encoded_msg = test_msg.encode_to_vec();
201 let mut prefixed_data = Vec::new();
202 prost::encoding::encode_varint(encoded_msg.len() as u64, &mut prefixed_data);
203 prefixed_data.extend_from_slice(&encoded_msg);
204
205 let half_point = prefixed_data.len() / 2;
207 decoder.add_data(&prefixed_data[..half_point]);
208
209 let result = decoder.decode_next().unwrap();
211 assert!(result.is_none());
212
213 decoder.add_data(&prefixed_data[half_point..]);
215
216 let decoded_msg = decoder.decode_next().unwrap().unwrap();
218 assert_eq!(decoded_msg.content, "Hello, Partial Data Test!");
219 assert_eq!(decoded_msg.id, 123);
220 }
221
222 #[test]
223 fn test_safe_protobuf_decode() {
224 let test_msg = TestMessage {
225 content: "Safe decode test".to_string(),
226 id: 999,
227 };
228
229 let encoded = test_msg.encode_to_vec();
230 let decoded = safe_protobuf_decode::<TestMessage>(&encoded).unwrap();
231
232 assert_eq!(decoded.content, "Safe decode test");
233 assert_eq!(decoded.id, 999);
234 }
235
236 #[test]
237 fn test_safe_string_decode() {
238 let valid_utf8 = b"Valid UTF-8 string";
239 let result = safe_string_decode(valid_utf8).unwrap();
240 assert_eq!(result, "Valid UTF-8 string");
241
242 let invalid_utf8 = &[0xFF, 0xFE, 0xFD]; let result = safe_string_decode(invalid_utf8);
245 assert!(result.is_err());
246 }
247}