Skip to main content

s2n_quic_dc/packet/control/
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    packet::{control::Tag, stream, WireVersion},
7};
8use core::fmt;
9use s2n_codec::{
10    decoder_invariant, CheckedRange, DecoderBufferMut, DecoderBufferMutResult as R, DecoderError,
11};
12use s2n_quic_core::{assume, frame::FrameMut, varint::VarInt};
13
14type PacketNumber = VarInt;
15
16pub trait Validator {
17    fn validate_tag(&mut self, tag: Tag) -> Result<(), DecoderError>;
18}
19
20impl Validator for () {
21    #[inline]
22    fn validate_tag(&mut self, _tag: Tag) -> Result<(), DecoderError> {
23        Ok(())
24    }
25}
26
27impl Validator for Tag {
28    #[inline]
29    fn validate_tag(&mut self, actual: Tag) -> Result<(), DecoderError> {
30        decoder_invariant!(*self == actual, "unexpected packet type");
31        Ok(())
32    }
33}
34
35impl<A, B> Validator for (A, B)
36where
37    A: Validator,
38    B: Validator,
39{
40    #[inline]
41    fn validate_tag(&mut self, tag: Tag) -> Result<(), DecoderError> {
42        self.0.validate_tag(tag)?;
43        self.1.validate_tag(tag)?;
44        Ok(())
45    }
46}
47
48pub struct Packet<'a> {
49    tag: Tag,
50    wire_version: WireVersion,
51    credentials: Credentials,
52    source_queue_id: Option<VarInt>,
53    stream_id: Option<stream::Id>,
54    packet_number: PacketNumber,
55    header: &'a mut [u8],
56    application_header: CheckedRange,
57    control_data: CheckedRange,
58    auth_tag: &'a mut [u8],
59}
60
61impl fmt::Debug for Packet<'_> {
62    fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
63        let header = &*self.header;
64
65        let mut s = f.debug_struct("control::Packet");
66
67        s.field("tag", &self.tag)
68            .field("wire_version", &self.wire_version)
69            .field("credentials", &self.credentials)
70            .field("source_queue_id", &self.source_queue_id)
71            .field("stream_id", &self.stream_id)
72            .field("packet_number", &self.packet_number);
73
74        if !self.application_header.is_empty() {
75            s.field("application_header", &self.application_header.get(header));
76        }
77
78        if !self.control_data.is_empty() {
79            s.field("control_data", &self.control_data.get(header));
80        }
81
82        s.field("auth_tag", &self.auth_tag).finish()
83    }
84}
85
86impl Packet<'_> {
87    #[inline]
88    pub fn tag(&self) -> Tag {
89        self.tag
90    }
91
92    #[inline]
93    pub fn wire_version(&self) -> WireVersion {
94        self.wire_version
95    }
96
97    #[inline]
98    pub fn credentials(&self) -> &Credentials {
99        &self.credentials
100    }
101
102    #[inline]
103    pub fn source_queue_id(&self) -> Option<VarInt> {
104        self.source_queue_id
105    }
106
107    #[inline]
108    pub fn stream_id(&self) -> Option<&stream::Id> {
109        self.stream_id.as_ref()
110    }
111
112    #[inline]
113    pub fn packet_number(&self) -> PacketNumber {
114        self.packet_number
115    }
116
117    #[inline]
118    pub fn application_header(&self) -> &[u8] {
119        self.application_header.get(self.header)
120    }
121
122    #[inline]
123    pub fn control_data(&self) -> &[u8] {
124        self.control_data.get(self.header)
125    }
126
127    #[inline]
128    pub fn control_data_mut(&mut self) -> &mut [u8] {
129        self.control_data.get_mut(self.header)
130    }
131
132    #[inline]
133    pub fn control_frames_mut(&mut self) -> ControlFramesMut<'_> {
134        ControlFramesMut {
135            buffer: self.control_data.get_mut(self.header),
136        }
137    }
138
139    #[inline]
140    pub fn header(&self) -> &[u8] {
141        self.header
142    }
143
144    #[inline]
145    pub fn auth_tag(&self) -> &[u8] {
146        self.auth_tag
147    }
148
149    #[inline]
150    pub fn total_len(&self) -> usize {
151        self.header.len() + self.auth_tag.len()
152    }
153
154    #[inline(always)]
155    pub fn decode<V: Validator>(
156        buffer: DecoderBufferMut,
157        mut validator: V,
158        crypto_tag_len: usize,
159    ) -> R<Packet> {
160        let (
161            tag,
162            wire_version,
163            credentials,
164            source_queue_id,
165            stream_id,
166            packet_number,
167            header_len,
168            total_header_len,
169            application_header_len,
170            control_data_len,
171        ) = {
172            let buffer = buffer.peek();
173
174            unsafe {
175                assume!(
176                    crypto_tag_len >= 16,
177                    "tag len needs to be at least 16 bytes"
178                );
179            }
180
181            let start_len = buffer.len();
182
183            let (tag, buffer) = buffer.decode()?;
184            validator.validate_tag(tag)?;
185
186            let (credentials, buffer) = buffer.decode()?;
187            let (wire_version, buffer) = buffer.decode()?;
188
189            let (stream_id, buffer) = if tag.is_stream() {
190                let (stream_id, buffer) = buffer.decode()?;
191                (Some(stream_id), buffer)
192            } else {
193                (None, buffer)
194            };
195
196            let (source_queue_id, buffer) = if tag.has_source_queue_id() {
197                let (v, buffer) = buffer.decode()?;
198                (Some(v), buffer)
199            } else {
200                (None, buffer)
201            };
202
203            let (packet_number, buffer) = buffer.decode::<VarInt>()?;
204            let (control_data_len, buffer) = buffer.decode::<VarInt>()?;
205
206            let (application_header_len, buffer) = if tag.has_application_header() {
207                let (application_header_len, buffer) = buffer.decode::<VarInt>()?;
208                ((*application_header_len) as usize, buffer)
209            } else {
210                (0, buffer)
211            };
212
213            let header_len = start_len - buffer.len();
214
215            let buffer = buffer.skip(application_header_len)?;
216            let buffer = buffer.skip(*control_data_len as _)?;
217
218            let total_header_len = start_len - buffer.len();
219
220            let buffer = buffer.skip(crypto_tag_len)?;
221
222            let _ = buffer;
223
224            (
225                tag,
226                wire_version,
227                credentials,
228                source_queue_id,
229                stream_id,
230                packet_number,
231                header_len,
232                total_header_len,
233                application_header_len,
234                control_data_len,
235            )
236        };
237
238        unsafe {
239            assume!(buffer.len() >= total_header_len);
240        }
241        let (header, buffer) = buffer.decode_slice(total_header_len)?;
242
243        let (application_header, control_data) = {
244            let buffer = header.peek();
245            unsafe {
246                assume!(buffer.len() >= header_len);
247            }
248            let buffer = buffer.skip(header_len)?;
249            unsafe {
250                assume!(buffer.len() >= application_header_len);
251            }
252            let (application_header, buffer) =
253                buffer.skip_into_range(application_header_len, &header)?;
254            unsafe {
255                assume!(buffer.len() >= *control_data_len as usize);
256            }
257            let (control_data, _) = buffer.skip_into_range(*control_data_len as usize, &header)?;
258
259            (application_header, control_data)
260        };
261        let header = header.into_less_safe_slice();
262
263        let (auth_tag, buffer) = buffer.decode_slice(crypto_tag_len)?;
264        let auth_tag = auth_tag.into_less_safe_slice();
265
266        let packet = Packet {
267            tag,
268            wire_version,
269            credentials,
270            source_queue_id,
271            stream_id,
272            packet_number,
273            header,
274            application_header,
275            control_data,
276            auth_tag,
277        };
278
279        Ok((packet, buffer))
280    }
281}
282
283pub struct ControlFramesMut<'a> {
284    buffer: &'a mut [u8],
285}
286
287impl<'a> ControlFramesMut<'a> {
288    #[inline]
289    pub(crate) fn new(buffer: &'a mut [u8]) -> Self {
290        Self { buffer }
291    }
292}
293
294impl<'a> Iterator for ControlFramesMut<'a> {
295    type Item = Result<FrameMut<'a>, s2n_codec::DecoderError>;
296
297    #[inline]
298    fn next(&mut self) -> Option<Self::Item> {
299        if self.buffer.is_empty() {
300            return None;
301        }
302
303        let buffer = unsafe {
304            // extend the lifetime of the buffer
305            core::mem::transmute::<&mut [u8], &mut [u8]>(self.buffer)
306        };
307        match DecoderBufferMut::new(buffer).decode::<FrameMut>() {
308            Ok((frame, remaining)) => {
309                self.buffer = remaining.into_less_safe_slice();
310                Some(Ok(frame))
311            }
312            Err(err) => {
313                // clear out the buffer and return an error
314                self.buffer = &mut [];
315                Some(Err(err))
316            }
317        }
318    }
319}