use std::net::{SocketAddr, SocketAddrV4, UdpSocket};
use std::collections::{HashMap, hash_map};
use std::time::{Duration, Instant};
use std::sync::{Arc, Mutex};
use std::io::{self, Cursor};
use blowfish::Blowfish;
use super::filter::{BlowfishReader, BlowfishWriter, blowfish::BLOCK_SIZE};
use super::packet::{Packet, RawPacket, PacketConfig, PacketConfigError};
use super::bundle::Bundle;
const FRAGMENT_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Clone)]
pub struct BundleSocket {
shared: Arc<Shared>,
}
struct Shared {
addr: SocketAddrV4,
socket: UdpSocket,
mutable: Mutex<SharedMutable>,
}
struct SharedMutable {
fragments: HashMap<(SocketAddr, u32), BundleFragments>,
next_sequence_num: u32,
channels: HashMap<SocketAddr, Channel>,
encryption_packet: Box<RawPacket>,
rejected_packets: Vec<(SocketAddr, Box<Packet>, PacketRejectionError)>,
}
impl BundleSocket {
pub fn new(addr: SocketAddrV4) -> io::Result<Self> {
let socket = UdpSocket::bind(SocketAddr::V4(addr))?;
Ok(Self {
shared: Arc::new(Shared {
addr,
socket,
mutable: Mutex::new(SharedMutable {
fragments: HashMap::new(),
next_sequence_num: 0,
channels: HashMap::new(),
encryption_packet: Box::new(RawPacket::new()),
rejected_packets: Vec::new(),
}),
}),
})
}
#[inline]
pub fn addr(&self) -> SocketAddrV4 {
self.shared.addr
}
pub fn set_channel(&mut self, addr: SocketAddr, blowfish: Arc<Blowfish>) {
self.shared.mutable.lock().unwrap()
.channels.insert(addr, Channel::new(blowfish));
}
pub fn send(&mut self, bundle: &mut Bundle, to: SocketAddr) -> io::Result<usize> {
if bundle.is_empty() {
return Ok(0)
}
let mut mutable = self.shared.mutable.lock().unwrap();
let SharedMutable {
next_sequence_num,
channels,
encryption_packet,
..
} = &mut *mutable;
let mut channel = channels.get_mut(&to);
let sequence_first_num = *next_sequence_num;
*next_sequence_num = next_sequence_num.checked_add(bundle.len() as u32).expect("sequence num overflow");
let sequence_last_num = *next_sequence_num - 1;
let mut packet_config = PacketConfig::new();
if let Some(channel) = channel.as_deref_mut() {
packet_config.set_on_channel(true);
packet_config.set_reliable(true);
if packet_config.cumulative_ack().is_some() {
channel.take_auto_ack();
}
}
if sequence_last_num > sequence_first_num {
packet_config.set_sequence_range(sequence_first_num, sequence_last_num);
}
let mut size = 0;
let mut sequence_num = sequence_first_num;
for packet in bundle.packets_mut() {
if sequence_num == sequence_last_num {
if let Some(channel) = channel.as_mut() {
if let Some(num) = channel.get_cumulative_ack_exclusive() {
packet_config.set_cumulative_ack(num);
} else {
packet_config.clear_cumulative_ack();
}
}
}
packet_config.set_sequence_num(sequence_num);
packet.write_config(&mut packet_config);
let raw_packet;
if let Some(channel) = channel.as_deref_mut() {
channel.add_sent_ack(sequence_num);
encrypt_packet(packet.raw(), &channel.blowfish, &mut **encryption_packet);
raw_packet = &**encryption_packet;
} else {
raw_packet = packet.raw()
}
size += self.shared.socket.send_to(raw_packet.data(), to)?;
sequence_num += 1;
}
Ok(size)
}
pub fn recv(&mut self) -> io::Result<Option<Bundle>> {
let mut packet = Packet::new_boxed();
let (len, addr) = self.shared.socket.recv_from(packet.raw_mut().raw_data_mut())?;
packet.raw_mut().set_data_len(len);
let mut mutable = self.shared.mutable.lock().unwrap();
let SharedMutable {
channels,
fragments,
rejected_packets,
..
} = &mut *mutable;
let mut channel = channels.get_mut(&addr);
if let Some(channel) = channel.as_deref_mut() {
match decrypt_packet(&packet, &channel.blowfish) {
Ok(clear_packet) => packet = clear_packet,
Err(()) => {
mutable.rejected_packets.push((addr, packet, PacketRejectionError::InvalidEncryption));
return Ok(None);
}
}
}
let len = packet.raw().data_len();
let mut packet_config = PacketConfig::new();
if let Err(error) = packet.read_config(len, &mut packet_config) {
mutable.rejected_packets.push((addr, packet, PacketRejectionError::Config(error)));
return Ok(None);
}
if let Some(channel) = channel.as_deref_mut() {
if packet_config.reliable() {
channel.add_received_ack(packet_config.sequence_num());
channel.set_auto_ack();
}
if let Some(ack) = packet_config.cumulative_ack() {
channel.remove_cumulative_ack(ack);
}
}
if packet_config.unk_1000().is_none() {
let instant = Instant::now();
match packet_config.sequence_range() {
Some((first_num, last_num)) if last_num > first_num => {
let num = packet_config.sequence_num();
match fragments.entry((addr, first_num)) {
hash_map::Entry::Occupied(mut o) => {
if o.get().is_old(instant, FRAGMENT_TIMEOUT) {
rejected_packets.extend(o.get_mut().drain()
.map(|packet| (addr, packet, PacketRejectionError::TimedOut)));
}
o.get_mut().set(num, packet);
if o.get().is_full() {
return Ok(Some(o.remove().into_bundle()));
}
},
hash_map::Entry::Vacant(v) => {
let mut fragments = BundleFragments::new(last_num - first_num + 1);
fragments.set(num, packet);
v.insert(fragments);
}
}
}
_ => {
return Ok(Some(Bundle::with_single(packet)));
}
}
}
Ok(None)
}
pub fn send_auto_ack(&mut self) {
let mut mutable = self.shared.mutable.lock().unwrap();
let SharedMutable {
channels,
encryption_packet,
..
} = &mut *mutable;
for (addr, channel) in channels {
if channel.take_auto_ack() {
let mut packet_config = PacketConfig::new();
let ack = channel.get_cumulative_ack_exclusive().expect("incoherent");
packet_config.set_sequence_num(ack);
packet_config.set_cumulative_ack(ack);
packet_config.set_on_channel(true);
packet_config.set_unk_1000(0);
let mut packet = Packet::new_boxed();
packet.write_config(&mut packet_config);
encrypt_packet(packet.raw(), &channel.blowfish, &mut **encryption_packet);
self.shared.socket.send_to(encryption_packet.data(), *addr).unwrap();
}
}
}
pub fn take_rejected_packets(&mut self) -> Vec<(SocketAddr, Box<Packet>, PacketRejectionError)> {
let mut mutable = self.shared.mutable.lock().unwrap();
let SharedMutable {
fragments,
rejected_packets,
..
} = &mut *mutable;
let instant = Instant::now();
fragments.retain(|(addr, _), fragments| {
if fragments.is_old(instant, FRAGMENT_TIMEOUT) {
rejected_packets.extend(fragments.drain()
.map(|packet| (*addr, packet, PacketRejectionError::TimedOut)));
false
} else {
true
}
});
std::mem::take(rejected_packets)
}
}
const ENCRYPTION_MAGIC: [u8; 4] = 0xDEADBEEFu32.to_le_bytes();
const ENCRYPTION_FOOTER_LEN: usize = ENCRYPTION_MAGIC.len() + 1;
fn decrypt_packet(packet: &Packet, bf: &Blowfish) -> Result<Box<Packet>, ()> {
let len = packet.raw().data_len();
let mut clear_packet = Packet::new_boxed();
clear_packet.raw_mut().set_data_len(len);
let src = packet.raw().body();
let dst = clear_packet.raw_mut().body_mut();
if src.len() % BLOCK_SIZE != 0 || src.len() < ENCRYPTION_FOOTER_LEN {
return Err(())
}
io::copy(
&mut BlowfishReader::new(Cursor::new(src), &bf),
&mut Cursor::new(&mut *dst),
).unwrap();
let wastage_begin = src.len() - 1;
let magic_begin = wastage_begin - 4;
if &dst[magic_begin..wastage_begin] != &ENCRYPTION_MAGIC {
return Err(())
}
let wastage = dst[wastage_begin];
assert!(wastage <= BLOCK_SIZE as u8, "temporary check that wastage is not greater than block size");
clear_packet.raw_mut().set_data_len(len - wastage as usize - ENCRYPTION_MAGIC.len());
clear_packet.raw_mut().write_prefix(packet.raw().read_prefix());
Ok(clear_packet)
}
fn encrypt_packet(src_packet: &RawPacket, bf: &Blowfish, dst_packet: &mut RawPacket) {
let mut len = src_packet.body_len() + ENCRYPTION_FOOTER_LEN;
let padding = (BLOCK_SIZE - (len % BLOCK_SIZE)) % BLOCK_SIZE;
len += padding;
let mut clear_data = Vec::from(src_packet.body());
clear_data.reserve_exact(padding + ENCRYPTION_FOOTER_LEN);
clear_data.extend_from_slice(&[0u8; BLOCK_SIZE - 1][..padding]); clear_data.extend_from_slice(&ENCRYPTION_MAGIC); clear_data.push(padding as u8 + 1);
debug_assert_eq!(clear_data.len(), len, "incoherent length");
debug_assert_eq!(clear_data.len() % 8, 0, "data not padded as expected");
dst_packet.set_data_len(clear_data.len() + 4);
io::copy(
&mut Cursor::new(&clear_data[..]),
&mut BlowfishWriter::new(Cursor::new(dst_packet.body_mut()), bf),
).unwrap();
dst_packet.write_prefix(src_packet.read_prefix());
}
#[derive(Debug)]
pub struct Channel {
blowfish: Arc<Blowfish>,
sent_acks: Vec<u32>,
received_acks: Vec<u32>,
auto_ack: bool,
}
impl Channel {
fn new(blowfish: Arc<Blowfish>) -> Self {
Self {
blowfish,
sent_acks: Vec::new(),
received_acks: Vec::new(),
auto_ack: false,
}
}
fn add_sent_ack(&mut self, sequence_num: u32) {
debug_assert!(
self.sent_acks.is_empty() || *self.sent_acks.last().unwrap() < sequence_num,
"sequence number is not ordered"
);
self.sent_acks.push(sequence_num);
println!("[AFTER ADD] sent_acks: {:?}", self.sent_acks);
}
fn remove_cumulative_ack(&mut self, ack: u32) {
let discard_offset = match self.sent_acks.binary_search(&ack) {
Ok(index) => index,
Err(index) => index,
};
self.sent_acks.drain(..discard_offset);
println!("[AFTER REM] sent_acks: {:?}", self.sent_acks);
}
fn add_received_ack(&mut self, sequence_num: u32) {
match self.received_acks.binary_search(&sequence_num) {
Ok(_) => {
}
Err(index) => {
self.received_acks.insert(index, sequence_num);
}
}
println!("[AFTER ADD] received_acks: {:?}", self.received_acks);
}
#[inline]
fn set_auto_ack(&mut self) {
self.auto_ack = true;
}
#[inline]
fn take_auto_ack(&mut self) -> bool {
std::mem::replace(&mut self.auto_ack, false)
}
fn get_cumulative_ack(&mut self) -> Option<u32> {
let first_ack = *self.received_acks.get(0)?;
let mut cumulative_ack = first_ack;
for &sequence_num in &self.received_acks[1..] {
if sequence_num == cumulative_ack + 1 {
cumulative_ack += 1;
} else {
break
}
}
if cumulative_ack > first_ack {
let diff = cumulative_ack - first_ack;
self.received_acks.drain(..diff as usize);
}
Some(cumulative_ack)
}
#[inline]
fn get_cumulative_ack_exclusive(&mut self) -> Option<u32> {
self.get_cumulative_ack().map(|n| n + 1)
}
}
struct BundleFragments {
fragments: Vec<Option<Box<Packet>>>, seq_count: u32,
last_update: Instant,
}
impl BundleFragments {
fn new(seq_len: u32) -> Self {
Self {
fragments: (0..seq_len).map(|_| None).collect(),
seq_count: 0,
last_update: Instant::now()
}
}
fn drain(&mut self) -> impl Iterator<Item = Box<Packet>> {
self.seq_count = 0;
std::mem::take(&mut self.fragments)
.into_iter()
.filter_map(|slot| slot)
}
fn set(&mut self, num: u32, packet: Box<Packet>) {
let frag = &mut self.fragments[num as usize];
if frag.is_none() {
self.seq_count += 1;
}
self.last_update = Instant::now();
*frag = Some(packet);
}
#[inline]
fn is_old(&self, instant: Instant, timeout: Duration) -> bool {
instant - self.last_update > timeout
}
#[inline]
fn is_full(&self) -> bool {
self.seq_count as usize == self.fragments.len()
}
#[inline]
fn into_bundle(self) -> Bundle {
assert!(self.is_full());
let packets = self.fragments.into_iter()
.map(|o| o.unwrap())
.collect();
Bundle::with_multiple(packets)
}
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum PacketRejectionError {
#[error("timed out")]
TimedOut,
#[error("invalid encryption")]
InvalidEncryption,
#[error("sync error: {0}")]
Config(#[from] PacketConfigError),
}