1use 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 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 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 self.header[0] &= !super::Tag::IS_RECOVERY_PACKET;
312
313 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 before.tag.set_packet_space(space);
344 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 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 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 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}