1use alloc::collections::{BTreeSet, VecDeque};
26use alloc::vec::Vec;
27use core::time::Duration;
28
29use crate::packet::{
30 AckAckPacket, AckCif, AckPacket, ControlPacket, DataPacket, EncryptionKeyField, LossListEntry,
31 NakPacket, PacketPosition,
32};
33
34use super::rtt::RttEstimator;
35use super::{duration_to_wire_us, seq};
36
37const MAX_RANGE_EXPANSION: u32 = 1 << 16;
42
43#[derive(Debug, Clone)]
46struct SentPacket {
47 seq: u32,
48 message_number: u32,
49 payload: Vec<u8>,
50 resend_count: u32,
52}
53
54#[derive(Debug)]
57pub struct Sender {
58 dest_socket_id: u32,
59 buffer: VecDeque<SentPacket>,
61 pending_retransmit: BTreeSet<u32>,
64 rtt: RttEstimator,
65}
66
67impl Sender {
68 pub fn new(dest_socket_id: u32) -> Self {
71 Sender {
72 dest_socket_id,
73 buffer: VecDeque::new(),
74 pending_retransmit: BTreeSet::new(),
75 rtt: RttEstimator::new(),
76 }
77 }
78
79 pub fn rtt(&self) -> Duration {
82 self.rtt.rtt()
83 }
84
85 pub fn rtt_var(&self) -> Duration {
87 self.rtt.rtt_var()
88 }
89
90 pub fn buffered_count(&self) -> usize {
92 self.buffer.len()
93 }
94
95 pub fn pending_retransmit_count(&self) -> usize {
97 self.pending_retransmit.len()
98 }
99
100 pub fn on_data(
104 &mut self,
105 seq: u32,
106 message_number: u32,
107 payload: &[u8],
108 now: Duration,
109 ) -> Vec<u8> {
110 self.buffer.push_back(SentPacket {
111 seq,
112 message_number,
113 payload: payload.to_vec(),
114 resend_count: 0,
115 });
116 let pkt = DataPacket {
117 seq_number: seq,
118 position: PacketPosition::Solo,
119 in_order: true,
120 key_flag: EncryptionKeyField::NotEncrypted,
121 retransmitted: false,
122 message_number,
123 timestamp: duration_to_wire_us(now),
124 dest_socket_id: self.dest_socket_id,
125 data: payload,
126 };
127 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
128 pkt.serialize_into(&mut buf)
129 .expect("buffer sized from serialized_len");
130 buf
131 }
132
133 pub fn on_nak(&mut self, nak: &NakPacket<'_>) {
138 for entry in nak.entries() {
139 let Ok(entry) = entry else { continue };
140 for seq in expand_loss_entry(entry) {
141 if self.buffer.iter().any(|p| p.seq == seq) {
142 self.pending_retransmit.insert(seq);
143 }
144 }
145 }
146 }
147
148 pub fn tick(&mut self, now: Duration) -> Vec<Vec<u8>> {
154 let seqs: Vec<u32> = core::mem::take(&mut self.pending_retransmit)
155 .into_iter()
156 .collect();
157 let mut out = Vec::with_capacity(seqs.len());
158 for seq in seqs {
159 let Some(sent) = self.buffer.iter_mut().find(|p| p.seq == seq) else {
160 continue; };
162 sent.resend_count += 1;
163 let pkt = DataPacket {
164 seq_number: sent.seq,
165 position: PacketPosition::Solo,
166 in_order: true,
167 key_flag: EncryptionKeyField::NotEncrypted,
168 retransmitted: true,
169 message_number: sent.message_number,
170 timestamp: duration_to_wire_us(now),
171 dest_socket_id: self.dest_socket_id,
172 data: &sent.payload,
173 };
174 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
175 pkt.serialize_into(&mut buf)
176 .expect("buffer sized from serialized_len");
177 out.push(buf);
178 }
179 out
180 }
181
182 pub fn on_ack(&mut self, ack: &AckPacket, now: Duration) -> Option<Vec<u8>> {
186 let last_ack_seq = match ack.cif {
187 AckCif::Full { last_ack_seq, .. }
188 | AckCif::Small { last_ack_seq, .. }
189 | AckCif::Light { last_ack_seq } => last_ack_seq,
190 };
191 while let Some(front) = self.buffer.front() {
193 if seq::seq_lt(front.seq, last_ack_seq) {
194 let freed = self.buffer.pop_front().expect("front just matched");
195 self.pending_retransmit.remove(&freed.seq); } else {
197 break;
198 }
199 }
200 let buffered: BTreeSet<u32> = self.buffer.iter().map(|p| p.seq).collect();
204 self.pending_retransmit.retain(|s| buffered.contains(s));
205
206 if let AckCif::Full { rtt_us, .. } = ack.cif {
207 self.rtt.update(Duration::from_micros(u64::from(rtt_us)));
210
211 let pkt = ControlPacket::AckAck(AckAckPacket {
212 ack_number: ack.ack_number,
213 timestamp: duration_to_wire_us(now),
214 dest_socket_id: self.dest_socket_id,
215 });
216 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
217 pkt.serialize_into(&mut buf)
218 .expect("buffer sized from serialized_len");
219 Some(buf)
220 } else {
221 None
228 }
229 }
230}
231
232fn expand_loss_entry(entry: LossListEntry) -> Vec<u32> {
233 match entry {
234 LossListEntry::Single(s) => alloc::vec![s],
235 LossListEntry::Range(first, last) => {
236 let mut out = Vec::new();
237 let mut s = first;
238 let mut n = 0u32;
239 loop {
240 out.push(s);
241 if s == last || n >= MAX_RANGE_EXPANSION {
242 break;
243 }
244 s = seq::seq_add(s, 1);
245 n += 1;
246 }
247 out
248 }
249 }
250}
251
252#[cfg(test)]
253mod tests {
254 use super::*;
255 use crate::packet::nak::build_loss_list;
256
257 const PEER: u32 = 0xAAAA;
258
259 fn nak_bytes(entries: &[LossListEntry]) -> Vec<u8> {
260 let raw = build_loss_list(entries).unwrap();
261 let pkt = ControlPacket::Nak(NakPacket {
262 timestamp: 0,
263 dest_socket_id: PEER,
264 raw_loss_list: &raw,
265 });
266 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
267 pkt.serialize_into(&mut buf).unwrap();
268 buf
269 }
270
271 #[test]
272 fn on_data_buffers_and_returns_wire_bytes() {
273 let mut s = Sender::new(PEER);
274 let bytes = s.on_data(5, 5, b"hello", Duration::from_millis(1));
275 assert_eq!(s.buffered_count(), 1);
276 let dp = DataPacket::parse(&bytes).unwrap();
277 assert_eq!(dp.seq_number, 5);
278 assert!(!dp.retransmitted);
279 assert_eq!(dp.data, b"hello");
280 }
281
282 #[test]
283 fn nak_then_tick_retransmits_with_r_flag_set() {
284 let mut s = Sender::new(PEER);
285 s.on_data(0, 0, b"a", Duration::ZERO);
286 s.on_data(1, 1, b"b", Duration::ZERO);
287
288 let raw = nak_bytes(&[LossListEntry::Single(1)]);
289 let ControlPacket::Nak(nak) = ControlPacket::parse(&raw).unwrap() else {
290 panic!("expected NAK");
291 };
292 s.on_nak(&nak);
293 assert_eq!(s.pending_retransmit_count(), 1);
294
295 let out = s.tick(Duration::from_millis(5));
296 assert_eq!(out.len(), 1);
297 let dp = DataPacket::parse(&out[0]).unwrap();
298 assert_eq!(dp.seq_number, 1);
299 assert!(dp.retransmitted);
300 assert_eq!(s.pending_retransmit_count(), 0);
301 }
302
303 #[test]
304 fn nak_for_unbuffered_seq_is_ignored() {
305 let mut s = Sender::new(PEER);
306 s.on_data(0, 0, b"a", Duration::ZERO);
307 let raw = nak_bytes(&[LossListEntry::Single(99)]);
308 let ControlPacket::Nak(nak) = ControlPacket::parse(&raw).unwrap() else {
309 panic!("expected NAK");
310 };
311 s.on_nak(&nak);
312 assert_eq!(s.pending_retransmit_count(), 0);
313 assert!(s.tick(Duration::ZERO).is_empty());
314 }
315
316 #[test]
317 fn full_ack_frees_buffer_and_updates_rtt_and_replies_ackack() {
318 let mut s = Sender::new(PEER);
319 s.on_data(0, 0, b"a", Duration::ZERO);
320 s.on_data(1, 1, b"b", Duration::ZERO);
321 s.on_data(2, 2, b"c", Duration::ZERO);
322
323 let ack = AckPacket {
324 ack_number: 1,
325 timestamp: 0,
326 dest_socket_id: PEER,
327 cif: AckCif::Full {
328 last_ack_seq: 2,
329 rtt_us: 20_000,
330 rtt_var_us: 5_000,
331 avail_buf_size: 0,
332 pkt_recv_rate: 0,
333 est_link_capacity: 0,
334 recv_rate_bps: 0,
335 },
336 };
337 let reply = s.on_ack(&ack, Duration::from_millis(1)).unwrap();
338 assert_eq!(s.buffered_count(), 1); assert!(s.rtt() < Duration::from_millis(100)); let ControlPacket::AckAck(ackack) = ControlPacket::parse(&reply).unwrap() else {
342 panic!("expected ACKACK");
343 };
344 assert_eq!(ackack.ack_number, 1);
345 }
346
347 #[test]
348 fn light_ack_does_not_trigger_ackack() {
349 let mut s = Sender::new(PEER);
350 s.on_data(0, 0, b"a", Duration::ZERO);
351 let ack = AckPacket {
352 ack_number: 0,
353 timestamp: 0,
354 dest_socket_id: PEER,
355 cif: AckCif::Light { last_ack_seq: 1 },
356 };
357 assert!(s.on_ack(&ack, Duration::ZERO).is_none());
358 assert_eq!(s.buffered_count(), 0);
359 }
360}