1use async_stream::try_stream;
2use futures_util::Stream;
3use tokio::io::{AsyncRead, AsyncReadExt};
4
5use crate::error::Error;
6
7pub struct WireEntry {
8 pub prefix: [u8; 5],
9 pub content: Vec<u8>,
10}
11
12pub fn decode_segment_stream<R: AsyncRead + Unpin + Send + 'static>(
13 mut reader: R,
14) -> impl Stream<Item = Result<WireEntry, Error>> + Send + 'static {
15 try_stream! {
16 let mut count_buf = [0u8; 4];
17 reader.read_exact(&mut count_buf).await.map_err(|e| Error::Decode(format!("failed to read count: {e}")))?;
18 let count = u32::from_le_bytes(count_buf) as usize;
19
20 const MAX_SEGMENT_ENTRIES: usize = 1_048_576;
21 if count > MAX_SEGMENT_ENTRIES {
22 Err(Error::Decode(format!(
23 "entry count {count} exceeds maximum {MAX_SEGMENT_ENTRIES}"
24 )))?;
25 }
26
27 for i in 0..count {
28 let mut prefix = [0u8; 5];
29 reader.read_exact(&mut prefix).await.map_err(|e| Error::Decode(format!("entry {i}: truncated header: {e}")))?;
30
31 let mut len_buf = [0u8; 4];
32 reader.read_exact(&mut len_buf).await.map_err(|e| Error::Decode(format!("entry {i}: truncated header: {e}")))?;
33 let content_len = u32::from_le_bytes(len_buf) as usize;
34
35 let mut content = vec![0u8; content_len];
36 reader.read_exact(&mut content).await.map_err(|e| Error::Decode(format!("entry {i}: content truncated: {e}")))?;
37
38 yield WireEntry { prefix, content };
39 }
40 }
41}
42
43#[cfg(test)]
44mod tests {
45 use super::*;
46
47 fn wire_bytes(entries: &[([u8; 5], &[u8])]) -> Vec<u8> {
48 let mut buf = Vec::new();
49 buf.extend_from_slice(&(entries.len() as u32).to_le_bytes());
50 for (prefix, content) in entries {
51 buf.extend_from_slice(prefix);
52 buf.extend_from_slice(&(content.len() as u32).to_le_bytes());
53 buf.extend_from_slice(content);
54 }
55 buf
56 }
57
58 #[tokio::test]
59 async fn decode_empty_segment() {
60 let buf = wire_bytes(&[]);
61 let reader = tokio::io::BufReader::new(std::io::Cursor::new(buf));
62 let mut stream = Box::pin(decode_segment_stream(reader));
63 use futures_util::StreamExt;
64 assert!(stream.next().await.is_none());
65 }
66
67 #[tokio::test]
68 async fn decode_single_entry() {
69 let buf = wire_bytes(&[(*b"A3C01", b"hello world" as &[u8])]);
70 let reader = tokio::io::BufReader::new(std::io::Cursor::new(buf));
71 let mut stream = Box::pin(decode_segment_stream(reader));
72 use futures_util::StreamExt;
73
74 let entry = stream.next().await.unwrap().unwrap();
75 assert_eq!(&entry.prefix, b"A3C01");
76 assert_eq!(entry.content, b"hello world");
77 assert!(stream.next().await.is_none());
78 }
79
80 #[tokio::test]
81 async fn decode_truncated_header() {
82 let mut buf = Vec::new();
84 buf.extend_from_slice(&1u32.to_le_bytes());
85 buf.extend_from_slice(b"A3C0"); let reader = tokio::io::BufReader::new(std::io::Cursor::new(buf));
88 let mut stream = Box::pin(decode_segment_stream(reader));
89 use futures_util::StreamExt;
90
91 assert!(stream.next().await.unwrap().is_err());
92 }
93
94 #[tokio::test]
95 async fn decode_truncated_content() {
96 let mut buf = Vec::new();
97 buf.extend_from_slice(&1u32.to_le_bytes());
98 buf.extend_from_slice(b"A3C01");
99 buf.extend_from_slice(&100u32.to_le_bytes()); buf.extend_from_slice(b"short"); let reader = tokio::io::BufReader::new(std::io::Cursor::new(buf));
103 let mut stream = Box::pin(decode_segment_stream(reader));
104 use futures_util::StreamExt;
105
106 assert!(stream.next().await.unwrap().is_err());
107 }
108}