Skip to main content

s2n_quic_dc/packet/stream/
decoder.rs

1// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use crate::{
5    credentials::Credentials,
6    crypto,
7    packet::{
8        control::decoder::ControlFramesMut,
9        stream::{self, RelativeRetransmissionOffset, Tag},
10        WireVersion,
11    },
12};
13use core::{fmt, mem::size_of};
14use s2n_codec::{
15    decoder_invariant, CheckedRange, DecoderBufferMut, DecoderBufferMutResult as R, DecoderError,
16};
17use s2n_quic_core::{assume, ensure, varint::VarInt};
18
19type PacketNumber = VarInt;
20
21pub trait Validator {
22    fn validate_tag(&mut self, tag: Tag) -> Result<(), DecoderError>;
23}
24
25impl Validator for () {
26    #[inline]
27    fn validate_tag(&mut self, _tag: Tag) -> Result<(), DecoderError> {
28        Ok(())
29    }
30}
31
32impl Validator for Tag {
33    #[inline]
34    fn validate_tag(&mut self, actual: Tag) -> Result<(), DecoderError> {
35        decoder_invariant!(*self == actual, "unexpected packet type");
36        Ok(())
37    }
38}
39
40impl<A, B> Validator for (A, B)
41where
42    A: Validator,
43    B: Validator,
44{
45    #[inline]
46    fn validate_tag(&mut self, tag: Tag) -> Result<(), DecoderError> {
47        self.0.validate_tag(tag)?;
48        self.1.validate_tag(tag)?;
49        Ok(())
50    }
51}
52
53#[derive(Clone, Debug, PartialEq, Eq)]
54pub struct Owned {
55    pub tag: Tag,
56    pub wire_version: WireVersion,
57    pub credentials: Credentials,
58    pub source_queue_id: Option<VarInt>,
59    pub stream_id: stream::Id,
60    pub original_packet_number: PacketNumber,
61    pub packet_number: PacketNumber,
62    pub retransmission_packet_number_offset: u8,
63    pub next_expected_control_packet: PacketNumber,
64    pub stream_offset: VarInt,
65    pub final_offset: Option<VarInt>,
66    pub application_header: Vec<u8>,
67    pub control_data: Vec<u8>,
68    pub payload: Vec<u8>,
69    pub auth_tag: Vec<u8>,
70}
71
72impl<'a> From<Packet<'a>> for Owned {
73    fn from(packet: Packet<'a>) -> Self {
74        let application_header = packet.application_header().to_vec();
75        let control_data = packet.control_data().to_vec();
76
77        Self {
78            tag: packet.tag,
79            wire_version: packet.wire_version,
80            credentials: packet.credentials,
81            source_queue_id: packet.source_queue_id,
82            stream_id: packet.stream_id,
83            original_packet_number: packet.original_packet_number,
84            packet_number: packet.packet_number,
85            retransmission_packet_number_offset: packet.retransmission_packet_number_offset,
86            next_expected_control_packet: packet.next_expected_control_packet,
87            stream_offset: packet.stream_offset,
88            final_offset: packet.final_offset,
89            application_header,
90            control_data,
91            payload: packet.payload.to_vec(),
92            auth_tag: packet.auth_tag.to_vec(),
93        }
94    }
95}
96
97pub struct Packet<'a> {
98    tag: Tag,
99    wire_version: WireVersion,
100    credentials: Credentials,
101    source_queue_id: Option<VarInt>,
102    stream_id: stream::Id,
103    original_packet_number: PacketNumber,
104    packet_number: PacketNumber,
105    retransmission_packet_number_offset: u8,
106    next_expected_control_packet: PacketNumber,
107    stream_offset: VarInt,
108    final_offset: Option<VarInt>,
109    header: &'a mut [u8],
110    application_header: CheckedRange,
111    control_data: CheckedRange,
112    payload: &'a mut [u8],
113    auth_tag: &'a mut [u8],
114}
115
116impl fmt::Debug for Packet<'_> {
117    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
118        f.debug_struct("stream::Packet")
119            .field("tag", &self.tag)
120            .field("wire_version", &self.wire_version)
121            .field("credentials", &self.credentials)
122            .field("source_queue_id", &self.source_queue_id)
123            .field("stream_id", &self.stream_id)
124            .field("packet_number", &self.packet_number())
125            .field("stream_offset", &self.stream_offset)
126            .field("final_offset", &self.final_offset)
127            .field("header_len", &self.header.len())
128            .field("payload_len", &self.payload.len())
129            .field("auth_tag_len", &self.auth_tag.len())
130            .finish()
131    }
132}
133
134impl Packet<'_> {
135    #[inline]
136    pub fn tag(&self) -> Tag {
137        self.tag
138    }
139
140    #[inline]
141    pub fn wire_version(&self) -> WireVersion {
142        self.wire_version
143    }
144
145    #[inline]
146    pub fn credentials(&self) -> &Credentials {
147        &self.credentials
148    }
149
150    #[inline]
151    pub fn source_queue_id(&self) -> Option<VarInt> {
152        self.source_queue_id
153    }
154
155    #[inline]
156    pub fn stream_id(&self) -> &stream::Id {
157        &self.stream_id
158    }
159
160    #[inline]
161    pub fn packet_number(&self) -> PacketNumber {
162        self.packet_number
163    }
164
165    #[inline]
166    pub fn is_retransmission(&self) -> bool {
167        self.packet_number != self.original_packet_number
168    }
169
170    #[inline]
171    pub fn next_expected_control_packet(&self) -> PacketNumber {
172        self.next_expected_control_packet
173    }
174
175    #[inline]
176    pub fn stream_offset(&self) -> VarInt {
177        self.stream_offset
178    }
179
180    #[inline]
181    pub fn final_offset(&self) -> Option<VarInt> {
182        self.final_offset
183    }
184
185    #[inline]
186    pub fn is_fin(&self) -> bool {
187        self.final_offset()
188            .and_then(|offset| offset.checked_sub(self.stream_offset))
189            .and_then(|offset| {
190                let len = VarInt::try_from(self.payload.len()).ok()?;
191                offset.checked_sub(len)
192            })
193            .is_some_and(|v| *v == 0)
194    }
195
196    #[inline]
197    pub fn application_header(&self) -> &[u8] {
198        self.application_header.get(self.header)
199    }
200
201    #[inline]
202    pub fn control_data(&self) -> &[u8] {
203        self.control_data.get(self.header)
204    }
205
206    #[inline]
207    pub fn control_frames_mut(&mut self) -> ControlFramesMut<'_> {
208        ControlFramesMut::new(self.control_data.get_mut(self.header))
209    }
210
211    #[inline]
212    pub fn header(&self) -> &[u8] {
213        self.header
214    }
215
216    #[inline]
217    pub fn payload(&self) -> &[u8] {
218        self.payload
219    }
220
221    #[inline]
222    pub fn auth_tag(&self) -> &[u8] {
223        self.auth_tag
224    }
225
226    #[inline]
227    pub fn payload_mut(&mut self) -> &mut [u8] {
228        self.payload
229    }
230
231    #[inline]
232    pub fn total_len(&self) -> usize {
233        self.header.len() + self.payload.len() + self.auth_tag.len()
234    }
235
236    #[inline]
237    pub fn decrypt<D, C>(
238        &mut self,
239        d: &D,
240        c: &C,
241        payload_out: &mut crypto::UninitSlice,
242    ) -> Result<(), crypto::open::Error>
243    where
244        D: crypto::open::Application,
245        C: crypto::open::control::Stream,
246    {
247        let key_phase = self.tag.key_phase();
248        let space = self.remove_retransmit(c)?;
249
250        let nonce = self.original_packet_number.as_u64();
251        let header = &self.header;
252        let payload = &self.payload;
253        let auth_tag = &self.auth_tag;
254
255        match space {
256            stream::PacketSpace::Stream => {
257                d.decrypt(key_phase, nonce, header, payload, auth_tag, payload_out)?;
258            }
259            stream::PacketSpace::Recovery => {
260                // recovery/probe packets cannot have payloads
261                ensure!(payload.is_empty(), Err(crypto::open::Error::MacOnly));
262                c.verify(header, auth_tag)?;
263            }
264        }
265
266        Ok(())
267    }
268
269    #[inline]
270    pub fn decrypt_in_place<D, C>(&mut self, d: &D, c: &C) -> Result<(), crypto::open::Error>
271    where
272        D: crypto::open::Application,
273        C: crypto::open::control::Stream,
274    {
275        let key_phase = self.tag.key_phase();
276        let space = self.remove_retransmit(c)?;
277
278        let nonce = self.original_packet_number.as_u64();
279        let header = &self.header;
280
281        match space {
282            stream::PacketSpace::Stream => {
283                d.decrypt_in_place(key_phase, nonce, header, self.payload, self.auth_tag)?;
284            }
285            stream::PacketSpace::Recovery => {
286                // recovery/probe packets cannot have payloads
287                ensure!(self.payload.is_empty(), Err(crypto::open::Error::MacOnly));
288                c.verify(header, self.auth_tag)?;
289            }
290        }
291
292        Ok(())
293    }
294
295    #[inline]
296    fn remove_retransmit<C>(&mut self, c: &C) -> Result<stream::PacketSpace, crypto::open::Error>
297    where
298        C: crypto::open::control::Stream,
299    {
300        let space = self.tag.packet_space();
301        let original_packet_number = self.original_packet_number;
302        let retransmission_packet_number = self.packet_number;
303
304        if original_packet_number != retransmission_packet_number {
305            c.retransmission_tag(
306                original_packet_number.as_u64(),
307                retransmission_packet_number.as_u64(),
308                self.auth_tag,
309            )?;
310            // clear the recovery packet bit, since this is a retransmission
311            self.header[0] &= !super::Tag::IS_RECOVERY_PACKET;
312
313            // update the retransmission offset to the zero value
314            let offset = self.retransmission_packet_number_offset as usize;
315            let range = offset..offset + size_of::<RelativeRetransmissionOffset>();
316            self.header[range].copy_from_slice(&[0; size_of::<RelativeRetransmissionOffset>()]);
317
318            Ok(stream::PacketSpace::Stream)
319        } else {
320            Ok(space)
321        }
322    }
323
324    #[inline]
325    #[cfg(debug_assertions)]
326    #[expect(
327        clippy::panic_in_result_fn,
328        reason = "debug-assertions-only build that round-trip checks the retransmission re-encode against a snapshot; the asserts validate internal encode/decode consistency"
329    )]
330    pub fn retransmit<K>(
331        buffer: DecoderBufferMut,
332        space: stream::PacketSpace,
333        retransmission_packet_number: VarInt,
334        key: &K,
335    ) -> Result<(), DecoderError>
336    where
337        K: crypto::seal::control::Stream,
338    {
339        let buffer = buffer.into_less_safe_slice();
340
341        let mut before = Self::snapshot(buffer, key.tag_len());
342        // update the expected packet space with the new one
343        before.tag.set_packet_space(space);
344        // the auth tag will have changed so clear it
345        before.auth_tag.clear();
346
347        Self::retransmit_impl(
348            DecoderBufferMut::new(buffer),
349            space,
350            retransmission_packet_number,
351            key,
352        )?;
353
354        let mut after = Self::snapshot(buffer, key.tag_len());
355        assert_eq!(after.packet_number, retransmission_packet_number);
356        after.packet_number = before.packet_number;
357        // the auth tag will have changed so clear it
358        after.auth_tag.clear();
359
360        assert_eq!(before, after);
361
362        Ok(())
363    }
364
365    #[inline]
366    #[cfg(not(debug_assertions))]
367    pub fn retransmit<K>(
368        buffer: DecoderBufferMut,
369        space: stream::PacketSpace,
370        retransmission_packet_number: VarInt,
371        key: &K,
372    ) -> Result<(), DecoderError>
373    where
374        K: crypto::seal::control::Stream,
375    {
376        Self::retransmit_impl(buffer, space, retransmission_packet_number, key)
377    }
378
379    #[inline]
380    #[cfg(debug_assertions)]
381    fn snapshot(buffer: &mut [u8], crypto_tag_len: usize) -> Owned {
382        let buffer = DecoderBufferMut::new(buffer);
383        #[expect(
384            clippy::unwrap_used,
385            reason = "debug-assertions-only snapshot helper decoding a buffer that was just produced by this crate's encoder, so it is guaranteed to be well-formed"
386        )]
387        let (packet, _buffer) = Self::decode(buffer, (), crypto_tag_len).unwrap();
388        packet.into()
389    }
390
391    #[inline(always)]
392    fn retransmit_impl<K>(
393        buffer: DecoderBufferMut,
394        space: stream::PacketSpace,
395        retransmission_packet_number: VarInt,
396        key: &K,
397    ) -> Result<(), DecoderError>
398    where
399        K: crypto::seal::control::Stream,
400    {
401        unsafe {
402            assume!(key.tag_len() >= 16, "tag len needs to be at least 16 bytes");
403        }
404
405        let (tag_slice, buffer) = buffer.decode_slice(1)?;
406
407        let tag: super::Tag = {
408            let tag_slice = tag_slice.into_less_safe_slice();
409
410            match space {
411                stream::PacketSpace::Stream => {
412                    tag_slice[0] &= !super::Tag::IS_RECOVERY_PACKET;
413                }
414                stream::PacketSpace::Recovery => {
415                    tag_slice[0] |= super::Tag::IS_RECOVERY_PACKET;
416                }
417            }
418
419            let tag_slice = DecoderBufferMut::new(tag_slice);
420            let (tag, _) = tag_slice.decode()?;
421            tag
422        };
423
424        let (_credentials, buffer) = buffer.decode::<Credentials>()?;
425        let (_wire_version, buffer) = buffer.decode::<WireVersion>()?;
426
427        let (_source_control_port, buffer) = buffer.decode::<u16>()?;
428
429        let (stream_id, buffer) = buffer.decode::<stream::Id>()?;
430
431        decoder_invariant!(
432            stream_id.is_reliable,
433            "only reliable streams can be retransmitted"
434        );
435
436        let (_source_queue_id, buffer) = if tag.has_source_queue_id() {
437            let (v, buffer) = buffer.decode::<VarInt>()?;
438            (Some(v), buffer)
439        } else {
440            (None, buffer)
441        };
442
443        let (original_packet_number, buffer) = buffer.decode::<VarInt>()?;
444        let (retransmission_packet_number_buffer, buffer) =
445            buffer.decode_slice(size_of::<RelativeRetransmissionOffset>())?;
446
447        let (_next_expected_control_packet, buffer) = buffer.decode::<VarInt>()?;
448        let (_stream_offset, buffer) = buffer.decode::<VarInt>()?;
449
450        let auth_tag_offset = buffer
451            .len()
452            .checked_sub(key.tag_len())
453            .ok_or(DecoderError::InvariantViolation("missing auth tag"))?;
454        let buffer = buffer.skip(auth_tag_offset)?;
455        let auth_tag = buffer.into_less_safe_slice();
456
457        let relative = retransmission_packet_number
458            .checked_sub(original_packet_number)
459            .ok_or(DecoderError::InvariantViolation(
460                "invalid retransmission packet number",
461            ))?;
462
463        let relative: RelativeRetransmissionOffset = relative
464            .as_u64()
465            .try_into()
466            .map_err(|_| DecoderError::InvariantViolation("packet is too old"))?;
467
468        // undo the previous retransmission if needed
469        let prev_value = RelativeRetransmissionOffset::from_be_bytes(
470            retransmission_packet_number_buffer.peek().decode_exact()?,
471        );
472        if prev_value != 0 {
473            let retransmission_packet_number =
474                original_packet_number + VarInt::from_u32(prev_value);
475            key.retransmission_tag(
476                original_packet_number.as_u64(),
477                retransmission_packet_number.as_u64(),
478                auth_tag,
479            );
480        }
481
482        retransmission_packet_number_buffer
483            .into_less_safe_slice()
484            .copy_from_slice(&relative.to_be_bytes());
485
486        key.retransmission_tag(
487            original_packet_number.as_u64(),
488            retransmission_packet_number.as_u64(),
489            auth_tag,
490        );
491
492        Ok(())
493    }
494
495    #[inline(always)]
496    pub fn decode<V: Validator>(
497        buffer: DecoderBufferMut,
498        mut validator: V,
499        crypto_tag_len: usize,
500    ) -> R<Packet> {
501        let (
502            tag,
503            wire_version,
504            credentials,
505            source_queue_id,
506            stream_id,
507            original_packet_number,
508            packet_number,
509            retransmission_packet_number_offset,
510            next_expected_control_packet,
511            stream_offset,
512            final_offset,
513            header_len,
514            total_header_len,
515            application_header_len,
516            control_data_len,
517            payload_len,
518        ) = {
519            let buffer = buffer.peek();
520
521            unsafe {
522                assume!(
523                    crypto_tag_len >= 16,
524                    "tag len needs to be at least 16 bytes"
525                );
526            }
527
528            let start_len = buffer.len();
529
530            let (tag, buffer) = buffer.decode()?;
531            validator.validate_tag(tag)?;
532
533            let (credentials, buffer) = buffer.decode()?;
534            let (wire_version, buffer) = buffer.decode()?;
535
536            // unused space - was source_control_port when we did port migration but that has
537            // been replaced with `source_queue_id`, which is more flexible
538            let (_source_control_port, buffer) = buffer.decode::<u16>()?;
539
540            let (stream_id, buffer) = buffer.decode::<stream::Id>()?;
541
542            let (source_queue_id, buffer) = if tag.has_source_queue_id() {
543                let (v, buffer) = buffer.decode()?;
544                (Some(v), buffer)
545            } else {
546                (None, buffer)
547            };
548
549            let (original_packet_number, buffer) = buffer.decode::<VarInt>()?;
550
551            let retransmission_packet_number_offset = (start_len - buffer.len()) as u8;
552            let (packet_number, buffer) = if stream_id.is_reliable {
553                let (rel, buffer) = buffer.decode::<RelativeRetransmissionOffset>()?;
554                let rel = VarInt::from_u32(rel);
555                let pn = original_packet_number.checked_add(rel).ok_or(
556                    DecoderError::InvariantViolation("retransmission packet number overflow"),
557                )?;
558                (pn, buffer)
559            } else {
560                (original_packet_number, buffer)
561            };
562
563            let (next_expected_control_packet, buffer) = buffer.decode()?;
564            let (stream_offset, buffer) = buffer.decode()?;
565            let (final_offset, buffer) = if tag.has_final_offset() {
566                let (final_offset, buffer) = buffer.decode()?;
567                (Some(final_offset), buffer)
568            } else {
569                (None, buffer)
570            };
571            let (control_data_len, buffer) = if tag.has_control_data() {
572                buffer.decode()?
573            } else {
574                (VarInt::ZERO, buffer)
575            };
576            let (payload_len, buffer) = buffer.decode::<VarInt>()?;
577
578            let (application_header_len, buffer) = if tag.has_application_header() {
579                let (application_header_len, buffer) = buffer.decode::<VarInt>()?;
580                ((*application_header_len) as usize, buffer)
581            } else {
582                (0, buffer)
583            };
584
585            let header_len = start_len - buffer.len();
586
587            let buffer = buffer.skip(application_header_len)?;
588            let buffer = buffer.skip(*control_data_len as _)?;
589
590            let total_header_len = start_len - buffer.len();
591
592            let buffer = buffer.skip(*payload_len as _)?;
593            let buffer = buffer.skip(crypto_tag_len)?;
594
595            let _ = buffer;
596
597            (
598                tag,
599                wire_version,
600                credentials,
601                source_queue_id,
602                stream_id,
603                original_packet_number,
604                packet_number,
605                retransmission_packet_number_offset,
606                next_expected_control_packet,
607                stream_offset,
608                final_offset,
609                header_len,
610                total_header_len,
611                application_header_len,
612                control_data_len,
613                payload_len,
614            )
615        };
616
617        unsafe {
618            assume!(buffer.len() >= total_header_len);
619        }
620        let (header, buffer) = buffer.decode_slice(total_header_len)?;
621
622        let (application_header, control_data) = {
623            let buffer = header.peek();
624            unsafe {
625                assume!(buffer.len() >= header_len);
626            }
627            let buffer = buffer.skip(header_len)?;
628            unsafe {
629                assume!(buffer.len() >= application_header_len);
630            }
631            let (application_header, buffer) =
632                buffer.skip_into_range(application_header_len, &header)?;
633            unsafe {
634                assume!(buffer.len() >= *control_data_len as usize);
635            }
636            let (control_data, _) = buffer.skip_into_range(*control_data_len as usize, &header)?;
637
638            (application_header, control_data)
639        };
640        let header = header.into_less_safe_slice();
641
642        let (payload, buffer) = buffer.decode_slice(*payload_len as usize)?;
643        let payload = payload.into_less_safe_slice();
644
645        let (auth_tag, buffer) = buffer.decode_slice(crypto_tag_len)?;
646        let auth_tag = auth_tag.into_less_safe_slice();
647
648        let packet = Packet {
649            tag,
650            wire_version,
651            credentials,
652            source_queue_id,
653            stream_id,
654            original_packet_number,
655            packet_number,
656            retransmission_packet_number_offset,
657            next_expected_control_packet,
658            stream_offset,
659            final_offset,
660            header,
661            application_header,
662            control_data,
663            payload,
664            auth_tag,
665        };
666
667        Ok((packet, buffer))
668    }
669}