use std::collections::HashMap;
use std::io;
use super::consts::{CONTROL_CHANNEL_MTU, TLS_RELIABLE_N_REC_BUFFERS};
use super::packet_ctrl::ControlPacket;
use super::Opcode;
#[derive(Debug, Default)]
pub struct RecvOutcome {
pub tls_bytes: Vec<u8>,
pub got_client_reset: bool,
}
#[derive(Debug)]
pub struct Reliable {
pub local_id: [u8; 8],
pub peer_id: [u8; 8],
out_counter: u32,
unacked: HashMap<u32, ControlPacket>,
in_counter: u32, in_buf: HashMap<u32, ControlPacket>,
pending_ack: Vec<u32>,
}
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.unacked.insert(pid, pkt.clone());
pkt
}
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.unacked.insert(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().cloned().collect()
}
#[allow(dead_code)]
pub fn unacked_count(&self) -> usize {
self.unacked.len()
}
}
#[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());
}
}