Skip to main content

iris_core/protocols/stream/quic/
parser.rs

1//! Quic Header parser
2//! Custom Quic Parser with many design choices borrowed from
3//! [Wireshark Quic Disector](https://gitlab.com/wireshark/wireshark/-/blob/master/epan/dissectors/packet-quic.c)
4//!
5use crate::protocols::stream::quic::crypto::calc_init_keys;
6use crate::protocols::stream::quic::frame::QuicFrame;
7use crate::protocols::stream::quic::header::{
8    LongHeaderPacketType, QuicLongHeader, QuicShortHeader,
9};
10use crate::protocols::stream::quic::{QuicError, QuicPacket};
11use crate::protocols::stream::tls::Tls;
12use crate::protocols::stream::{
13    ConnParsable, L4Pdu, ParseResult, ParsingState, ProbeResult, Session, SessionData,
14};
15use byteorder::{BigEndian, ByteOrder};
16use nom::Err as NomErr;
17use std::collections::HashSet;
18use tls_parser::parse_tls_message_handshake;
19
20use super::QuicConn;
21
22#[derive(Debug)]
23pub struct QuicParser {
24    // /// Maps session ID to Quic transaction
25    // sessions: HashMap<usize, QuicPacket>,
26    // /// Total sessions ever seen (Running session ID)
27    // cnt: usize,
28    sessions: Vec<QuicConn>,
29}
30
31impl Default for QuicParser {
32    fn default() -> Self {
33        QuicParser {
34            sessions: vec![QuicConn::new()],
35        }
36    }
37}
38
39impl ConnParsable for QuicParser {
40    fn parse(&mut self, pdu: &L4Pdu) -> ParseResult {
41        let offset = pdu.offset();
42        let length = pdu.length();
43        if length == 0 {
44            return ParseResult::Skipped;
45        }
46
47        if let Ok(data) = (pdu.mbuf_ref()).get_data_slice(offset, length) {
48            if !self.sessions.is_empty() {
49                return self.sessions[0].parse_packet(data, pdu.dir);
50            }
51            ParseResult::Skipped
52        } else {
53            log::warn!("Malformed packet on parse");
54            ParseResult::Skipped
55        }
56    }
57
58    fn probe(&self, pdu: &L4Pdu) -> ProbeResult {
59        if pdu.length() < 5 {
60            return ProbeResult::Unsure;
61        }
62
63        let offset = pdu.offset();
64        let length = pdu.length();
65
66        if let Ok(data) = (pdu.mbuf).get_data_slice(offset, length) {
67            // Check if Fixed Bit is set
68            if (data[0] & 0x40) == 0 {
69                return ProbeResult::NotForUs;
70            }
71
72            if (data[0] & 0x80) != 0 {
73                // Potential Long Header
74                if data.len() < 6 {
75                    return ProbeResult::Unsure;
76                }
77
78                // Check if version is known
79                let version = ((data[1] as u32) << 24)
80                    | ((data[2] as u32) << 16)
81                    | ((data[3] as u32) << 8)
82                    | (data[4] as u32);
83                match QuicVersion::from_u32(version) {
84                    QuicVersion::Unknown => ProbeResult::NotForUs,
85                    _ => ProbeResult::Certain,
86                }
87            } else {
88                ProbeResult::Unsure
89            }
90        } else {
91            log::warn!("Malformed packet");
92            ProbeResult::Error
93        }
94    }
95
96    fn remove_session(&mut self, session_id: usize) -> Option<Session> {
97        self.sessions.pop().map(|quic| Session {
98            data: SessionData::Quic(Box::new(quic)),
99            id: session_id,
100        })
101    }
102
103    fn drain_sessions(&mut self) -> Vec<Session> {
104        self.sessions
105            .drain(..)
106            .map(|quic| Session {
107                data: SessionData::Quic(Box::new(quic)),
108                id: 0,
109            })
110            .collect()
111    }
112
113    fn session_parsed_state(&self) -> ParsingState {
114        ParsingState::Parsing
115    }
116
117    // Temporary - not supported for QUIC parser.
118    fn body_offset(&mut self) -> Option<usize> {
119        None
120    }
121}
122
123/// Supported Quic Versions
124#[derive(Debug, PartialEq, Eq, Hash)]
125#[repr(u32)]
126pub enum QuicVersion {
127    ReservedNegotiation = 0x00000000,
128    Rfc9000 = 0x00000001, // Quic V1
129    Rfc9369 = 0x6b3343cf, // Quic V2
130    Draft27 = 0xff00001b, // Quic draft 27
131    Draft28 = 0xff00001c, // Quic draft 28
132    Draft29 = 0xff00001d, // Quic draft 29
133    Mvfst27 = 0xfaceb002, // Facebook Implementation of draft 27
134    Unknown,
135}
136
137impl QuicVersion {
138    pub fn from_u32(version: u32) -> Self {
139        match version {
140            0x00000000 => QuicVersion::ReservedNegotiation,
141            0x00000001 => QuicVersion::Rfc9000,
142            0x6b3343cf => QuicVersion::Rfc9369,
143            0xff00001b => QuicVersion::Draft27,
144            0xff00001c => QuicVersion::Draft28,
145            0xff00001d => QuicVersion::Draft29,
146            0xfaceb002 => QuicVersion::Mvfst27,
147            _ => QuicVersion::Unknown,
148        }
149    }
150}
151
152impl QuicPacket {
153    /// Processes the connection ID bytes array to a hex string
154    pub fn vec_u8_to_hex_string(vec: &[u8]) -> String {
155        vec.iter()
156            .map(|&byte| format!("{:02x}", byte))
157            .collect::<Vec<String>>()
158            .join("")
159    }
160
161    // Calculate the length of a variable length encoding
162    // See RFC 9000 Section 16 for details
163    pub fn get_var_len(a: u8) -> Result<usize, QuicError> {
164        let two_msb = a >> 6;
165        match two_msb {
166            0b00 => Ok(1),
167            0b01 => Ok(2),
168            0b10 => Ok(4),
169            0b11 => Ok(8),
170            _ => Err(QuicError::UnsupportedVarLen),
171        }
172    }
173
174    // Masks variable length encoding and returns u64 value for remainder of field
175    pub fn slice_to_u64(data: &[u8]) -> Result<u64, QuicError> {
176        if data.len() > 8 {
177            return Err(QuicError::UnsupportedVarLen);
178        }
179
180        let mut result: u64 = 0;
181        for &byte in data {
182            result = (result << 8) | u64::from(byte);
183        }
184        result &= !(0b11 << ((data.len() * 8) - 2)); // Var length encoding mask
185        Ok(result)
186    }
187
188    pub fn access_data(data: &[u8], start: usize, end: usize) -> Result<&[u8], QuicError> {
189        if end < start {
190            return Err(QuicError::InvalidDataIndices);
191        }
192        if data.len() < end {
193            return Err(QuicError::PacketTooShort);
194        }
195        Ok(&data[start..end])
196    }
197
198    /// Parses Quic packet from bytes
199    pub fn parse_from(
200        conn: &mut QuicConn,
201        data: &[u8],
202        mut offset: usize,
203        dir: bool,
204    ) -> Result<(QuicPacket, usize), QuicError> {
205        let packet_header_byte = QuicPacket::access_data(data, offset, offset + 1)?[0];
206        offset += 1;
207        // Check the fixed bit
208        if (packet_header_byte & 0x40) == 0 {
209            return Err(QuicError::FixedBitNotSet);
210        }
211        // Check the Header form
212        if (packet_header_byte & 0x80) != 0 {
213            // Long Header
214            // Parse packet type
215            let packet_type = LongHeaderPacketType::from_u8((packet_header_byte & 0x30) >> 4)?;
216            let type_specific = packet_header_byte & 0x0F; // Remainder of information from header byte, Reserved and protected packet number length
217                                                           // Parse version
218            let version_bytes = QuicPacket::access_data(data, offset, offset + 4)?;
219            let version = ((version_bytes[0] as u32) << 24)
220                | ((version_bytes[1] as u32) << 16)
221                | ((version_bytes[2] as u32) << 8)
222                | (version_bytes[3] as u32);
223            if QuicVersion::from_u32(version) == QuicVersion::Unknown {
224                return Err(QuicError::UnknownVersion);
225            }
226            offset += 4;
227            // Parse DCID
228            let dcid_len = QuicPacket::access_data(data, offset, offset + 1)?[0];
229            offset += 1;
230            let dcid_bytes = QuicPacket::access_data(data, offset, offset + dcid_len as usize)?;
231            let dcid = QuicPacket::vec_u8_to_hex_string(dcid_bytes);
232            if dcid_len > 0 && !conn.cids.contains(&dcid) {
233                conn.cids.insert(dcid.clone());
234            }
235            offset += dcid_len as usize;
236            // Parse SCID
237            let scid_len = QuicPacket::access_data(data, offset, offset + 1)?[0];
238            offset += 1;
239            let scid_bytes = QuicPacket::access_data(data, offset, offset + scid_len as usize)?;
240            let scid = QuicPacket::vec_u8_to_hex_string(scid_bytes);
241            if scid_len > 0 && !conn.cids.contains(&scid) {
242                conn.cids.insert(scid.clone());
243            }
244            offset += scid_len as usize;
245
246            let token_len;
247            let token;
248            let packet_len;
249            let retry_tag;
250            let decrypted_payload;
251            // Parse packet type specific fields
252            match packet_type {
253                LongHeaderPacketType::Initial => {
254                    retry_tag = None;
255                    // Parse token
256                    let token_len_len = QuicPacket::get_var_len(
257                        QuicPacket::access_data(data, offset, offset + 1)?[0],
258                    )?;
259                    let token_len_bytes =
260                        QuicPacket::access_data(data, offset, offset + token_len_len)?;
261                    token_len = Some(QuicPacket::slice_to_u64(token_len_bytes)?);
262                    offset += token_len_len;
263                    let token_bytes = QuicPacket::access_data(
264                        data,
265                        offset,
266                        offset + token_len.unwrap() as usize,
267                    )?;
268                    token = Some(QuicPacket::vec_u8_to_hex_string(token_bytes));
269                    offset += token_len.unwrap() as usize;
270                    // Parse payload length
271                    let packet_len_len = QuicPacket::get_var_len(
272                        QuicPacket::access_data(data, offset, offset + 1)?[0],
273                    )?;
274                    let packet_len_bytes =
275                        QuicPacket::access_data(data, offset, offset + packet_len_len)?;
276                    packet_len = Some(QuicPacket::slice_to_u64(packet_len_bytes)?);
277                    offset += packet_len_len;
278                    if conn.client_opener.is_none() {
279                        // Derive initial keys
280                        let [client_opener, server_opener] = calc_init_keys(dcid_bytes, version)?;
281                        conn.client_opener = Some(client_opener);
282                        conn.server_opener = Some(server_opener);
283                    }
284                    // Calculate HP
285                    let sample_len = conn.client_opener.as_ref().unwrap().sample_len();
286                    let hp_sample =
287                        QuicPacket::access_data(data, offset + 4, offset + 4 + sample_len)?;
288                    let mask = if dir {
289                        conn.client_opener.as_ref().unwrap().new_mask(hp_sample)?
290                    } else {
291                        conn.server_opener.as_ref().unwrap().new_mask(hp_sample)?
292                    };
293                    // Remove HP from packet header byte
294                    let unprotected_header = packet_header_byte ^ (mask[0] & 0b00001111);
295                    if (unprotected_header >> 2) & 0b00000011 != 0 {
296                        return Err(QuicError::FailedHeaderProtection);
297                    }
298                    // Parse packet number
299                    let packet_num_len = ((unprotected_header & 0b00000011) + 1) as usize;
300                    let packet_number_bytes =
301                        QuicPacket::access_data(data, offset, offset + packet_num_len)?;
302                    let mut packet_number = vec![0; 4 - packet_num_len];
303                    for i in 0..packet_num_len {
304                        packet_number.push(packet_number_bytes[i] ^ mask[i + 1]);
305                    }
306
307                    let initial_packet_number_bytes = &packet_number[4 - packet_num_len..];
308                    let packet_number_int = BigEndian::read_i32(&packet_number);
309                    offset += packet_num_len;
310                    // Parse the encrypted payload
311                    let tag_len = conn.client_opener.as_ref().unwrap().alg().tag_len();
312                    if (packet_len.unwrap() as usize) < (tag_len + packet_num_len) {
313                        return Err(QuicError::PacketTooShort);
314                    }
315                    let cipher_text_len = packet_len.unwrap() as usize - tag_len - packet_num_len;
316                    let mut encrypted_payload =
317                        QuicPacket::access_data(data, offset, offset + cipher_text_len)?.to_vec();
318                    offset += cipher_text_len;
319                    // Parse auth tag
320                    let tag = QuicPacket::access_data(data, offset, offset + tag_len)?;
321                    offset += tag_len;
322                    // Reconstruct authenticated data
323                    let mut ad = Vec::new();
324                    ad.append(&mut [unprotected_header].to_vec());
325                    ad.append(&mut version_bytes.to_vec());
326                    ad.append(&mut [dcid_len].to_vec());
327                    ad.append(&mut dcid_bytes.to_vec());
328                    ad.append(&mut [scid_len].to_vec());
329                    ad.append(&mut scid_bytes.to_vec());
330                    ad.append(&mut token_len_bytes.to_vec());
331                    ad.append(&mut token_bytes.to_vec());
332                    ad.append(&mut packet_len_bytes.to_vec());
333                    ad.append(&mut initial_packet_number_bytes.to_vec());
334                    // Decrypt payload with proper keys based on traffic direction
335                    if dir {
336                        decrypted_payload =
337                            Some(conn.client_opener.as_ref().unwrap().open_with_u64_counter(
338                                packet_number_int as u64,
339                                &ad,
340                                &mut encrypted_payload,
341                                tag,
342                            )?);
343                    } else {
344                        decrypted_payload =
345                            Some(conn.server_opener.as_ref().unwrap().open_with_u64_counter(
346                                packet_number_int as u64,
347                                &ad,
348                                &mut encrypted_payload,
349                                tag,
350                            )?);
351                    }
352                }
353                LongHeaderPacketType::ZeroRTT | LongHeaderPacketType::Handshake => {
354                    token_len = None;
355                    token = None;
356                    retry_tag = None;
357                    decrypted_payload = None;
358                    // Parse payload length
359                    let packet_len_len = QuicPacket::get_var_len(
360                        QuicPacket::access_data(data, offset, offset + 1)?[0],
361                    )?;
362                    packet_len = Some(QuicPacket::slice_to_u64(QuicPacket::access_data(
363                        data,
364                        offset,
365                        offset + packet_len_len,
366                    )?)?);
367                    offset += packet_len_len;
368                    offset += packet_len.unwrap() as usize;
369                }
370                LongHeaderPacketType::Retry => {
371                    packet_len = None;
372                    decrypted_payload = None;
373                    if data.len() > (offset + 16) {
374                        token_len = Some((data.len() - offset - 16) as u64);
375                    } else {
376                        return Err(QuicError::PacketTooShort);
377                    }
378                    // Parse retry token
379                    let token_bytes = QuicPacket::access_data(
380                        data,
381                        offset,
382                        offset + token_len.unwrap() as usize,
383                    )?;
384                    token = Some(QuicPacket::vec_u8_to_hex_string(token_bytes));
385                    offset += token_len.unwrap() as usize;
386                    // Parse retry tag
387                    let retry_tag_bytes = QuicPacket::access_data(data, offset, offset + 16)?;
388                    retry_tag = Some(QuicPacket::vec_u8_to_hex_string(retry_tag_bytes));
389                    offset += 16;
390                }
391            }
392
393            let mut frames: Option<Vec<QuicFrame>> = None;
394            // If decrypted payload is not None, parse the frames
395            if let Some(frame_bytes) = decrypted_payload {
396                let (q_frames, crypto_chunks) = QuicFrame::parse_frames(&frame_bytes)?;
397                frames = Some(q_frames);
398
399                // Drop CRYPTO chunks at their absolute cryptostream offsets
400                // into the per-direction sparse map. Multiple CRYPTO frames
401                // at non-contiguous offsets within a single packet are
402                // common (Chrome QUIC interleaves PING/PADDING with split
403                // ClientHello CRYPTO frames).
404                let crypto_map: &mut std::collections::BTreeMap<u64, Vec<u8>> = if dir {
405                    &mut conn.client_crypto
406                } else {
407                    &mut conn.server_crypto
408                };
409                for (off, data) in crypto_chunks {
410                    crypto_map.entry(off).or_insert(data);
411                }
412
413                // Walk the map from `consumed` onward, building the longest
414                // contiguous buffer. Feed it to the TLS parser; on Ok,
415                // advance `consumed` past the bytes it consumed.
416                let consumed_ref = if dir {
417                    &mut conn.client_consumed
418                } else {
419                    &mut conn.server_consumed
420                };
421                let mut contiguous: Vec<u8> = Vec::new();
422                let mut next = *consumed_ref;
423                for (&off, data) in crypto_map.iter() {
424                    let chunk_end = off + data.len() as u64;
425                    if chunk_end <= next {
426                        continue;
427                    }
428                    if off > next {
429                        break;
430                    }
431                    let skip = (next - off) as usize;
432                    contiguous.extend_from_slice(&data[skip..]);
433                    next = chunk_end;
434                }
435                if !contiguous.is_empty() {
436                    match parse_tls_message_handshake(&contiguous) {
437                        Ok((rest, msg)) => {
438                            conn.tls.parse_message_level(&msg, dir);
439                            let used = (contiguous.len() - rest.len()) as u64;
440                            *consumed_ref += used;
441                        }
442                        Err(NomErr::Incomplete(_)) => {
443                            // Wait for more bytes from later packets.
444                        }
445                        Err(_) => {
446                            // Skip this message rather than block the
447                            // connection.
448                            *consumed_ref = next;
449                        }
450                    }
451                }
452            }
453
454            Ok((
455                QuicPacket {
456                    payload_bytes_count: packet_len,
457                    short_header: None,
458                    long_header: Some(QuicLongHeader {
459                        packet_type,
460                        type_specific,
461                        version,
462                        dcid_len,
463                        dcid,
464                        scid_len,
465                        scid,
466                        token_len,
467                        token,
468                        retry_tag,
469                    }),
470                    frames,
471                },
472                offset,
473            ))
474        } else {
475            // Short Header
476            let mut dcid_len = 20;
477            if data.len() < 1 + dcid_len {
478                dcid_len = data.len() - 1;
479            }
480            // Parse DCID
481            let dcid_hex = QuicPacket::vec_u8_to_hex_string(QuicPacket::access_data(
482                data,
483                offset,
484                offset + dcid_len,
485            )?);
486            let mut dcid = None;
487            for cid in &conn.cids {
488                if dcid_hex.starts_with(cid) {
489                    dcid_len = cid.chars().count() / 2;
490                    dcid = Some(cid.clone());
491                }
492            }
493            offset += dcid_len;
494            // Counts all bytes remaining
495            let payload_bytes_count = (data.len() - offset) as u64;
496            offset += payload_bytes_count as usize;
497            Ok((
498                QuicPacket {
499                    short_header: Some(QuicShortHeader { dcid }),
500                    long_header: None,
501                    payload_bytes_count: Some(payload_bytes_count),
502                    frames: None,
503                },
504                offset,
505            ))
506        }
507    }
508}
509
510impl QuicConn {
511    pub(crate) fn new() -> QuicConn {
512        QuicConn {
513            packets: Vec::new(),
514            cids: HashSet::new(),
515            tls: Tls::new(),
516            client_opener: None,
517            server_opener: None,
518            client_crypto: std::collections::BTreeMap::new(),
519            server_crypto: std::collections::BTreeMap::new(),
520            client_consumed: 0,
521            server_consumed: 0,
522        }
523    }
524
525    fn parse_packet(&mut self, data: &[u8], direction: bool) -> ParseResult {
526        let mut offset = 0;
527        // Iterate over all of the data in the datagram
528        // Parse as many QUIC packets as possible
529        // NICE-TO-HAVE: identify padding appended to datagram
530        while data.len() > offset {
531            if let Ok((quic, off)) = QuicPacket::parse_from(self, data, offset, direction) {
532                self.packets.push(quic);
533                offset = off;
534            } else {
535                return ParseResult::Skipped;
536            }
537        }
538        if self
539            .packets
540            .last()
541            .is_some_and(|p| p.short_header.is_some())
542        {
543            return ParseResult::HeadersDone(0);
544        }
545        ParseResult::Continue(0)
546    }
547}