pub(crate) const MAX_RELIABLE_SEND_ATTEMPTS: u8 = 8;
#[derive(Debug, Clone, Copy)]
struct BufferedPacket<const S: usize> {
packet: DataPacket<S>,
connection_id: u64,
try_count: u8,
reliable: bool,
}
impl<const S: usize> BufferedPacket<S> {
fn new(packet: DataPacket<S>, connection_id: u64, reliable: bool) -> Self {
Self {
packet,
connection_id,
try_count: 0,
reliable,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct PacketBuffer<const N: usize, const S: usize> {
buffer: [Option<BufferedPacket<S>>; N],
}
impl<const N: usize, const S: usize> PacketBuffer<N, S> {
pub fn new() -> Self {
Self { buffer: [None; N] }
}
pub fn add(&mut self, packet: DataPacket<S>, connection_id: u64, reliable: bool) {
for i in 0..N {
if self.buffer[i].is_none() {
self.buffer[i] = Some(BufferedPacket::new(packet, connection_id, reliable));
return;
}
}
for i in 0..N {
if let Some(p) = self.buffer[i] {
if !p.reliable {
self.buffer[i] = Some(BufferedPacket::new(packet, connection_id, reliable));
return;
}
}
}
if reliable {
if let Some(p) = self.buffer.iter_mut().max_by_key(|packet| {
packet.as_ref().map_or((std::cmp::Reverse(0), 0), |p| {
(std::cmp::Reverse(p.packet.connection_status.sequence), p.try_count)
})
}) {
*p = Some(BufferedPacket::new(packet, connection_id, reliable));
}
} else {
if let Some(p) = self
.buffer
.iter_mut()
.filter(|packet| packet.as_ref().is_some_and(|p| !p.reliable))
.max_by_key(|packet| packet.as_ref().map_or(0, |p| p.try_count))
{
*p = Some(BufferedPacket::new(packet, connection_id, reliable));
}
}
}
pub fn remove(&mut self, sequence: u16) {
for i in 0..N {
if let Some(packet) = self.buffer[i] {
if packet.packet.connection_status.sequence == sequence {
self.buffer[i] = None;
break;
}
}
}
}
pub fn acknowledge_packets(&mut self, ack: u16, ack_bitfield: u32) {
for bit in 0..u32::BITS {
if (ack_bitfield >> bit) & 1 == 1 {
self.remove(ack.wrapping_sub(bit as u16));
}
}
}
pub fn gather_unsent_packets_for_retry(&mut self) -> ArrayVec<[DataPacket<S>; N]> {
let mut packets = ArrayVec::new();
for slot in &mut self.buffer {
let Some(mut buffered) = *slot else {
continue;
};
let _ = buffered.connection_id;
buffered.try_count += 1;
packets.push(buffered.packet);
*slot = if buffered.reliable && buffered.try_count < MAX_RELIABLE_SEND_ATTEMPTS {
Some(buffered)
} else {
None
};
}
packets
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::packets::ConnectionStatus;
fn make_packet<const S: usize>(sequence: u16, fill: u8) -> DataPacket<S> {
DataPacket::new(1, ConnectionStatus::new(sequence, 0, 0), [fill; S])
}
#[test]
fn test_new_buffer_is_empty() {
let buffer = PacketBuffer::<4, 8>::new();
assert!(buffer.buffer.iter().all(|packet| packet.is_none()));
}
#[test]
fn test_add_packets() {
let mut buffer = PacketBuffer::<4, 16>::new();
assert_eq!(buffer.gather_unsent_packets_for_retry().len(), 0);
buffer.add(DataPacket::new(1, ConnectionStatus::new(0, 0, 0), [0; 16]), 1, false);
buffer.add(DataPacket::new(1, ConnectionStatus::new(1, 0, 0), [0; 16]), 1, false);
buffer.add(DataPacket::new(1, ConnectionStatus::new(2, 0, 0), [0; 16]), 1, false);
buffer.add(DataPacket::new(1, ConnectionStatus::new(3, 0, 0), [0; 16]), 1, false);
assert_eq!(buffer.buffer.iter().filter(|p| p.is_some()).count(), 4);
}
#[test]
fn test_add_fills_first_empty_slot() {
let mut buffer = PacketBuffer::<4, 8>::new();
buffer.add(make_packet::<8>(10, 1), 7, true);
let entry = buffer.buffer[0].expect("expected a packet in the first slot");
assert_eq!(entry.packet.connection_status.sequence, 10);
assert_eq!(entry.connection_id, 7);
assert!(entry.reliable);
assert_eq!(entry.try_count, 0);
}
#[test]
fn test_add_replaces_first_unreliable_even_for_reliable_packet() {
let mut buffer = PacketBuffer::<3, 8>::new();
buffer.add(make_packet::<8>(1, 1), 1, true);
buffer.add(make_packet::<8>(2, 2), 1, false);
buffer.add(make_packet::<8>(3, 3), 1, true);
buffer.add(make_packet::<8>(9, 9), 1, true);
assert_eq!(buffer.buffer[1].unwrap().packet.connection_status.sequence, 9);
}
#[test]
fn test_add_unreliable_dropped_when_buffer_full_of_reliable() {
let mut buffer = PacketBuffer::<3, 8>::new();
buffer.add(make_packet::<8>(1, 1), 1, true);
buffer.add(make_packet::<8>(2, 2), 1, true);
buffer.add(make_packet::<8>(3, 3), 1, true);
let sequences_before: [u16; 3] = [
buffer.buffer[0].unwrap().packet.connection_status.sequence,
buffer.buffer[1].unwrap().packet.connection_status.sequence,
buffer.buffer[2].unwrap().packet.connection_status.sequence,
];
buffer.add(make_packet::<8>(9, 9), 1, false);
let sequences_after: [u16; 3] = [
buffer.buffer[0].unwrap().packet.connection_status.sequence,
buffer.buffer[1].unwrap().packet.connection_status.sequence,
buffer.buffer[2].unwrap().packet.connection_status.sequence,
];
assert_eq!(sequences_after, sequences_before);
}
#[test]
fn test_add_reliable_replaces_packet_with_most_retries() {
let mut buffer = PacketBuffer::<3, 8>::new();
buffer.add(make_packet::<8>(1, 1), 1, true);
buffer.add(make_packet::<8>(2, 2), 1, true);
buffer.add(make_packet::<8>(3, 3), 1, true);
buffer.gather_unsent_packets_for_retry();
buffer.remove(2);
buffer.add(make_packet::<8>(4, 4), 1, true);
buffer.gather_unsent_packets_for_retry();
assert_eq!(buffer.buffer[0].unwrap().try_count, 2);
assert_eq!(buffer.buffer[1].unwrap().try_count, 1);
assert_eq!(buffer.buffer[2].unwrap().try_count, 2);
buffer.add(make_packet::<8>(9, 9), 1, true);
assert_eq!(buffer.buffer[0].unwrap().packet.connection_status.sequence, 9);
}
#[test]
fn test_replace_packets() {
let mut buffer = PacketBuffer::<4, 16>::new();
assert_eq!(buffer.gather_unsent_packets_for_retry().len(), 0);
buffer.add(DataPacket::new(1, ConnectionStatus::new(0, 0, 0), [0; 16]), 1, false);
buffer.add(DataPacket::new(1, ConnectionStatus::new(1, 0, 0), [0; 16]), 1, false);
buffer.add(DataPacket::new(1, ConnectionStatus::new(2, 0, 0), [0; 16]), 1, false);
buffer.add(DataPacket::new(1, ConnectionStatus::new(3, 0, 0), [0; 16]), 1, false);
assert_eq!(buffer.buffer.iter().filter(|p| p.is_some()).count(), 4);
buffer.add(DataPacket::new(1, ConnectionStatus::new(4, 0, 0), [0; 16]), 1, false);
assert_eq!(buffer.buffer[0].unwrap().packet.connection_status.sequence, 4); }
#[test]
fn test_remove_existing_and_missing_sequence() {
let mut buffer = PacketBuffer::<4, 8>::new();
buffer.add(make_packet::<8>(1, 1), 1, false);
buffer.add(make_packet::<8>(2, 2), 1, false);
buffer.remove(2);
assert!(buffer.buffer.iter().any(|packet| {
packet
.map(|packet| packet.packet.connection_status.sequence == 1)
.unwrap_or(false)
}));
assert!(buffer.buffer.iter().all(|packet| {
packet
.map(|packet| packet.packet.connection_status.sequence != 2)
.unwrap_or(true)
}));
let count_before = buffer.buffer.iter().filter(|packet| packet.is_some()).count();
buffer.remove(99);
let count_after = buffer.buffer.iter().filter(|packet| packet.is_some()).count();
assert_eq!(count_after, count_before);
}
#[test]
fn unreliable_packets_are_sent_once_while_reliable_packets_remain_for_retry() {
let mut buffer = PacketBuffer::<3, 8>::new();
buffer.add(make_packet::<8>(1, 1), 1, false);
buffer.add(make_packet::<8>(2, 2), 1, true);
let first = buffer.gather_unsent_packets_for_retry();
let second = buffer.gather_unsent_packets_for_retry();
assert_eq!(first.len(), 2);
assert_eq!(second.len(), 1);
assert!(buffer.buffer[0].is_none());
assert_eq!(buffer.buffer[1].unwrap().try_count, 2);
let sequences: ArrayVec<[u16; 3]> = buffer
.buffer
.iter()
.filter_map(|packet| packet.map(|packet| packet.packet.connection_status.sequence))
.collect();
assert!(!sequences.contains(&1));
assert!(sequences.contains(&2));
}
#[test]
fn reliable_retry_eventually_quiesces_when_acknowledgements_never_arrive() {
let mut buffer = PacketBuffer::<1, 8>::new();
buffer.add(make_packet::<8>(1, 1), 1, true);
for _ in 0..MAX_RELIABLE_SEND_ATTEMPTS {
assert_eq!(buffer.gather_unsent_packets_for_retry().len(), 1);
}
assert_eq!(buffer.gather_unsent_packets_for_retry().len(), 0);
assert!(buffer.buffer[0].is_none());
}
#[test]
fn acknowledgement_window_removes_exactly_the_modeled_sequences() {
let mut buffer = PacketBuffer::<8, 8>::new();
for sequence in [u16::MAX - 1, u16::MAX, 0, 1, 2, 3, 4, 5] {
buffer.add(make_packet::<8>(sequence, 1), 1, true);
}
buffer.acknowledge_packets(2, 0b1_0101);
for sequence in [2, 0, u16::MAX - 1] {
assert!(buffer
.buffer
.iter()
.all(|entry| { entry.is_none_or(|entry| entry.packet.connection_status.sequence != sequence) }));
}
for sequence in [u16::MAX, 1, 3, 4, 5] {
assert!(buffer
.buffer
.iter()
.any(|entry| { entry.is_some_and(|entry| entry.packet.connection_status.sequence == sequence) }));
}
}
}
use tinyvec::ArrayVec;
use crate::packets::DataPacket;