s2n_quic_dc/packet/control/
decoder.rs1use 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 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 self.buffer = &mut [];
315 Some(Err(err))
316 }
317 }
318 }
319}