Skip to content
File

Blob: firmware/vendor/sctp-proto/src/queue/reassembly_queue.rs

rust642 lines
1use crate::StreamId;
2use crate::chunk::chunk_payload_data::{ChunkPayloadData, PayloadProtocolIdentifier};
3use crate::error::{Error, Result};
4use crate::util::*;
5 
6use alloc::vec::Vec;
7use bytes::{Bytes, BytesMut};
8use core::cmp::Ordering;
9 
10fn sort_chunks_by_tsn(c: &mut [ChunkPayloadData]) {
11 c.sort_by(|a, b| {
12 if sna32lt(a.tsn, b.tsn) {
13 Ordering::Less
14 } else {
15 Ordering::Greater
16 }
17 });
18}
19 
20fn sort_chunks_by_ssn(c: &mut [Chunks]) {
21 c.sort_by(|a, b| {
22 if sna16lt(a.ssn, b.ssn) {
23 Ordering::Less
24 } else {
25 Ordering::Greater
26 }
27 });
28}
29 
30/// A chunk of data from the stream
31#[derive(Debug, PartialEq)]
32pub struct Chunk {
33 /// The contents of the chunk
34 pub bytes: Bytes,
35}
36 
37/// Chunks is a set of chunks that share the same SSN
38#[derive(Default, Debug, Clone)]
39pub struct Chunks {
40 /// used only with the ordered chunks
41 pub(crate) ssn: u16,
42 pub ppi: PayloadProtocolIdentifier,
43 pub chunks: Vec<ChunkPayloadData>,
44 offset: usize,
45 index: usize,
46}
47 
48impl Chunks {
49 pub fn is_empty(&self) -> bool {
50 self.len() == 0
51 }
52 
53 pub fn len(&self) -> usize {
54 let mut l = 0;
55 for c in &self.chunks {
56 l += c.user_data.len();
57 }
58 l
59 }
60 
61 // Concat all fragments into the buffer
62 pub fn read(&self, buf: &mut [u8]) -> Result<usize> {
63 let mut n_written = 0;
64 for c in &self.chunks {
65 let to_copy = c.user_data.len();
66 let n = core::cmp::min(to_copy, buf.len() - n_written);
67 buf[n_written..n_written + n].copy_from_slice(&c.user_data[..n]);
68 n_written += n;
69 if n < to_copy {
70 return Err(Error::ErrShortBuffer);
71 }
72 }
73 Ok(n_written)
74 }
75 
76 pub fn next(&mut self, max_length: usize) -> Option<Chunk> {
77 if self.index >= self.chunks.len() {
78 return None;
79 }
80 
81 let mut buf = BytesMut::with_capacity(max_length);
82 
83 let mut n_written = 0;
84 while self.index < self.chunks.len() {
85 let to_copy = self.chunks[self.index].user_data[self.offset..].len();
86 let n = core::cmp::min(to_copy, max_length - n_written);
87 buf.extend_from_slice(&self.chunks[self.index].user_data[self.offset..self.offset + n]);
88 n_written += n;
89 if n < to_copy {
90 self.offset += n;
91 return Some(Chunk {
92 bytes: buf.freeze(),
93 });
94 }
95 self.index += 1;
96 self.offset = 0;
97 }
98 
99 Some(Chunk {
100 bytes: buf.freeze(),
101 })
102 }
103 
104 pub(crate) fn new(
105 ssn: u16,
106 ppi: PayloadProtocolIdentifier,
107 chunks: Vec<ChunkPayloadData>,
108 ) -> Self {
109 Chunks {
110 ssn,
111 ppi,
112 chunks,
113 offset: 0,
114 index: 0,
115 }
116 }
117 
118 pub(crate) fn push(&mut self, chunk: ChunkPayloadData) -> bool {
119 // check if dup
120 for c in &self.chunks {
121 if c.tsn == chunk.tsn {
122 return false;
123 }
124 }
125 
126 // append and sort
127 self.chunks.push(chunk);
128 sort_chunks_by_tsn(&mut self.chunks);
129 
130 // Check if we now have a complete set
131 self.is_complete()
132 }
133 
134 pub(crate) fn is_complete(&self) -> bool {
135 // Condition for complete set
136 // 0. Has at least one chunk.
137 // 1. Begins with beginningFragment set to true
138 // 2. Ends with endingFragment set to true
139 // 3. TSN monotinically increase by 1 from beginning to end
140 
141 // 0.
142 let n_chunks = self.chunks.len();
143 if n_chunks == 0 {
144 return false;
145 }
146 
147 // 1.
148 if !self.chunks[0].beginning_fragment {
149 return false;
150 }
151 
152 // 2.
153 if !self.chunks[n_chunks - 1].ending_fragment {
154 return false;
155 }
156 
157 // 3.
158 let mut last_tsn = 0u32;
159 for (i, c) in self.chunks.iter().enumerate() {
160 if i > 0 {
161 // Fragments must have contiguous TSN
162 // From RFC 4960 Section 3.3.1:
163 // When a user message is fragmented into multiple chunks, the TSNs are
164 // used by the receiver to reassemble the message. This means that the
165 // TSNs for each fragment of a fragmented user message MUST be strictly
166 // sequential.
167 if c.tsn != last_tsn.wrapping_add(1) {
168 // mid or end fragment is missing
169 return false;
170 }
171 }
172 
173 last_tsn = c.tsn;
174 }
175 
176 true
177 }
178 
179 fn is_missing_tsn_at_or_before(&self, cumulative_tsn: u32) -> bool {
180 let Some(first) = self.chunks.first() else {
181 return false;
182 };
183 if !first.beginning_fragment && sna32lte(first.tsn.wrapping_sub(1), cumulative_tsn) {
184 return true;
185 }
186 if self.chunks.windows(2).any(|pair| {
187 let missing_tsn = pair[0].tsn.wrapping_add(1);
188 missing_tsn != pair[1].tsn && sna32lte(missing_tsn, cumulative_tsn)
189 }) {
190 return true;
191 }
192 
193 let last = self.chunks.last().expect("a first chunk exists");
194 !last.ending_fragment && sna32lte(last.tsn.wrapping_add(1), cumulative_tsn)
195 }
196}
197 
198#[derive(Default, Debug)]
199pub(crate) struct ReassemblyQueue {
200 pub(crate) si: StreamId,
201 pub(crate) next_ssn: u16,
202 /// expected SSN for next ordered chunk
203 pub(crate) ordered: Vec<Chunks>,
204 pub(crate) unordered: Vec<Chunks>,
205 pub(crate) unordered_chunks: Vec<ChunkPayloadData>,
206 pub(crate) n_bytes: usize,
207 pub(crate) max_message_size: u32,
208}
209 
210impl ReassemblyQueue {
211 /// From RFC 4960 Sec 6.5:
212 /// The Stream Sequence Number in all the streams MUST start from 0 when
213 /// the association is Established. Also, when the Stream Sequence
214 /// Number reaches the value 65535 the next Stream Sequence Number MUST
215 /// be set to 0.
216 pub(crate) fn new(si: StreamId, max_message_size: u32) -> Self {
217 ReassemblyQueue {
218 si,
219 next_ssn: 0, // From RFC 4960 Sec 6.5:
220 ordered: vec![],
221 unordered: vec![],
222 unordered_chunks: vec![],
223 n_bytes: 0,
224 max_message_size,
225 }
226 }
227 
228 /// Every retained DATA fragment, complete or still awaiting reassembly.
229 pub(crate) fn chunks(&self) -> impl Iterator<Item = &ChunkPayloadData> {
230 self.ordered
231 .iter()
232 .chain(&self.unordered)
233 .flat_map(|message| &message.chunks)
234 .chain(&self.unordered_chunks)
235 }
236 
237 pub(crate) fn push(&mut self, chunk: ChunkPayloadData) -> Result<bool> {
238 if chunk.stream_identifier != self.si {
239 return Ok(false);
240 }
241 
242 if chunk.unordered {
243 // Check if adding this chunk would exceed limit for unordered
244 let projected_size = self.calculate_unordered_message_size(&chunk);
245 if projected_size > self.max_message_size as usize {
246 return Err(Error::ErrInboundPacketTooLarge);
247 }
248 
249 // First, insert into unordered_chunks array
250 //atomic.AddUint64(&r.n_bytes, uint64(len(chunk.userData)))
251 self.n_bytes += chunk.user_data.len();
252 self.unordered_chunks.push(chunk);
253 sort_chunks_by_tsn(&mut self.unordered_chunks);
254 
255 // Scan unordered_chunks that are contiguous (in TSN)
256 // If found, append the complete set to the unordered array
257 if let Some(cset) = self.find_complete_unordered_chunk_set() {
258 self.unordered.push(cset);
259 return Ok(true);
260 }
261 
262 Ok(false)
263 } else {
264 // Check if adding this chunk would exceed limit for ordered
265 let projected_size = self.calculate_ordered_message_size(&chunk);
266 if projected_size > self.max_message_size as usize {
267 return Err(Error::ErrInboundPacketTooLarge);
268 }
269 
270 // This is an ordered chunk
271 if sna16lt(chunk.stream_sequence_number, self.next_ssn) {
272 return Ok(false);
273 }
274 
275 self.n_bytes += chunk.user_data.len();
276 
277 // Check if a chunkSet with the SSN already exists
278 for s in &mut self.ordered {
279 if s.ssn == chunk.stream_sequence_number {
280 return Ok(s.push(chunk));
281 }
282 }
283 
284 // If not found, create a new chunkSet
285 let mut cset = Chunks::new(chunk.stream_sequence_number, chunk.payload_type, vec![]);
286 let unordered = chunk.unordered;
287 let ok = cset.push(chunk);
288 self.ordered.push(cset);
289 if !unordered {
290 sort_chunks_by_ssn(&mut self.ordered);
291 }
292 
293 Ok(ok)
294 }
295 }
296 
297 fn calculate_ordered_message_size(&self, new_chunk: &ChunkPayloadData) -> usize {
298 let ssn = new_chunk.stream_sequence_number;
299 let existing: usize = self
300 .ordered
301 .iter()
302 .find(|s| s.ssn == ssn)
303 .map(|s| s.len())
304 .unwrap_or(0);
305 existing + new_chunk.user_data.len()
306 }
307 
308 fn calculate_unordered_message_size(&self, new_chunk: &ChunkPayloadData) -> usize {
309 // For unordered, calculate size of contiguous chunk set this belongs to
310 // This is more complex - need to find the message boundary
311 
312 // first find the set of TSNs that preceeds this chunk
313 let prefix = if new_chunk.beginning_fragment {
314 0
315 } else if let Some(mut p) = self
316 .unordered_chunks
317 .iter()
318 .rposition(|f| f.tsn == new_chunk.tsn.wrapping_sub(1))
319 {
320 let mut cnt = 0;
321 let mut tsn = new_chunk.tsn;
322 loop {
323 if self.unordered_chunks[p].tsn == tsn.wrapping_sub(1) {
324 cnt += self.unordered_chunks[p].user_data.len();
325 tsn = self.unordered_chunks[p].tsn;
326 } else {
327 break;
328 }
329 
330 if self.unordered_chunks[p].beginning_fragment || (p == 0) {
331 break;
332 }
333 p -= 1;
334 }
335 cnt
336 } else {
337 0
338 };
339 
340 // next find the set of TSNs that succeeds this chunk
341 let suffix = if new_chunk.ending_fragment {
342 0
343 } else if let Some(mut p) = self
344 .unordered_chunks
345 .iter()
346 .rposition(|f| f.tsn == new_chunk.tsn.wrapping_add(1))
347 {
348 let mut cnt = 0;
349 let mut tsn = new_chunk.tsn;
350 while p < self.unordered_chunks.len() {
351 if self.unordered_chunks[p].tsn == tsn.wrapping_add(1) {
352 cnt += self.unordered_chunks[p].user_data.len();
353 tsn = self.unordered_chunks[p].tsn;
354 } else {
355 break;
356 }
357 
358 if self.unordered_chunks[p].ending_fragment {
359 break;
360 }
361 p += 1;
362 }
363 cnt
364 } else {
365 0
366 };
367 
368 // now sum the lengths together
369 prefix + new_chunk.user_data.len() + suffix
370 }
371 
372 pub(crate) fn find_complete_unordered_chunk_set(&mut self) -> Option<Chunks> {
373 let mut start_idx = -1isize;
374 let mut n_chunks = 0usize;
375 let mut last_tsn = 0u32;
376 let mut found = false;
377 
378 for (i, c) in self.unordered_chunks.iter().enumerate() {
379 // seek beginning
380 if c.beginning_fragment {
381 start_idx = i as isize;
382 n_chunks = 1;
383 last_tsn = c.tsn;
384 
385 if c.ending_fragment {
386 found = true;
387 break;
388 }
389 continue;
390 }
391 
392 if start_idx < 0 {
393 continue;
394 }
395 
396 // Check if contiguous in TSN
397 if c.tsn != last_tsn.wrapping_add(1) {
398 start_idx = -1;
399 continue;
400 }
401 
402 last_tsn = c.tsn;
403 n_chunks += 1;
404 
405 if c.ending_fragment {
406 found = true;
407 break;
408 }
409 }
410 
411 if !found {
412 return None;
413 }
414 
415 // Extract the range of chunks
416 let chunks: Vec<ChunkPayloadData> = self
417 .unordered_chunks
418 .drain(start_idx as usize..(start_idx as usize) + n_chunks)
419 .collect();
420 Some(Chunks::new(0, chunks[0].payload_type, chunks))
421 }
422 
423 pub(crate) fn is_readable(&self) -> bool {
424 // Check unordered first
425 if !self.unordered.is_empty() {
426 // The chunk sets in r.unordered should all be complete.
427 return true;
428 }
429 
430 // Check ordered sets
431 if !self.ordered.is_empty() {
432 let cset = &self.ordered[0];
433 if cset.is_complete() && sna16lte(cset.ssn, self.next_ssn) {
434 return true;
435 }
436 }
437 false
438 }
439 
440 pub(crate) fn read(&mut self) -> Option<Chunks> {
441 // Check unordered first
442 let chunks = if !self.unordered.is_empty() {
443 self.unordered.remove(0)
444 } else if !self.ordered.is_empty() {
445 // Now, check ordered
446 let chunks = &self.ordered[0];
447 if !chunks.is_complete() {
448 return None;
449 }
450 if sna16gt(chunks.ssn, self.next_ssn) {
451 return None;
452 }
453 if chunks.ssn == self.next_ssn {
454 self.next_ssn = self.next_ssn.wrapping_add(1);
455 }
456 self.ordered.remove(0)
457 } else {
458 return None;
459 };
460 
461 self.subtract_num_bytes(chunks.len());
462 
463 Some(chunks)
464 }
465 
466 /// Use last_ssn to locate a chunkSet then remove it if the set has
467 /// not been complete
468 pub(crate) fn forward_tsn_for_ordered(&mut self, last_ssn: u16) {
469 let num_bytes = self
470 .ordered
471 .iter()
472 .filter(|s| sna16lte(s.ssn, last_ssn) && !s.is_complete())
473 .fold(0, |n, s| {
474 n + s.chunks.iter().fold(0, |acc, c| acc + c.user_data.len())
475 });
476 self.subtract_num_bytes(num_bytes);
477 
478 self.ordered
479 .retain(|s| !sna16lte(s.ssn, last_ssn) || s.is_complete());
480 
481 // Finally, forward next_ssn
482 if sna16lte(self.next_ssn, last_ssn) {
483 self.next_ssn = last_ssn.wrapping_add(1);
484 }
485 }
486 
487 /// Apply an ordered Forward-TSN retained across a stream reset without
488 /// crossing into messages whose TSNs are newer than its cumulative point.
489 pub(crate) fn forward_tsn_for_ordered_bounded(
490 &mut self,
491 last_ssn: u16,
492 new_cumulative_tsn: u32,
493 ) {
494 let first_unaffected_ssn = self
495 .ordered
496 .iter()
497 .filter(|chunks| {
498 sna16lte(chunks.ssn, last_ssn)
499 && if chunks.is_complete() {
500 chunks
501 .chunks
502 .iter()
503 .any(|chunk| sna32gt(chunk.tsn, new_cumulative_tsn))
504 } else {
505 !chunks.is_missing_tsn_at_or_before(new_cumulative_tsn)
506 }
507 })
508 .map(|chunks| chunks.ssn)
509 .reduce(|first, candidate| {
510 if sna16lt(candidate, first) {
511 candidate
512 } else {
513 first
514 }
515 });
516 
517 let is_abandoned_partial = |chunks: &&Chunks| {
518 sna16lte(chunks.ssn, last_ssn)
519 && !chunks.is_complete()
520 && chunks.is_missing_tsn_at_or_before(new_cumulative_tsn)
521 };
522 let num_bytes = self
523 .ordered
524 .iter()
525 .filter(is_abandoned_partial)
526 .flat_map(|chunks| chunks.chunks.iter())
527 .map(|chunk| chunk.user_data.len())
528 .sum();
529 self.subtract_num_bytes(num_bytes);
530 self.ordered.retain(|chunks| {
531 !sna16lte(chunks.ssn, last_ssn)
532 || chunks.is_complete()
533 || !chunks.is_missing_tsn_at_or_before(new_cumulative_tsn)
534 });
535 
536 let next_ssn = first_unaffected_ssn.unwrap_or_else(|| last_ssn.wrapping_add(1));
537 if sna16lt(self.next_ssn, next_ssn) {
538 self.next_ssn = next_ssn;
539 }
540 }
541 
542 /// The highest prefix of an ambiguous deferred ordered skip that has
543 /// enough TSN context to advance without overtaking a missing successor.
544 pub(crate) fn applicable_forward_tsn_for_ordered_bounded(
545 &self,
546 last_ssn: u16,
547 new_cumulative_tsn: u32,
548 peer_last_tsn: u32,
549 ) -> Option<u16> {
550 if sna16gt(self.next_ssn, last_ssn) {
551 return Some(last_ssn);
552 }
553 
554 let is_covered = |chunks: &Chunks| {
555 if chunks.is_complete() {
556 chunks
557 .chunks
558 .iter()
559 .all(|chunk| sna32lte(chunk.tsn, new_cumulative_tsn))
560 } else {
561 chunks.is_missing_tsn_at_or_before(new_cumulative_tsn)
562 }
563 };
564 if self.ordered.iter().any(|chunks| {
565 chunks.ssn == last_ssn && (chunks.ssn == self.next_ssn || is_covered(chunks))
566 }) {
567 return Some(last_ssn);
568 }
569 
570 let has_resolved_later_tsn = self
571 .ordered
572 .iter()
573 .filter(|chunks| chunks.ssn == last_ssn || sna16gt(chunks.ssn, last_ssn))
574 .flat_map(|chunks| chunks.chunks.iter())
575 .map(|chunk| chunk.tsn)
576 .reduce(|first, candidate| {
577 if sna32lt(candidate, first) {
578 candidate
579 } else {
580 first
581 }
582 })
583 .is_some_and(|first_tsn| sna32lte(first_tsn, peer_last_tsn));
584 if has_resolved_later_tsn {
585 return Some(last_ssn);
586 }
587 
588 self.ordered
589 .iter()
590 .filter(|chunks| {
591 (chunks.ssn == self.next_ssn || sna16gt(chunks.ssn, self.next_ssn))
592 && sna16lt(chunks.ssn, last_ssn)
593 && (is_covered(chunks)
594 || chunks
595 .chunks
596 .first()
597 .is_some_and(|chunk| sna32lte(chunk.tsn, peer_last_tsn)))
598 })
599 .map(|chunks| chunks.ssn)
600 .reduce(|last, candidate| {
601 if sna16gt(candidate, last) {
602 candidate
603 } else {
604 last
605 }
606 })
607 }
608 
609 /// Remove all fragments in the unordered sets that contains chunks
610 /// equal to or older than `new_cumulative_tsn`.
611 /// We know all sets in the r.unordered are complete ones.
612 /// Just remove chunks that are equal to or older than new_cumulative_tsn
613 /// from the unordered_chunks
614 pub(crate) fn forward_tsn_for_unordered(&mut self, new_cumulative_tsn: u32) {
615 let mut last_idx: isize = -1;
616 for (i, c) in self.unordered_chunks.iter().enumerate() {
617 if sna32gt(c.tsn, new_cumulative_tsn) {
618 break;
619 }
620 last_idx = i as isize;
621 }
622 if last_idx >= 0 {
623 for i in 0..(last_idx + 1) as usize {
624 self.subtract_num_bytes(self.unordered_chunks[i].user_data.len());
625 }
626 self.unordered_chunks.drain(..(last_idx + 1) as usize);
627 }
628 }
629 
630 pub(crate) fn subtract_num_bytes(&mut self, n_bytes: usize) {
631 if self.n_bytes >= n_bytes {
632 self.n_bytes -= n_bytes;
633 } else {
634 self.n_bytes = 0;
635 }
636 }
637 
638 pub(crate) fn get_num_bytes(&self) -> usize {
639 self.n_bytes
640 }
641}