Skip to main content

flare_core/common/
protobuf_decoder.rs

1//! # 安全的 Protobuf 解码器(含粘包处理)
2//!
3//! 用于处理带长度前缀的 Protobuf 消息,防止将 varint 长度前缀误认为字符串内容
4//! 这解决了 protobuf string 字段前出现 "\x0c" 的问题,该问题是由于 length varint
5//! 被当成字符串解码导致的
6//!
7//! 此模块提供了专门的 Protobuf 序列化器实现,以替代基础的 ProtobufSerializer
8
9use bytes::BytesMut;
10use prost::Message;
11use std::io::{self, Cursor};
12
13/// Protobuf 消息解码器,支持粘包处理
14pub struct ProtobufDecoder<T> {
15    /// 消息类型
16    _phantom: std::marker::PhantomData<T>,
17    /// 缓冲区,用于处理粘包
18    buffer: BytesMut,
19}
20
21impl<T> ProtobufDecoder<T>
22where
23    T: Message + Default,
24{
25    /// 创建新的解码器
26    pub fn new() -> Self {
27        Self {
28            _phantom: std::marker::PhantomData,
29            buffer: BytesMut::new(),
30        }
31    }
32
33    /// 向解码器添加数据(处理粘包)
34    pub fn add_data(&mut self, data: &[u8]) {
35        self.buffer.extend_from_slice(data);
36    }
37
38    /// 解码下一个完整的消息
39    ///
40    /// 返回 Ok(Some(message)) 如果有足够的数据解码一个完整的消息
41    /// 返回 Ok(None) 如果数据不足(需要更多数据)
42    /// 返回 Err 如果解码失败
43    pub fn decode_next(&mut self) -> Result<Option<T>, Box<dyn std::error::Error + Send + Sync>> {
44        // 尝试解析长度前缀(varint 编码)
45        if let Some((message_len, offset)) = self.read_varint()? {
46            // 检查是否有足够的数据来解码完整的消息
47            if self.buffer.len() < offset + message_len {
48                // 数据不足,需要等待更多数据
49                return Ok(None);
50            }
51
52            // 提取消息数据
53            let message_data = self.buffer.split_to(offset + message_len).freeze();
54            let message_bytes = &message_data[offset..];
55
56            // 解码 protobuf 消息
57            let message = T::decode(message_bytes)?;
58            Ok(Some(message))
59        } else {
60            // 数据不足,无法读取完整的 varint
61            Ok(None)
62        }
63    }
64
65    /// 读取 varint 编码的长度
66    /// 返回 (length, bytes_consumed)
67    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            // varint 最多 10 个字节
74            if cursor.position() as usize >= self.buffer.len() {
75                // 数据不足,无法读取完整的 varint
76                return Ok(None);
77            }
78
79            let byte = cursor.get_ref()[cursor.position() as usize];
80            cursor.set_position(cursor.position() + 1);
81
82            // 提取低 7 位
83            let value = (byte & 0x7F) as u64;
84            result |= value << shift;
85
86            // 检查最高位是否为 1(表示还有后续字节)
87            if (byte & 0x80) == 0 {
88                // 完成了 varint 解码
89                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    /// 检查是否还有未处理的数据
102    pub fn has_remaining(&self) -> bool {
103        !self.buffer.is_empty()
104    }
105
106    /// 获取剩余数据长度
107    pub fn remaining_len(&self) -> usize {
108        self.buffer.len()
109    }
110
111    /// 清空缓冲区
112    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
126/// 安全的 protobuf 消息内容解码函数
127///
128/// 此函数安全地尝试解码 protobuf 数据,如果解码失败则返回错误而不是崩溃
129pub 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
139/// 安全的字符串解码函数(用于调试目的)
140///
141/// 此函数仅在确认数据是有效 UTF-8 时才进行转换
142pub 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    // 定义一个简单的测试消息
155    #[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        // 创建测试消息
168        let test_msg = TestMessage {
169            content: "Hello, World!".to_string(),
170            id: 42,
171        };
172
173        // 编码消息(带长度前缀)
174        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        // 添加数据到解码器
180        decoder.add_data(&prefixed_data);
181
182        // 解码消息
183        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        // 创建测试消息
194        let test_msg = TestMessage {
195            content: "Hello, Partial Data Test!".to_string(),
196            id: 123,
197        };
198
199        // 编码消息(带长度前缀)
200        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        // 只添加部分数据
206        let half_point = prefixed_data.len() / 2;
207        decoder.add_data(&prefixed_data[..half_point]);
208
209        // 尝试解码 - 应该返回 None(数据不足)
210        let result = decoder.decode_next().unwrap();
211        assert!(result.is_none());
212
213        // 添加剩余数据
214        decoder.add_data(&prefixed_data[half_point..]);
215
216        // 现在应该能够解码
217        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        // 测试无效的 UTF-8 数据
243        let invalid_utf8 = &[0xFF, 0xFE, 0xFD]; // 无效的 UTF-8 序列
244        let result = safe_string_decode(invalid_utf8);
245        assert!(result.is_err());
246    }
247}