1use alloc::collections::{BTreeMap, BTreeSet};
22use alloc::vec::Vec;
23use core::time::Duration;
24
25use crate::packet::nak::build_loss_list;
26use crate::packet::{AckAckPacket, AckCif, AckPacket, ControlPacket, LossListEntry, NakPacket};
27
28use super::rtt::RttEstimator;
29use super::{FULL_ACK_PERIOD, LIGHT_ACK_THRESHOLD, duration_to_wire_us, nak_interval, seq};
30
31const MAX_GAP_EXPANSION: u32 = 1 << 16;
37
38#[derive(Debug, Clone, PartialEq, Eq, Default)]
40#[non_exhaustive]
41pub struct FeedOutcome {
42 pub delivered: Vec<u32>,
47 pub nak: Option<Vec<u8>>,
50}
51
52#[derive(Debug)]
54pub struct Receiver {
55 dest_socket_id: u32,
56 next_expected: u32,
59 out_of_order: BTreeSet<u32>,
61 loss_list: BTreeSet<u32>,
63 highest_received: Option<u32>,
66 packets_since_ack: u32,
68 last_full_ack_at: Duration,
70 last_nak_at: Duration,
72 next_ack_number: u32,
74 outstanding_acks: BTreeMap<u32, Duration>,
77 rtt: RttEstimator,
78}
79
80impl Receiver {
81 pub fn new(dest_socket_id: u32, initial_seq: u32) -> Self {
84 Receiver {
85 dest_socket_id,
86 next_expected: initial_seq,
87 out_of_order: BTreeSet::new(),
88 loss_list: BTreeSet::new(),
89 highest_received: None,
90 packets_since_ack: 0,
91 last_full_ack_at: Duration::ZERO,
92 last_nak_at: Duration::ZERO,
93 next_ack_number: 1,
94 outstanding_acks: BTreeMap::new(),
95 rtt: RttEstimator::new(),
96 }
97 }
98
99 pub fn ack_point(&self) -> u32 {
102 self.next_expected
103 }
104
105 pub fn rtt(&self) -> Duration {
107 self.rtt.rtt()
108 }
109
110 pub fn rtt_var(&self) -> Duration {
112 self.rtt.rtt_var()
113 }
114
115 pub fn loss_list_len(&self) -> usize {
117 self.loss_list.len()
118 }
119
120 pub fn feed_data(&mut self, seq_number: u32, now: Duration) -> FeedOutcome {
122 self.packets_since_ack = self.packets_since_ack.saturating_add(1);
123
124 let mut newly_lost = Vec::new();
125 match self.highest_received {
126 None => self.highest_received = Some(seq_number),
127 Some(highest) if seq::seq_gt(seq_number, highest) => {
128 let mut s = seq::seq_next(highest);
129 let mut n = 0u32;
130 while s != seq_number && n < MAX_GAP_EXPANSION {
131 newly_lost.push(s);
132 self.loss_list.insert(s);
133 s = seq::seq_next(s);
134 n += 1;
135 }
136 self.highest_received = Some(seq_number);
137 }
138 _ => {}
139 }
140
141 self.loss_list.remove(&seq_number);
142
143 let mut delivered = Vec::new();
144 if seq_number == self.next_expected {
145 delivered.push(seq_number);
146 self.next_expected = seq::seq_next(seq_number);
147 while self.out_of_order.remove(&self.next_expected) {
148 delivered.push(self.next_expected);
149 self.next_expected = seq::seq_next(self.next_expected);
150 }
151 } else if seq::seq_gt(seq_number, self.next_expected) {
152 self.out_of_order.insert(seq_number);
153 }
154 let nak = if newly_lost.is_empty() {
159 None
160 } else {
161 Some(self.build_nak(&newly_lost, now))
162 };
163
164 FeedOutcome { delivered, nak }
165 }
166
167 fn build_nak(&self, seqs: &[u32], now: Duration) -> Vec<u8> {
168 let entries = coalesce(seqs);
169 let raw = build_loss_list(&entries).expect("seq numbers are 31-bit by construction");
170 let pkt = ControlPacket::Nak(NakPacket {
171 timestamp: duration_to_wire_us(now),
172 dest_socket_id: self.dest_socket_id,
173 raw_loss_list: &raw,
174 });
175 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
176 pkt.serialize_into(&mut buf)
177 .expect("buffer sized from serialized_len");
178 buf
179 }
180
181 pub fn tick(&mut self, now: Duration) -> Vec<Vec<u8>> {
188 let mut out = Vec::new();
189
190 if elapsed(now, self.last_full_ack_at) >= FULL_ACK_PERIOD {
191 out.push(self.build_full_ack(now));
192 self.last_full_ack_at = now;
193 self.packets_since_ack = 0;
194 } else if self.packets_since_ack >= LIGHT_ACK_THRESHOLD {
195 out.push(self.build_light_ack());
196 self.packets_since_ack = 0;
197 }
198
199 let interval = nak_interval(self.rtt.rtt(), self.rtt.rtt_var());
200 if !self.loss_list.is_empty() && elapsed(now, self.last_nak_at) >= interval {
201 let seqs: Vec<u32> = self.loss_list.iter().copied().collect();
202 out.push(self.build_nak(&seqs, now));
203 self.last_nak_at = now;
204 }
205
206 out
207 }
208
209 fn build_full_ack(&mut self, now: Duration) -> Vec<u8> {
210 let ack_number = self.next_ack_number;
211 self.next_ack_number = self.next_ack_number.wrapping_add(1);
212 self.outstanding_acks.insert(ack_number, now);
213 let pkt = ControlPacket::Ack(AckPacket {
214 ack_number,
215 timestamp: duration_to_wire_us(now),
216 dest_socket_id: self.dest_socket_id,
217 cif: AckCif::Full {
218 last_ack_seq: self.next_expected,
219 rtt_us: self.rtt.rtt_us(),
220 rtt_var_us: self.rtt.rtt_var_us(),
221 avail_buf_size: 0,
225 pkt_recv_rate: 0,
226 est_link_capacity: 0,
227 recv_rate_bps: 0,
228 },
229 });
230 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
231 pkt.serialize_into(&mut buf)
232 .expect("buffer sized from serialized_len");
233 buf
234 }
235
236 fn build_light_ack(&self) -> Vec<u8> {
237 let pkt = ControlPacket::Ack(AckPacket {
241 ack_number: 0,
242 timestamp: 0,
243 dest_socket_id: self.dest_socket_id,
244 cif: AckCif::Light {
245 last_ack_seq: self.next_expected,
246 },
247 });
248 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
249 pkt.serialize_into(&mut buf)
250 .expect("buffer sized from serialized_len");
251 buf
252 }
253
254 pub fn on_ackack(&mut self, ackack: &AckAckPacket, now: Duration) {
259 if let Some(sent_at) = self.outstanding_acks.remove(&ackack.ack_number) {
260 self.rtt.update(elapsed(now, sent_at));
261 }
262 }
263}
264
265fn elapsed(now: Duration, since: Duration) -> Duration {
269 now.checked_sub(since).unwrap_or(Duration::ZERO)
270}
271
272fn coalesce(seqs: &[u32]) -> Vec<LossListEntry> {
276 let mut out = Vec::new();
277 let mut i = 0;
278 while i < seqs.len() {
279 let start = seqs[i];
280 let mut end = start;
281 let mut j = i + 1;
282 while j < seqs.len() && seqs[j] == seq::seq_next(end) {
283 end = seqs[j];
284 j += 1;
285 }
286 if start == end {
287 out.push(LossListEntry::Single(start));
288 } else {
289 out.push(LossListEntry::Range(start, end));
290 }
291 i = j;
292 }
293 out
294}
295
296#[cfg(test)]
297mod tests {
298 use super::*;
299
300 const PEER: u32 = 0xBBBB;
301
302 #[test]
303 fn in_order_arrivals_deliver_immediately_without_nak() {
304 let mut r = Receiver::new(PEER, 0);
305 for seq_number in 0..5u32 {
306 let outcome = r.feed_data(seq_number, Duration::ZERO);
307 assert_eq!(outcome.delivered, alloc::vec![seq_number]);
308 assert!(outcome.nak.is_none());
309 }
310 assert_eq!(r.ack_point(), 5);
311 assert_eq!(r.loss_list_len(), 0);
312 }
313
314 #[test]
315 fn a_gap_triggers_an_immediate_nak_and_stalls_delivery() {
316 let mut r = Receiver::new(PEER, 0);
317 r.feed_data(0, Duration::ZERO);
318 r.feed_data(1, Duration::ZERO);
319 let outcome = r.feed_data(3, Duration::ZERO); assert!(outcome.delivered.is_empty());
321 let nak = outcome.nak.expect("gap must trigger an immediate NAK");
322 let ControlPacket::Nak(n) = ControlPacket::parse(&nak).unwrap() else {
323 panic!("expected NAK");
324 };
325 let entries: Vec<LossListEntry> = n.entries().map(|e| e.unwrap()).collect();
326 assert_eq!(entries, alloc::vec![LossListEntry::Single(2)]);
327 assert_eq!(r.ack_point(), 2); assert_eq!(r.loss_list_len(), 1);
329
330 let fill = r.feed_data(2, Duration::ZERO);
332 assert_eq!(fill.delivered, alloc::vec![2, 3]);
333 assert!(fill.nak.is_none());
334 assert_eq!(r.ack_point(), 4);
335 assert_eq!(r.loss_list_len(), 0);
336 }
337
338 #[test]
339 fn zero_loss_tick_never_emits_a_nak() {
340 let mut r = Receiver::new(PEER, 0);
341 for seq_number in 0..5u32 {
342 r.feed_data(seq_number, Duration::ZERO);
343 }
344 for ms in 1..200u64 {
345 let out = r.tick(Duration::from_millis(ms));
346 for bytes in &out {
347 assert!(!matches!(
348 ControlPacket::parse(bytes).unwrap(),
349 ControlPacket::Nak(_)
350 ));
351 }
352 }
353 }
354
355 #[test]
356 fn full_ack_fires_on_the_10ms_timer_and_light_ack_on_the_64_packet_threshold() {
357 let mut r = Receiver::new(PEER, 0);
358 let out = r.tick(FULL_ACK_PERIOD);
359 assert_eq!(out.len(), 1);
360 let ControlPacket::Ack(ack) = ControlPacket::parse(&out[0]).unwrap() else {
361 panic!("expected ACK");
362 };
363 assert!(matches!(ack.cif, AckCif::Full { .. }));
364 assert_eq!(ack.ack_number, 1);
365
366 for seq_number in 0..LIGHT_ACK_THRESHOLD {
368 r.feed_data(seq_number, Duration::ZERO);
369 }
370 let out = r.tick(FULL_ACK_PERIOD + Duration::from_millis(1));
371 let out2 = r.tick(FULL_ACK_PERIOD + Duration::from_millis(2));
374 let light =
375 out.into_iter()
376 .chain(out2)
377 .find_map(|b| match ControlPacket::parse(&b).unwrap() {
378 ControlPacket::Ack(a) if matches!(a.cif, AckCif::Light { .. }) => Some(a),
379 _ => None,
380 });
381 assert!(
382 light.is_some(),
383 "expected a Light ACK from the 64-packet threshold"
384 );
385 }
386
387 #[test]
388 fn ackack_updates_rtt_from_the_measured_round_trip() {
389 let mut r = Receiver::new(PEER, 0);
390 let out = r.tick(FULL_ACK_PERIOD);
391 let ControlPacket::Ack(ack) = ControlPacket::parse(&out[0]).unwrap() else {
392 panic!("expected ACK");
393 };
394 let ackack = AckAckPacket {
395 ack_number: ack.ack_number,
396 timestamp: 0,
397 dest_socket_id: PEER,
398 };
399 let sample = Duration::from_millis(20);
400 r.on_ackack(&ackack, FULL_ACK_PERIOD + sample);
401 assert!(r.rtt() < Duration::from_millis(100));
403 assert!(r.rtt() > sample);
404 }
405}