use crate::api::Message;
use crate::api::StreamId;
use crate::api::handover::HandoverOrderedStream;
use crate::api::handover::HandoverReadiness;
use crate::api::handover::HandoverUnorderedStream;
use crate::api::handover::SocketHandoverState;
use crate::packet::SkippedStream;
use crate::packet::data::Data;
use crate::rx::IntervalList;
use crate::rx::ReassemblyKey;
use crate::rx::reassembly_streams::ReassemblyStreams;
use crate::types::Fsn;
use crate::types::Mid;
use crate::types::SerialNumber;
use crate::types::StreamKey;
use crate::types::Tsn;
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct InterleavedKey {
pub mid: Mid, pub fsn: Fsn, }
impl PartialOrd for InterleavedKey {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for InterleavedKey {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
if self.mid == other.mid {
if self.fsn == other.fsn {
std::cmp::Ordering::Equal
} else if other.fsn.greater_than(self.fsn) {
std::cmp::Ordering::Less
} else {
std::cmp::Ordering::Greater
}
} else if other.mid.greater_than(self.mid) {
std::cmp::Ordering::Less
} else {
std::cmp::Ordering::Greater
}
}
}
impl ReassemblyKey for InterleavedKey {
fn next(&self) -> Self {
InterleavedKey { mid: self.mid, fsn: self.fsn + 1 }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UnorderedInterleavedKey {
pub mid: Mid,
pub fsn: Fsn,
}
impl PartialOrd for UnorderedInterleavedKey {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for UnorderedInterleavedKey {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.mid.0.cmp(&other.mid.0).then_with(|| self.fsn.0.cmp(&other.fsn.0))
}
}
impl ReassemblyKey for UnorderedInterleavedKey {
fn next(&self) -> Self {
UnorderedInterleavedKey { mid: self.mid, fsn: self.fsn + 1 }
}
}
pub struct OrderedStream {
stream_id: StreamId,
intervals: IntervalList<InterleavedKey>,
next_mid: Mid,
}
impl OrderedStream {
fn new(stream_id: StreamId, next_mid: Mid) -> Self {
Self { stream_id, intervals: IntervalList::default(), next_mid }
}
fn try_assemble_next(&mut self, on_reassembled: &mut dyn FnMut(Message)) -> usize {
let mut assembled_bytes = 0;
while let Some(interval) =
self.intervals.pop_front_if_complete_and(|i| i.start.mid == self.next_mid)
{
let stream_id = self.stream_id;
let ppid = interval.ppid;
let payload = interval.collect_payload();
assembled_bytes += payload.len();
on_reassembled(Message::new(stream_id, ppid, payload));
self.next_mid += 1;
}
assembled_bytes
}
fn add(&mut self, data: Data, on_reassembled: &mut dyn FnMut(Message)) -> isize {
if data.mid != self.next_mid && !data.mid.greater_than(self.next_mid) {
return 0;
}
if data.mid == self.next_mid && data.is_beginning && data.is_end {
on_reassembled(Message::new(self.stream_id, data.ppid, data.payload));
self.next_mid += 1;
let assembled = self.try_assemble_next(on_reassembled);
return -(assembled as isize);
}
let key = InterleavedKey { mid: data.mid, fsn: data.fsn };
let queued_bytes = data.payload.len() as isize;
self.intervals.add(key, data);
let assembled = self.try_assemble_next(on_reassembled);
queued_bytes - (assembled as isize)
}
}
pub struct UnorderedStream {
stream_id: StreamId,
intervals: IntervalList<UnorderedInterleavedKey>,
}
impl UnorderedStream {
fn new(stream_id: StreamId) -> Self {
Self { stream_id, intervals: IntervalList::default() }
}
fn add(&mut self, data: Data, on_reassembled: &mut dyn FnMut(Message)) -> isize {
if data.is_beginning && data.is_end {
on_reassembled(Message::new(data.stream_key.id(), data.ppid, data.payload));
return 0;
}
let key = UnorderedInterleavedKey { mid: data.mid, fsn: data.fsn };
let queued_bytes = data.payload.len() as isize;
let idx = self.intervals.add(key, data);
if let Some(interval) = self.intervals.pop_if_complete(idx) {
let stream_id = self.stream_id;
let ppid = interval.ppid;
let payload = interval.collect_payload();
let total_payload_len = payload.len();
on_reassembled(Message::new(stream_id, ppid, payload));
queued_bytes - (total_payload_len as isize)
} else {
queued_bytes
}
}
}
pub struct InterleavedReassemblyStreams {
ordered: HashMap<StreamId, OrderedStream>,
unordered: HashMap<StreamId, UnorderedStream>,
queued_bytes: usize,
}
impl InterleavedReassemblyStreams {
pub fn new() -> Self {
Self { ordered: HashMap::new(), unordered: HashMap::new(), queued_bytes: 0 }
}
}
impl ReassemblyStreams for InterleavedReassemblyStreams {
fn add(&mut self, _tsn: Tsn, data: Data, on_reassembled: &mut dyn FnMut(Message)) {
let diff = match data.stream_key {
StreamKey::Ordered(stream_id) => self
.ordered
.entry(stream_id)
.or_insert_with(|| OrderedStream::new(stream_id, Mid(0)))
.add(data, on_reassembled),
StreamKey::Unordered(stream_id) => self
.unordered
.entry(stream_id)
.or_insert_with(|| UnorderedStream::new(stream_id))
.add(data, on_reassembled),
};
self.queued_bytes = self.queued_bytes.strict_add_signed(diff);
}
fn handle_forward_tsn(
&mut self,
_new_cumulative_ack: Tsn,
skipped_streams: &[SkippedStream],
on_reassembled: &mut dyn FnMut(Message),
) {
for skipped_stream in skipped_streams {
if let SkippedStream::IForwardTsn(stream_key, mid) = skipped_stream {
match stream_key {
StreamKey::Ordered(stream_id) => {
let stream = self
.ordered
.entry(*stream_id)
.or_insert_with(|| OrderedStream::new(*stream_id, Mid(0)));
self.queued_bytes -= stream
.intervals
.retain(|interval| interval.start.mid.greater_than(*mid));
if stream.next_mid.less_than_or_equal(*mid) {
stream.next_mid = *mid + 1;
}
self.queued_bytes -= stream.try_assemble_next(on_reassembled);
}
StreamKey::Unordered(stream_id) => {
let stream = self
.unordered
.entry(*stream_id)
.or_insert_with(|| UnorderedStream::new(*stream_id));
self.queued_bytes -= stream
.intervals
.retain(|interval| interval.start.mid.greater_than(*mid));
}
}
}
}
}
fn reset_streams(&mut self, streams: &[StreamId]) {
if streams.is_empty() {
for stream in self.ordered.values_mut() {
self.queued_bytes -= stream.intervals.total_bytes();
stream.next_mid = Mid(0);
stream.intervals = IntervalList::default();
}
} else {
for stream_id in streams {
if let Some(stream) = self.ordered.get_mut(stream_id) {
self.queued_bytes -= stream.intervals.total_bytes();
stream.next_mid = Mid(0);
stream.intervals = IntervalList::default();
}
}
}
}
fn queued_bytes(&self) -> usize {
self.queued_bytes
}
fn get_handover_readiness(&self) -> HandoverReadiness {
let has_ordered_chunks = self.ordered.values().any(|s| !s.intervals.is_empty());
let has_unordered_chunks = self.unordered.values().any(|s| !s.intervals.is_empty());
HandoverReadiness::STREAM_HAS_UNASSEMBLED_CHUNKS
& (has_ordered_chunks | has_unordered_chunks)
}
fn add_to_handover_state(&self, state: &mut SocketHandoverState) {
for (stream_id, stream) in &self.ordered {
state
.rx
.ordered_streams
.push(HandoverOrderedStream { id: stream_id.0, next_ssn: stream.next_mid.0 });
}
for stream_id in self.unordered.keys() {
state.rx.unordered_streams.push(HandoverUnorderedStream { id: stream_id.0 });
}
}
fn restore_from_state(&mut self, state: &SocketHandoverState) {
for stream in &state.rx.ordered_streams {
self.ordered.insert(
StreamId(stream.id),
OrderedStream::new(StreamId(stream.id), Mid(stream.next_ssn)),
);
}
for stream in &state.rx.unordered_streams {
self.unordered.insert(StreamId(stream.id), UnorderedStream::new(StreamId(stream.id)));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::testing::data_sequencer::DataSequencer;
#[test]
fn add_unordered_message_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
s.add(Tsn(1), seq.unordered("a", "B"), &mut |_| {});
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(2), seq.unordered("bcd", ""), &mut |_| {});
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(3), seq.unordered("ef", ""), &mut |_| {});
assert_eq!(s.queued_bytes(), 6);
s.add(Tsn(4), seq.unordered("g", "E"), &mut |_| {});
assert_eq!(s.queued_bytes(), 0);
}
#[test]
fn add_unordered_message_out_of_order_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c1 = seq.unordered("a", "B");
let c2 = seq.unordered("bcd", "");
let c3 = seq.unordered("ef", "");
let c4 = seq.unordered("g", "E");
s.add(Tsn(1), c1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(2), c2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(4), c4, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 5);
assert!(messages.is_empty());
s.add(Tsn(3), c3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
}
#[test]
fn add_simple_ordered_message_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c1 = seq.ordered("a", "B");
let c2 = seq.ordered("bcd", "");
let c3 = seq.ordered("ef", "");
let c4 = seq.ordered("g", "E");
s.add(Tsn(1), c1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(2), c2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(3), c3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 6);
s.add(Tsn(4), c4, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
}
#[test]
fn add_more_complex_ordered_message_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c11 = seq.ordered("a", "B");
let c12 = seq.ordered("bcd", "");
let c13 = seq.ordered("ef", "");
let c14 = seq.ordered("g", "E");
let c21 = seq.ordered("h", "BE");
let c31 = seq.ordered("ij", "B");
let c32 = seq.ordered("k", "E");
s.add(Tsn(1), c11, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(3), c13, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 3);
s.add(Tsn(4), c14, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(5), c21, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 5);
s.add(Tsn(6), c31, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 7);
s.add(Tsn(7), c32, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 8);
assert!(messages.is_empty());
s.add(Tsn(2), c12, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 3);
}
#[test]
fn delete_unordered_message_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c1 = seq.unordered("a", "B");
let c2 = seq.unordered("bcd", "");
let c3 = seq.unordered("ef", "");
s.add(Tsn(1), c1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(2), c2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(3), c3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 6);
s.handle_forward_tsn(
Tsn(3),
&[SkippedStream::IForwardTsn(StreamKey::Unordered(StreamId(1)), Mid(0))],
&mut |m| messages.push(m),
);
assert_eq!(s.queued_bytes(), 0);
}
#[test]
fn delete_simple_ordered_message_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c1 = seq.ordered("a", "B");
let c2 = seq.ordered("bcd", "");
let c3 = seq.ordered("ef", "");
s.add(Tsn(1), c1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(2), c2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(3), c3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 6);
s.handle_forward_tsn(
Tsn(3),
&[SkippedStream::IForwardTsn(StreamKey::Ordered(StreamId(1)), Mid(0))],
&mut |m| messages.push(m),
);
assert_eq!(s.queued_bytes(), 0);
}
#[test]
fn delete_many_ordered_messages_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c1 = seq.ordered("a", "B");
seq.ordered("bcd", ""); let c3 = seq.ordered("ef", "");
let c4 = seq.ordered("g", "E");
let c5 = seq.ordered("h", "BE");
let c6 = seq.ordered("ij", "B");
let c7 = seq.ordered("k", "E");
s.add(Tsn(1), c1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(3), c3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 3);
s.add(Tsn(4), c4, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(5), c5, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 5);
s.add(Tsn(6), c6, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 7);
s.add(Tsn(7), c7, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 8);
s.handle_forward_tsn(
Tsn(8),
&[SkippedStream::IForwardTsn(StreamKey::Ordered(StreamId(1)), Mid(2))],
&mut |m| messages.push(m),
);
assert_eq!(s.queued_bytes(), 0);
}
#[test]
fn delete_ordered_message_delives_two_returns_correct_size() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let c1 = seq.ordered("a", "B");
seq.ordered("bcd", ""); let c3 = seq.ordered("ef", "");
let c4 = seq.ordered("g", "E");
let c5 = seq.ordered("h", "BE");
let c6 = seq.ordered("ij", "B");
let c7 = seq.ordered("k", "E");
s.add(Tsn(1), c1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
s.add(Tsn(3), c3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 3);
s.add(Tsn(4), c4, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 4);
s.add(Tsn(5), c5, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 5);
s.add(Tsn(6), c6, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 7);
s.add(Tsn(7), c7, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 8);
s.handle_forward_tsn(
Tsn(8),
&[SkippedStream::IForwardTsn(StreamKey::Ordered(StreamId(1)), Mid(0))],
&mut |m| messages.push(m),
);
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 2);
}
#[test]
fn can_delete_first_ordered_message() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
seq.ordered("abc", "BE"); let c2 = seq.ordered("def", "BE");
s.handle_forward_tsn(
Tsn(1),
&[SkippedStream::IForwardTsn(StreamKey::Ordered(StreamId(1)), Mid(0))],
&mut |m| messages.push(m),
);
assert_eq!(s.queued_bytes(), 0);
s.add(Tsn(2), c2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
}
#[test]
fn can_reassemble_fast_path_unordered() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let data1 = seq.unordered("a", "BE");
let data2 = seq.unordered("b", "BE");
let data3 = seq.unordered("c", "BE");
let data4 = seq.unordered("d", "BE");
s.add(Tsn(1), data1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
s.add(Tsn(3), data3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 2);
s.add(Tsn(2), data2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 3);
s.add(Tsn(4), data4, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 4);
}
#[test]
fn can_reassemble_fast_path_ordered() {
let mut s = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut messages = Vec::new();
let data1 = seq.ordered("a", "BE");
let data2 = seq.ordered("b", "BE");
let data3 = seq.ordered("c", "BE");
let data4 = seq.ordered("d", "BE");
s.add(Tsn(1), data1, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
s.add(Tsn(3), data3, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 1);
assert_eq!(messages.len(), 1);
s.add(Tsn(2), data2, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 3);
s.add(Tsn(4), data4, &mut |m| messages.push(m));
assert_eq!(s.queued_bytes(), 0);
assert_eq!(messages.len(), 4);
}
#[test]
fn can_handover_ordered_streams() {
let mut streams1 = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
streams1.add(Tsn(1), seq.ordered("a", "B"), &mut |_| {});
assert_eq!(streams1.queued_bytes(), 1);
assert!(
streams1
.get_handover_readiness()
.contains(HandoverReadiness::STREAM_HAS_UNASSEMBLED_CHUNKS)
);
streams1.add(Tsn(2), seq.ordered("bcd", "E"), &mut |_| {});
assert_eq!(streams1.queued_bytes(), 0);
assert!(streams1.get_handover_readiness().is_ready());
let mut state = SocketHandoverState::default();
streams1.add_to_handover_state(&mut state);
let mut streams2 = InterleavedReassemblyStreams::new();
let mut messages = Vec::new();
streams2.restore_from_state(&state);
let data = seq.ordered("efgh", "BE");
assert_eq!(data.mid, Mid(1));
streams2.add(Tsn(3), data, &mut |m| messages.push(m));
assert_eq!(streams2.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].payload, b"efgh");
}
#[test]
fn can_handover_unordered_streams() {
let mut streams1 = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
streams1.add(Tsn(1), seq.unordered("a", "B"), &mut |_| {});
assert_eq!(streams1.queued_bytes(), 1);
assert!(
streams1
.get_handover_readiness()
.contains(HandoverReadiness::STREAM_HAS_UNASSEMBLED_CHUNKS)
);
streams1.add(Tsn(2), seq.unordered("bcd", "E"), &mut |_| {});
assert_eq!(streams1.queued_bytes(), 0);
assert!(streams1.get_handover_readiness().is_ready());
let mut state = SocketHandoverState::default();
streams1.add_to_handover_state(&mut state);
let mut streams2 = InterleavedReassemblyStreams::new();
let mut messages = Vec::new();
streams2.restore_from_state(&state);
let data = seq.unordered("efgh", "BE");
streams2.add(Tsn(3), data, &mut |m| messages.push(m));
assert_eq!(streams2.queued_bytes(), 0);
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].payload, b"efgh");
}
#[test]
fn unordered_interleaved_streams_support_extreme_mid_distances() {
let mut streams = InterleavedReassemblyStreams::new();
let mut seq = DataSequencer::new(StreamId(1));
let mut data1 = seq.unordered("a", "B");
data1.mid = Mid(0);
streams.add(Tsn(1), data1, &mut |_| {});
let mut data2 = seq.unordered("b", "B");
data2.mid = Mid(1 << 31);
streams.add(Tsn(2), data2, &mut |_| {});
assert_eq!(streams.queued_bytes(), 2);
}
}