Skip to main content

hibp_sync_client/
wire.rs

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        // count=1 but only 4 bytes of the 9-byte entry header follow
83        let mut buf = Vec::new();
84        buf.extend_from_slice(&1u32.to_le_bytes());
85        buf.extend_from_slice(b"A3C0"); // only 4 bytes of prefix (needs 5+4)
86
87        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()); // claims 100 bytes
100        buf.extend_from_slice(b"short"); // only 5
101
102        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}