use alloc::collections::{BTreeSet, VecDeque};
use alloc::vec::Vec;
use core::time::Duration;
use crate::packet::{
AckAckPacket, AckCif, AckPacket, ControlPacket, DataPacket, EncryptionKeyField, LossListEntry,
NakPacket, PacketPosition,
};
use super::rtt::RttEstimator;
use super::{duration_to_wire_us, seq};
const MAX_RANGE_EXPANSION: u32 = 1 << 16;
#[derive(Debug, Clone)]
struct SentPacket {
seq: u32,
message_number: u32,
payload: Vec<u8>,
resend_count: u32,
}
#[derive(Debug)]
pub struct Sender {
dest_socket_id: u32,
buffer: VecDeque<SentPacket>,
pending_retransmit: BTreeSet<u32>,
rtt: RttEstimator,
}
impl Sender {
pub fn new(dest_socket_id: u32) -> Self {
Sender {
dest_socket_id,
buffer: VecDeque::new(),
pending_retransmit: BTreeSet::new(),
rtt: RttEstimator::new(),
}
}
pub fn rtt(&self) -> Duration {
self.rtt.rtt()
}
pub fn rtt_var(&self) -> Duration {
self.rtt.rtt_var()
}
pub fn buffered_count(&self) -> usize {
self.buffer.len()
}
pub fn pending_retransmit_count(&self) -> usize {
self.pending_retransmit.len()
}
pub fn on_data(
&mut self,
seq: u32,
message_number: u32,
payload: &[u8],
now: Duration,
) -> Vec<u8> {
self.buffer.push_back(SentPacket {
seq,
message_number,
payload: payload.to_vec(),
resend_count: 0,
});
let pkt = DataPacket {
seq_number: seq,
position: PacketPosition::Solo,
in_order: true,
key_flag: EncryptionKeyField::NotEncrypted,
retransmitted: false,
message_number,
timestamp: duration_to_wire_us(now),
dest_socket_id: self.dest_socket_id,
data: payload,
};
let mut buf = alloc::vec![0u8; pkt.serialized_len()];
pkt.serialize_into(&mut buf)
.expect("buffer sized from serialized_len");
buf
}
pub fn on_nak(&mut self, nak: &NakPacket<'_>) {
for entry in nak.entries() {
let Ok(entry) = entry else { continue };
for seq in expand_loss_entry(entry) {
if self.buffer.iter().any(|p| p.seq == seq) {
self.pending_retransmit.insert(seq);
}
}
}
}
pub fn tick(&mut self, now: Duration) -> Vec<Vec<u8>> {
let seqs: Vec<u32> = core::mem::take(&mut self.pending_retransmit)
.into_iter()
.collect();
let mut out = Vec::with_capacity(seqs.len());
for seq in seqs {
let Some(sent) = self.buffer.iter_mut().find(|p| p.seq == seq) else {
continue; };
sent.resend_count += 1;
let pkt = DataPacket {
seq_number: sent.seq,
position: PacketPosition::Solo,
in_order: true,
key_flag: EncryptionKeyField::NotEncrypted,
retransmitted: true,
message_number: sent.message_number,
timestamp: duration_to_wire_us(now),
dest_socket_id: self.dest_socket_id,
data: &sent.payload,
};
let mut buf = alloc::vec![0u8; pkt.serialized_len()];
pkt.serialize_into(&mut buf)
.expect("buffer sized from serialized_len");
out.push(buf);
}
out
}
pub fn on_ack(&mut self, ack: &AckPacket, now: Duration) -> Option<Vec<u8>> {
let last_ack_seq = match ack.cif {
AckCif::Full { last_ack_seq, .. }
| AckCif::Small { last_ack_seq, .. }
| AckCif::Light { last_ack_seq } => last_ack_seq,
};
while let Some(front) = self.buffer.front() {
if seq::seq_lt(front.seq, last_ack_seq) {
let freed = self.buffer.pop_front().expect("front just matched");
self.pending_retransmit.remove(&freed.seq); } else {
break;
}
}
let buffered: BTreeSet<u32> = self.buffer.iter().map(|p| p.seq).collect();
self.pending_retransmit.retain(|s| buffered.contains(s));
if let AckCif::Full { rtt_us, .. } = ack.cif {
self.rtt.update(Duration::from_micros(u64::from(rtt_us)));
let pkt = ControlPacket::AckAck(AckAckPacket {
ack_number: ack.ack_number,
timestamp: duration_to_wire_us(now),
dest_socket_id: self.dest_socket_id,
});
let mut buf = alloc::vec![0u8; pkt.serialized_len()];
pkt.serialize_into(&mut buf)
.expect("buffer sized from serialized_len");
Some(buf)
} else {
None
}
}
}
fn expand_loss_entry(entry: LossListEntry) -> Vec<u32> {
match entry {
LossListEntry::Single(s) => alloc::vec![s],
LossListEntry::Range(first, last) => {
let mut out = Vec::new();
let mut s = first;
let mut n = 0u32;
loop {
out.push(s);
if s == last || n >= MAX_RANGE_EXPANSION {
break;
}
s = seq::seq_add(s, 1);
n += 1;
}
out
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::packet::nak::build_loss_list;
const PEER: u32 = 0xAAAA;
fn nak_bytes(entries: &[LossListEntry]) -> Vec<u8> {
let raw = build_loss_list(entries).unwrap();
let pkt = ControlPacket::Nak(NakPacket {
timestamp: 0,
dest_socket_id: PEER,
raw_loss_list: &raw,
});
let mut buf = alloc::vec![0u8; pkt.serialized_len()];
pkt.serialize_into(&mut buf).unwrap();
buf
}
#[test]
fn on_data_buffers_and_returns_wire_bytes() {
let mut s = Sender::new(PEER);
let bytes = s.on_data(5, 5, b"hello", Duration::from_millis(1));
assert_eq!(s.buffered_count(), 1);
let dp = DataPacket::parse(&bytes).unwrap();
assert_eq!(dp.seq_number, 5);
assert!(!dp.retransmitted);
assert_eq!(dp.data, b"hello");
}
#[test]
fn nak_then_tick_retransmits_with_r_flag_set() {
let mut s = Sender::new(PEER);
s.on_data(0, 0, b"a", Duration::ZERO);
s.on_data(1, 1, b"b", Duration::ZERO);
let raw = nak_bytes(&[LossListEntry::Single(1)]);
let ControlPacket::Nak(nak) = ControlPacket::parse(&raw).unwrap() else {
panic!("expected NAK");
};
s.on_nak(&nak);
assert_eq!(s.pending_retransmit_count(), 1);
let out = s.tick(Duration::from_millis(5));
assert_eq!(out.len(), 1);
let dp = DataPacket::parse(&out[0]).unwrap();
assert_eq!(dp.seq_number, 1);
assert!(dp.retransmitted);
assert_eq!(s.pending_retransmit_count(), 0);
}
#[test]
fn nak_for_unbuffered_seq_is_ignored() {
let mut s = Sender::new(PEER);
s.on_data(0, 0, b"a", Duration::ZERO);
let raw = nak_bytes(&[LossListEntry::Single(99)]);
let ControlPacket::Nak(nak) = ControlPacket::parse(&raw).unwrap() else {
panic!("expected NAK");
};
s.on_nak(&nak);
assert_eq!(s.pending_retransmit_count(), 0);
assert!(s.tick(Duration::ZERO).is_empty());
}
#[test]
fn full_ack_frees_buffer_and_updates_rtt_and_replies_ackack() {
let mut s = Sender::new(PEER);
s.on_data(0, 0, b"a", Duration::ZERO);
s.on_data(1, 1, b"b", Duration::ZERO);
s.on_data(2, 2, b"c", Duration::ZERO);
let ack = AckPacket {
ack_number: 1,
timestamp: 0,
dest_socket_id: PEER,
cif: AckCif::Full {
last_ack_seq: 2,
rtt_us: 20_000,
rtt_var_us: 5_000,
avail_buf_size: 0,
pkt_recv_rate: 0,
est_link_capacity: 0,
recv_rate_bps: 0,
},
};
let reply = s.on_ack(&ack, Duration::from_millis(1)).unwrap();
assert_eq!(s.buffered_count(), 1); assert!(s.rtt() < Duration::from_millis(100));
let ControlPacket::AckAck(ackack) = ControlPacket::parse(&reply).unwrap() else {
panic!("expected ACKACK");
};
assert_eq!(ackack.ack_number, 1);
}
#[test]
fn light_ack_does_not_trigger_ackack() {
let mut s = Sender::new(PEER);
s.on_data(0, 0, b"a", Duration::ZERO);
let ack = AckPacket {
ack_number: 0,
timestamp: 0,
dest_socket_id: PEER,
cif: AckCif::Light { last_ack_seq: 1 },
};
assert!(s.on_ack(&ack, Duration::ZERO).is_none());
assert_eq!(s.buffered_count(), 0);
}
}