use std::collections::HashMap;
use std::io;
use std::time::{Duration, Instant};
use super::consts::{CONTROL_CHANNEL_MTU, TLS_RELIABLE_N_REC_BUFFERS};
use super::packet_ctrl::ControlPacket;
use super::Opcode;
pub const RETRANSMIT_INITIAL: Duration = Duration::from_secs(1);
pub const RETRANSMIT_MAX_INTERVAL: Duration = Duration::from_secs(8);
pub const RETRANSMIT_MAX_ATTEMPTS: u32 = 8;
#[derive(Debug, Default)]
pub struct RecvOutcome {
pub tls_bytes: Vec<u8>,
pub got_client_reset: bool,
}
#[derive(Debug, Clone)]
struct Unacked {
pkt: ControlPacket,
last_sent: Instant,
attempts: u32,
}
#[derive(Debug, Default)]
pub struct TickOutcome {
pub resend: Vec<Vec<u8>>,
pub timed_out: bool,
}
#[derive(Debug)]
pub struct Reliable {
pub local_id: [u8; 8],
pub peer_id: [u8; 8],
out_counter: u32,
unacked: HashMap<u32, Unacked>,
in_counter: u32, in_buf: HashMap<u32, ControlPacket>,
pending_ack: Vec<u32>,
}
fn backoff(attempts: u32) -> Duration {
let shift = attempts.saturating_sub(1).min(16);
let scaled = RETRANSMIT_INITIAL
.checked_mul(1u32 << shift)
.unwrap_or(RETRANSMIT_MAX_INTERVAL);
scaled.min(RETRANSMIT_MAX_INTERVAL)
}
fn invalid(msg: &str) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, msg.to_string())
}
impl Reliable {
pub fn new(local_id: [u8; 8]) -> Reliable {
Reliable {
local_id,
peer_id: [0u8; 8],
out_counter: 0,
unacked: HashMap::new(),
in_counter: 0,
in_buf: HashMap::new(),
pending_ack: Vec::new(),
}
}
pub fn recv(&mut self, data: &[u8]) -> io::Result<RecvOutcome> {
let pkt = ControlPacket::parse(data)?;
for pid in &pkt.acked_pids {
self.unacked.remove(pid);
}
let mut outcome = RecvOutcome::default();
if pkt.opcode == Opcode::ACK_V1 {
return Ok(outcome);
}
let pid = pkt
.pid
.ok_or_else(|| invalid("control packet missing packet id"))?;
self.pending_ack.push(pid);
if pkt.opcode == Opcode::CONTROL_HARD_RESET_CLIENT_V2
|| pkt.opcode == Opcode::CONTROL_HARD_RESET_SERVER_V2
{
self.peer_id = pkt.session_id;
if pkt.opcode == Opcode::CONTROL_HARD_RESET_CLIENT_V2 {
outcome.got_client_reset = true;
}
if pid == self.in_counter {
self.in_counter += 1;
}
return Ok(outcome);
}
if pkt.opcode != Opcode::CONTROL_V1 {
return Ok(outcome);
}
if pid < self.in_counter {
return Ok(outcome); }
if pid > self.in_counter + TLS_RELIABLE_N_REC_BUFFERS as u32 {
return Err(invalid("rejecting packet because pid looks invalid"));
}
self.in_buf.entry(pid).or_insert(pkt);
loop {
match self.in_buf.remove(&self.in_counter) {
Some(p) => {
self.in_counter += 1;
outcome.tls_bytes.extend_from_slice(&p.payload);
}
None => {
if self.in_buf.len() > TLS_RELIABLE_N_REC_BUFFERS {
return Err(invalid("received too many packets, dropping connection"));
}
break;
}
}
}
Ok(outcome)
}
pub fn take_pending_acks(&mut self) -> Vec<u32> {
std::mem::take(&mut self.pending_ack)
}
pub fn has_pending_acks(&self) -> bool {
!self.pending_ack.is_empty()
}
pub fn build_control(&mut self, payload: &[u8]) -> ControlPacket {
let mut pkt = ControlPacket::new(Opcode::CONTROL_V1, 0, self.local_id, self.peer_id);
pkt.payload = payload.to_vec();
let pid = self.out_counter;
pkt.set_pid(pid);
self.out_counter += 1;
self.track_unacked(pid, pkt.clone());
pkt
}
fn track_unacked(&mut self, pid: u32, pkt: ControlPacket) {
self.unacked.insert(
pid,
Unacked {
pkt,
last_sent: Instant::now(),
attempts: 1,
},
);
}
pub fn chunk_tls_stream(&mut self, data: &[u8]) -> Vec<ControlPacket> {
let mut packets = Vec::new();
let mut off = 0;
while off < data.len() {
let end = (off + CONTROL_CHANNEL_MTU).min(data.len());
packets.push(self.build_control(&data[off..end]));
off = end;
}
packets
}
pub fn build_hard_reset(&mut self) -> ControlPacket {
self.build_reset(Opcode::CONTROL_HARD_RESET_SERVER_V2)
}
#[allow(dead_code)]
pub fn build_client_hard_reset(&mut self) -> ControlPacket {
self.build_reset(Opcode::CONTROL_HARD_RESET_CLIENT_V2)
}
fn build_reset(&mut self, opcode: Opcode) -> ControlPacket {
let mut pkt = ControlPacket::new(opcode, 0, self.local_id, self.peer_id);
let pid = self.out_counter;
pkt.set_pid(pid);
self.out_counter += 1;
self.track_unacked(pid, pkt.clone());
pkt
}
pub fn build_ack(&self) -> ControlPacket {
ControlPacket::new(Opcode::ACK_V1, 0, self.local_id, self.peer_id)
}
#[allow(dead_code)]
pub fn unacked_packets(&self) -> Vec<ControlPacket> {
self.unacked.values().map(|u| u.pkt.clone()).collect()
}
#[allow(dead_code)]
pub fn unacked_count(&self) -> usize {
self.unacked.len()
}
pub fn tick(&mut self, now: Instant) -> TickOutcome {
let mut outcome = TickOutcome::default();
let mut due: Vec<u32> = Vec::new();
for (&pid, u) in self.unacked.iter() {
if now.duration_since(u.last_sent) >= backoff(u.attempts) {
due.push(pid);
}
}
due.sort_unstable();
for pid in due {
let u = self.unacked.get_mut(&pid).expect("pid just collected");
if u.attempts >= RETRANSMIT_MAX_ATTEMPTS {
outcome.timed_out = true;
self.unacked.remove(&pid);
continue;
}
u.attempts += 1;
u.last_sent = now;
outcome.resend.push(u.pkt.to_bytes(&[]));
}
outcome
}
}
#[cfg(test)]
mod tests {
use super::*;
fn local() -> [u8; 8] {
[1, 2, 3, 4, 5, 6, 7, 8]
}
fn client_control(pid: u32, peer_local: [u8; 8], payload: &[u8]) -> Vec<u8> {
let mut pkt = ControlPacket::new(Opcode::CONTROL_V1, 0, peer_local, [0u8; 8]);
pkt.set_pid(pid);
pkt.payload = payload.to_vec();
pkt.to_bytes(&[])
}
#[test]
fn hard_reset_sets_peer_id() {
let mut r = Reliable::new(local());
let client_sid = [9u8, 9, 9, 9, 9, 9, 9, 9];
let mut reset =
ControlPacket::new(Opcode::CONTROL_HARD_RESET_CLIENT_V2, 0, client_sid, [0; 8]);
reset.set_pid(0);
let out = r.recv(&reset.to_bytes(&[])).unwrap();
assert!(out.got_client_reset);
assert_eq!(r.peer_id, client_sid);
assert_eq!(r.take_pending_acks(), vec![0]);
}
#[test]
fn in_order_stream_reassembly() {
let mut r = Reliable::new(local());
let client_sid = [9u8; 8];
let mut reset =
ControlPacket::new(Opcode::CONTROL_HARD_RESET_CLIENT_V2, 0, client_sid, [0; 8]);
reset.set_pid(0);
r.recv(&reset.to_bytes(&[])).unwrap();
let out2 = r.recv(&client_control(2, client_sid, b"world")).unwrap();
assert!(out2.tls_bytes.is_empty(), "pid 2 should buffer");
let out1 = r.recv(&client_control(1, client_sid, b"hello")).unwrap();
assert_eq!(out1.tls_bytes, b"helloworld");
}
#[test]
fn duplicate_old_packet_ignored() {
let mut r = Reliable::new(local());
let sid = [9u8; 8];
let mut reset = ControlPacket::new(Opcode::CONTROL_HARD_RESET_CLIENT_V2, 0, sid, [0; 8]);
reset.set_pid(0);
r.recv(&reset.to_bytes(&[])).unwrap();
let out = r.recv(&client_control(1, sid, b"a")).unwrap();
assert_eq!(out.tls_bytes, b"a");
let dup = r.recv(&client_control(1, sid, b"a")).unwrap();
assert!(dup.tls_bytes.is_empty());
}
#[test]
fn outgoing_acks_remove_unacked() {
let mut r = Reliable::new(local());
let _p0 = r.build_control(b"x");
let _p1 = r.build_control(b"y");
assert_eq!(r.unacked_count(), 2);
let sid = [9u8; 8];
let mut ack = ControlPacket::new(Opcode::ACK_V1, 0, sid, r.local_id);
let data = ack_bytes(&mut ack, &[0]);
r.recv(&data).unwrap();
assert_eq!(r.unacked_count(), 1);
}
fn ack_bytes(pkt: &mut ControlPacket, acks: &[u32]) -> Vec<u8> {
pkt.to_bytes(acks)
}
#[test]
fn chunking_respects_mtu() {
let mut r = Reliable::new(local());
let big = vec![0u8; CONTROL_CHANNEL_MTU * 2 + 10];
let chunks = r.chunk_tls_stream(&big);
assert_eq!(chunks.len(), 3);
assert_eq!(chunks[0].payload.len(), CONTROL_CHANNEL_MTU);
assert_eq!(chunks[1].payload.len(), CONTROL_CHANNEL_MTU);
assert_eq!(chunks[2].payload.len(), 10);
assert_eq!(chunks[0].pid, Some(0));
assert_eq!(chunks[2].pid, Some(2));
}
#[test]
fn far_future_pid_rejected() {
let mut r = Reliable::new(local());
let sid = [9u8; 8];
let mut reset = ControlPacket::new(Opcode::CONTROL_HARD_RESET_CLIENT_V2, 0, sid, [0; 8]);
reset.set_pid(0);
r.recv(&reset.to_bytes(&[])).unwrap();
assert!(r.recv(&client_control(1000, sid, b"z")).is_err());
}
#[test]
fn backoff_doubles_and_caps() {
assert_eq!(backoff(1), RETRANSMIT_INITIAL);
assert_eq!(backoff(2), RETRANSMIT_INITIAL * 2);
assert_eq!(backoff(3), RETRANSMIT_INITIAL * 4);
assert_eq!(backoff(100), RETRANSMIT_MAX_INTERVAL);
assert!(backoff(50) <= RETRANSMIT_MAX_INTERVAL);
}
#[test]
fn unacked_past_deadline_is_resent() {
let mut r = Reliable::new(local());
let p = r.build_control(b"hello");
let start = Instant::now();
let early = r.tick(start + RETRANSMIT_INITIAL - Duration::from_millis(1));
assert!(early.resend.is_empty());
assert!(!early.timed_out);
let late = r.tick(start + RETRANSMIT_INITIAL + Duration::from_millis(1));
assert_eq!(late.resend.len(), 1);
assert_eq!(late.resend[0], p.to_bytes(&[]));
assert!(!late.timed_out);
}
#[test]
fn acked_packet_is_not_resent() {
let mut r = Reliable::new(local());
let _p = r.build_control(b"hello");
let start = Instant::now();
assert_eq!(r.unacked_count(), 1);
let sid = [9u8; 8];
let ack = ControlPacket::new(Opcode::ACK_V1, 0, sid, r.local_id);
r.recv(&ack.to_bytes(&[0])).unwrap();
assert_eq!(r.unacked_count(), 0);
let out = r.tick(start + RETRANSMIT_INITIAL * 4);
assert!(out.resend.is_empty());
}
#[test]
fn backoff_increases_between_retransmits() {
let mut r = Reliable::new(local());
let _p = r.build_control(b"x");
let start = Instant::now();
let t1 = start + RETRANSMIT_INITIAL;
assert_eq!(r.tick(t1).resend.len(), 1);
let t2 = t1 + RETRANSMIT_INITIAL;
assert!(r.tick(t2).resend.is_empty(), "backoff should have doubled");
let t3 = t1 + RETRANSMIT_INITIAL * 2;
assert_eq!(r.tick(t3).resend.len(), 1);
}
#[test]
fn retries_cap_and_signal_timeout() {
let mut r = Reliable::new(local());
let _p = r.build_control(b"x");
let mut t = Instant::now();
let mut timed_out = false;
for _ in 0..(RETRANSMIT_MAX_ATTEMPTS + 4) {
t += RETRANSMIT_MAX_INTERVAL * 2;
let out = r.tick(t);
if out.timed_out {
timed_out = true;
break;
}
}
assert!(timed_out, "packet should eventually time out");
assert_eq!(r.unacked_count(), 0);
let after = r.tick(t + RETRANSMIT_MAX_INTERVAL * 2);
assert!(after.resend.is_empty());
assert!(!after.timed_out);
}
}