Skip to main content

compression_codecs/gzip/
header.rs

1use compression_core::util::PartialBuffer;
2use flate2::Crc;
3use std::io;
4
5#[derive(Debug, Default)]
6struct Flags {
7    _ascii: bool,
8    crc: bool,
9    extra: bool,
10    filename: bool,
11    comment: bool,
12}
13
14#[derive(Debug, Default)]
15pub(super) struct Header {
16    flags: Flags,
17}
18
19#[derive(Debug)]
20enum State {
21    Fixed(PartialBuffer<[u8; 10]>),
22    ExtraLen(PartialBuffer<[u8; 2]>),
23    Extra(usize),
24    Filename,
25    Comment,
26    Crc(PartialBuffer<[u8; 2]>),
27    Done,
28}
29
30impl Default for State {
31    fn default() -> Self {
32        State::Fixed(<_>::default())
33    }
34}
35
36#[derive(Debug, Default)]
37pub(super) struct Parser {
38    state: State,
39    header: Header,
40}
41
42/// The fixed prefix every gzip member starts with: magic bytes plus the deflate method.
43const MAGIC: [u8; 3] = [0x1f, 0x8b, 0x08];
44
45impl Header {
46    fn parse(input: &[u8; 10]) -> io::Result<Self> {
47        if input[0..3] != MAGIC {
48            return Err(io::Error::new(
49                io::ErrorKind::InvalidData,
50                "Invalid gzip header",
51            ));
52        }
53
54        let flag = input[3];
55
56        let flags = Flags {
57            _ascii: (flag & 0b0000_0001) != 0,
58            crc: (flag & 0b0000_0010) != 0,
59            extra: (flag & 0b0000_0100) != 0,
60            filename: (flag & 0b0000_1000) != 0,
61            comment: (flag & 0b0001_0000) != 0,
62        };
63
64        Ok(Header { flags })
65    }
66}
67
68fn consume_input(crc: &mut Crc, n: usize, input: &mut PartialBuffer<&[u8]>) {
69    crc.update(&input.unwritten()[..n]);
70    input.advance(n);
71}
72
73fn consume_cstr(crc: &mut Crc, input: &mut PartialBuffer<&[u8]>) -> Option<()> {
74    if let Some(len) = memchr::memchr(0, input.unwritten()) {
75        consume_input(crc, len + 1, input);
76        Some(())
77    } else {
78        consume_input(crc, input.unwritten().len(), input);
79        None
80    }
81}
82
83impl Parser {
84    pub(super) fn input(
85        &mut self,
86        crc: &mut Crc,
87        input: &mut PartialBuffer<&[u8]>,
88    ) -> io::Result<Option<Header>> {
89        loop {
90            match &mut self.state {
91                State::Fixed(data) => {
92                    data.copy_unwritten_from(input);
93
94                    // `MAGIC` decides validity, so reject as soon as the bytes seen so far
95                    // contradict it instead of waiting for all 10. A reader that stalls part
96                    // way through the header would otherwise never resolve, which matters for
97                    // `multiple_members` over a live stream.
98                    let seen = data.written();
99                    let checked = seen.len().min(MAGIC.len());
100                    if seen[..checked] != MAGIC[..checked] {
101                        return Err(io::Error::new(
102                            io::ErrorKind::InvalidData,
103                            "Invalid gzip header",
104                        ));
105                    }
106
107                    if data.unwritten().is_empty() {
108                        let data = data.get_mut();
109                        crc.update(data);
110                        self.header = Header::parse(data)?;
111                        self.state = State::ExtraLen(<_>::default());
112                    } else {
113                        break Ok(None);
114                    }
115                }
116
117                State::ExtraLen(data) => {
118                    if !self.header.flags.extra {
119                        self.state = State::Filename;
120                        continue;
121                    }
122
123                    data.copy_unwritten_from(input);
124
125                    if data.unwritten().is_empty() {
126                        let data = data.get_mut();
127                        crc.update(data);
128                        let len = u16::from_le_bytes(*data);
129                        self.state = State::Extra(len.into());
130                    } else {
131                        break Ok(None);
132                    }
133                }
134
135                State::Extra(bytes_to_consume) => {
136                    let n = input.unwritten().len().min(*bytes_to_consume);
137                    *bytes_to_consume -= n;
138                    consume_input(crc, n, input);
139
140                    if *bytes_to_consume == 0 {
141                        self.state = State::Filename;
142                    } else {
143                        break Ok(None);
144                    }
145                }
146
147                State::Filename => {
148                    if !self.header.flags.filename {
149                        self.state = State::Comment;
150                        continue;
151                    }
152
153                    if consume_cstr(crc, input).is_none() {
154                        break Ok(None);
155                    }
156
157                    self.state = State::Comment;
158                }
159
160                State::Comment => {
161                    if !self.header.flags.comment {
162                        self.state = State::Crc(<_>::default());
163                        continue;
164                    }
165
166                    if consume_cstr(crc, input).is_none() {
167                        break Ok(None);
168                    }
169
170                    self.state = State::Crc(<_>::default());
171                }
172
173                State::Crc(data) => {
174                    let header = std::mem::take(&mut self.header);
175
176                    if !self.header.flags.crc {
177                        self.state = State::Done;
178                        break Ok(Some(header));
179                    }
180
181                    data.copy_unwritten_from(input);
182
183                    break if data.unwritten().is_empty() {
184                        let data = data.take().into_inner();
185                        self.state = State::Done;
186                        let checksum = crc.sum().to_le_bytes();
187
188                        if data == checksum[..2] {
189                            Ok(Some(header))
190                        } else {
191                            Err(io::Error::new(
192                                io::ErrorKind::InvalidData,
193                                "CRC computed for header does not match",
194                            ))
195                        }
196                    } else {
197                        Ok(None)
198                    };
199                }
200
201                State::Done => break Err(io::Error::other("parser used after done")),
202            }
203        }
204    }
205}