1use std::collections::{BTreeMap, BTreeSet};
28
29const KIND_DATA: u8 = 0x00;
30const KIND_ACK: u8 = 0x01;
31const HEADER: usize = 5;
32
33pub const DEFAULT_RTO_MS: u64 = 500;
35pub const DEFAULT_MAX_ATTEMPTS: u32 = 8;
37const DEDUP_WINDOW: u32 = 4096;
40
41struct Unacked {
43 payload: Vec<u8>,
44 last_sent_ms: u64,
45 attempts: u32,
46}
47
48#[derive(Debug, Default, PartialEq)]
50pub struct Inbound {
51 pub delivered: Option<Vec<u8>>,
53 pub to_send: Vec<Vec<u8>>,
55}
56
57pub struct ReliableEndpoint {
60 next_seq: u32,
61 unacked: BTreeMap<u32, Unacked>,
62 received: BTreeSet<u32>,
64 highest_received: Option<u32>,
65 rto_ms: u64,
66 max_attempts: u32,
67}
68
69impl Default for ReliableEndpoint {
70 fn default() -> Self {
71 Self::new(DEFAULT_RTO_MS, DEFAULT_MAX_ATTEMPTS)
72 }
73}
74
75fn data_frame(seq: u32, payload: &[u8]) -> Vec<u8> {
76 let mut f = Vec::with_capacity(HEADER + payload.len());
77 f.push(KIND_DATA);
78 f.extend_from_slice(&seq.to_be_bytes());
79 f.extend_from_slice(payload);
80 f
81}
82
83fn ack_frame(seq: u32) -> Vec<u8> {
84 let mut f = Vec::with_capacity(HEADER);
85 f.push(KIND_ACK);
86 f.extend_from_slice(&seq.to_be_bytes());
87 f
88}
89
90impl ReliableEndpoint {
91 pub fn new(rto_ms: u64, max_attempts: u32) -> ReliableEndpoint {
93 ReliableEndpoint {
94 next_seq: 0,
95 unacked: BTreeMap::new(),
96 received: BTreeSet::new(),
97 highest_received: None,
98 rto_ms,
99 max_attempts,
100 }
101 }
102
103 pub fn send(&mut self, payload: &[u8], now_ms: u64) -> Vec<u8> {
106 let seq = self.next_seq;
107 self.next_seq = self.next_seq.wrapping_add(1);
108 let frame = data_frame(seq, payload);
109 self.unacked.insert(
110 seq,
111 Unacked {
112 payload: payload.to_vec(),
113 last_sent_ms: now_ms,
114 attempts: 1,
115 },
116 );
117 frame
118 }
119
120 pub fn pending(&self) -> usize {
122 self.unacked.len()
123 }
124
125 fn remember_received(&mut self, seq: u32) {
126 self.received.insert(seq);
127 let hi = self.highest_received.map_or(seq, |h| h.max(seq));
128 self.highest_received = Some(hi);
129 let cutoff = hi.saturating_sub(DEDUP_WINDOW);
131 while let Some(&low) = self.received.iter().next() {
132 if low < cutoff {
133 self.received.remove(&low);
134 } else {
135 break;
136 }
137 }
138 }
139
140 pub fn on_datagram(&mut self, frame: &[u8], _now_ms: u64) -> Inbound {
144 if frame.len() < HEADER {
145 return Inbound::default();
146 }
147 let seq = u32::from_be_bytes([frame[1], frame[2], frame[3], frame[4]]);
148 match frame[0] {
149 KIND_DATA => {
150 let is_new = !self.received.contains(&seq);
151 let delivered = if is_new {
152 self.remember_received(seq);
153 Some(frame[HEADER..].to_vec())
154 } else {
155 None };
157 Inbound {
158 delivered,
159 to_send: vec![ack_frame(seq)],
160 }
161 }
162 KIND_ACK => {
163 self.unacked.remove(&seq);
164 Inbound::default()
165 }
166 _ => Inbound::default(),
167 }
168 }
169
170 pub fn on_tick(&mut self, now_ms: u64) -> (Vec<Vec<u8>>, Vec<u32>) {
173 let mut resend = Vec::new();
174 let mut gave_up = Vec::new();
175 let mut abandon = Vec::new();
176
177 for (&seq, u) in self.unacked.iter_mut() {
178 if now_ms.saturating_sub(u.last_sent_ms) < self.rto_ms {
179 continue;
180 }
181 if u.attempts >= self.max_attempts {
182 abandon.push(seq);
183 continue;
184 }
185 u.attempts += 1;
186 u.last_sent_ms = now_ms;
187 resend.push(data_frame(seq, &u.payload));
188 }
189 for seq in abandon {
190 self.unacked.remove(&seq);
191 gave_up.push(seq);
192 }
193 (resend, gave_up)
194 }
195}
196
197#[cfg(test)]
198mod tests {
199 use super::*;
200
201 fn seq_of(frame: &[u8]) -> u32 {
202 u32::from_be_bytes([frame[1], frame[2], frame[3], frame[4]])
203 }
204
205 #[test]
206 fn deliver_and_ack_clears_unacked() {
207 let mut a = ReliableEndpoint::default();
208 let mut b = ReliableEndpoint::default();
209
210 let data = a.send(b"hello", 0);
211 assert_eq!(a.pending(), 1);
212
213 let inb = b.on_datagram(&data, 1);
215 assert_eq!(inb.delivered.as_deref(), Some(&b"hello"[..]));
216 assert_eq!(inb.to_send.len(), 1);
217
218 let ack = &inb.to_send[0];
220 let back = a.on_datagram(ack, 2);
221 assert!(back.delivered.is_none());
222 assert_eq!(a.pending(), 0);
223 }
224
225 #[test]
226 fn duplicate_data_is_delivered_once_but_acked_each_time() {
227 let mut a = ReliableEndpoint::default();
228 let mut b = ReliableEndpoint::default();
229 let data = a.send(b"dup", 0);
230
231 let first = b.on_datagram(&data, 1);
232 assert_eq!(first.delivered.as_deref(), Some(&b"dup"[..]));
233 assert_eq!(first.to_send.len(), 1, "acked");
234
235 let second = b.on_datagram(&data, 2);
237 assert!(second.delivered.is_none(), "not delivered twice");
238 assert_eq!(second.to_send.len(), 1, "still acked so sender stops");
239 }
240
241 #[test]
242 fn retransmits_after_rto_then_stops_once_acked() {
243 let mut a = ReliableEndpoint::new(100, 8);
244 let _data = a.send(b"x", 0);
245
246 let (resend, gave_up) = a.on_tick(50);
248 assert!(resend.is_empty() && gave_up.is_empty());
249
250 let (resend, _) = a.on_tick(150);
252 assert_eq!(resend.len(), 1);
253 assert_eq!(seq_of(&resend[0]), 0);
254
255 let mut b = ReliableEndpoint::default();
257 let inb = b.on_datagram(&resend[0], 160);
258 a.on_datagram(&inb.to_send[0], 170);
259 let (resend, _) = a.on_tick(1000);
260 assert!(resend.is_empty(), "acked → no retransmit");
261 }
262
263 #[test]
264 fn gives_up_after_max_attempts() {
265 let mut a = ReliableEndpoint::new(10, 3);
266 a.send(b"lost", 0); let (_r, g) = a.on_tick(20); assert!(g.is_empty());
270 let (_r, g) = a.on_tick(40); assert!(g.is_empty());
272 let (r, g) = a.on_tick(60); assert!(r.is_empty());
274 assert_eq!(g, vec![0]);
275 assert_eq!(a.pending(), 0, "abandoned datagram dropped");
276 }
277
278 #[test]
279 fn out_of_order_and_lossy_stream_delivers_every_payload_once() {
280 let mut a = ReliableEndpoint::new(100, 8);
283 let f0 = a.send(b"m0", 0);
284 let _f1 = a.send(b"m1", 0);
287 let f2 = a.send(b"m2", 0);
288
289 let mut b = ReliableEndpoint::default();
290 let mut delivered: Vec<Vec<u8>> = Vec::new();
291 let mut push = |inb: Inbound| {
292 if let Some(p) = inb.delivered {
293 delivered.push(p);
294 }
295 };
296
297 push(b.on_datagram(&f2, 1)); push(b.on_datagram(&f0, 2));
299 push(b.on_datagram(&f2, 3)); let (resend, _) = a.on_tick(200);
302 assert!(resend.iter().any(|f| seq_of(f) == 1));
303 for f in &resend {
304 if seq_of(f) == 1 {
305 push(b.on_datagram(f, 210));
306 }
307 }
308
309 delivered.sort();
310 assert_eq!(
311 delivered,
312 vec![b"m0".to_vec(), b"m1".to_vec(), b"m2".to_vec()]
313 );
314 }
315}