const TICK_PERIOD: std::time::Duration = std::time::Duration::from_millis(250);
pub struct Server<D: literustlib::packet::PacketData + 'static, U: Send + Sync, H: super::EventHandler<PacketData=D, UserData=U>> {
connections: std::collections::HashMap<core::net::SocketAddr, (std::sync::Arc<super::Connection<D>>, std::sync::Arc<U>)>,
handler: H,
socket: std::sync::Arc<tokio::net::UdpSocket>,
sender: std::sync::Arc<super::DataSender<D>>,
window: u16,
filter: super::BadnessFilter,
conf: std::sync::Arc<super::ServerConfig>,
}
impl <D: literustlib::packet::PacketData + 'static, U: Send + Sync, H: super::EventHandler<PacketData=D, UserData=U>> Server<D, U, H> {
pub async fn new<A: tokio::net::ToSocketAddrs>(handler: H, addr: A, config: super::ServerConfig) -> std::io::Result<Self> {
let socket = std::sync::Arc::new(tokio::net::UdpSocket::bind(addr).await?);
Ok(Self {
connections: std::collections::HashMap::new(),
handler,
sender: std::sync::Arc::new(super::DataSender::new(config.mtu)),
socket,
window: 64,
filter: super::BadnessFilter::new(config.dos_protection),
conf: std::sync::Arc::new(config),
})
}
pub async fn listen(mut self) -> ! {
let mut last_tick = std::time::SystemTime::now();
let mut udp_buf = [0u8; u16::MAX as _];
loop {
let now = std::time::SystemTime::now();
if let Ok(duration) = now.duration_since(last_tick) {
if duration >= TICK_PERIOD {
last_tick = now;
self.filter.tick();
self.handle_disconnects().await;
}
}
match self.socket.recv_from(&mut udp_buf).await {
Err(e) => log::error!("Error on UDP receive: {}", e),
Ok((size, from_addr)) => {
if self.filter.is_bad(from_addr) {
log::trace!("Ignoring bad packet from {}", from_addr);
continue;
}
let valid_buf = &udp_buf[0..size];
self.handle_packet(valid_buf, from_addr).await;
}
}
}
}
fn is_acked(data: &[u8], pos: usize) -> bool {
data.len() > pos / 8 && data[pos / 8] & (1 << pos % 8) != 0
}
async fn handle_disconnects(&mut self) {
let mut to_remove = std::collections::HashSet::new();
for (key, (c_data, _)) in self.connections.iter() {
if !c_data.is_connected.load(literustlib::serdes::ATOMIC_ORDERING) {
to_remove.insert(*key);
}
}
for key in to_remove {
if let Some((c_data, u_data)) = self.connections.remove(&key) {
self.handler.on_disconnect(&c_data, &u_data).await;
if c_data.is_certified() {
self.filter.reset(c_data.addr);
}
}
}
}
async fn handle_acks(c_data: &super::Connection<D>, sender: &super::DataSender<D>, packet: &literustlib::packet::Packet, max_bulk_resends: usize) -> bool {
if let Some(seq) = packet.header.sequence {
let channel = c_data.channel_by_ack_prop(packet.header.property);
let window_start = channel.window_start.load(literustlib::serdes::ATOMIC_ORDERING);
if seq > 32768 {
log::debug!("Dropping {:?} packet due to unacceptable sequence number", packet.header.property);
return false; }
let mut last_contiguous_acked = window_start;
let mut highest_acked = 0;
for i in 0..channel.window_size {
let window_pos = (seq + i) % 32768;
let is_acked = Self::is_acked(&packet.data, i as usize);
if is_acked && last_contiguous_acked == window_pos {
last_contiguous_acked = (last_contiguous_acked + 1) % 32768;
}
if is_acked {
highest_acked = i;
}
}
let mut pending_lock = channel.pending.lock().await;
pending_lock.retain(|p| {
if let Some(p_seq) = p.header.sequence {
let window_offset = if p_seq < seq {
(32768 - seq) + p_seq
} else {
p_seq - seq
};
!Self::is_acked(&packet.data, window_offset as usize)
} else {
false }
});
let mut resend_count = 0;
for pending_packet in pending_lock.iter() {
let p_seq = pending_packet.header.sequence.unwrap();
let window_offset = if p_seq < seq {
(32768 - seq) + p_seq
} else {
p_seq - seq
};
if window_offset < highest_acked {
if let Err(e) = sender.raw_send_to(&pending_packet, c_data).await {
log::error!("Failed to resend NAck packet #{} {}", p_seq, e);
return false;
} else {
resend_count += 1;
if resend_count > max_bulk_resends { break; }
}
}
}
log::trace!("Now contiguously {:?}-ed up to {} (exclusive)", packet.header.property, last_contiguous_acked);
channel.window_start.store(last_contiguous_acked, literustlib::serdes::ATOMIC_ORDERING);
drop(pending_lock);
return last_contiguous_acked != window_start;
} else {
log::warn!("Got {:?} packet without sequence number!", packet.header.property);
return false;
}
}
async fn handle_reliable_tracking(c_data: &super::Connection<D>, sender: &super::DataSender<D>, header: &literustlib::packet::Header) -> bool {
let seq = header.sequence.unwrap();
let channel = c_data.channel_by_prop(header.property).unwrap();
let remote_window_start = channel.remote_window_start.load(literustlib::serdes::ATOMIC_ORDERING);
let window_offset = if seq < remote_window_start {
let offset = (32768 - remote_window_start) + seq;
if offset > channel.window_size {
log::trace!("Ignoring old {:?} packet", header.property);
let window_bytes = ((channel.window_size - 1) / 8) + 1;
let ack_prop = header.property.ack_equivalent().unwrap();
let relevant_seen_bytes = &channel.seen.lock().await[0..window_bytes as usize];
let mut packet = literustlib::packet::Packet::with_data(ack_prop, relevant_seen_bytes);
packet.header.sequence = Some(remote_window_start);
if let Err(e) = sender.raw_send_to(&packet, c_data).await {
log::error!("Failed to re-send {:?} for reliable packet: {}", ack_prop, e);
}
return false;
} else {
log::trace!("Packet window is wrapping; remote_window_start: {} seq:{} window_offset:{}", remote_window_start, seq, offset);
offset
}
} else {
seq - remote_window_start
};
let mut seen_lock = channel.seen.lock().await;
let pos = (window_offset / 8) as usize;
if pos > seen_lock.len() {
log::error!("Got {:?} packet with a sequence number too high for the ack/seen buffer!", header.property);
return false;
}
let offset = window_offset % 8;
seen_lock[pos] |= 1 << offset;
let window_bytes = ((channel.window_size - 1) / 8) + 1;
let ack_prop = header.property.ack_equivalent().unwrap();
let relevant_seen_bytes = &seen_lock[0..window_bytes as usize];
let mut packet = literustlib::packet::Packet::with_data(ack_prop, relevant_seen_bytes);
packet.header.sequence = Some(remote_window_start);
if let Err(e) = sender.raw_send_to(&packet, c_data).await {
log::error!("Failed to send {:?} for reliable packet: {}", ack_prop, e);
false
} else {
c_data.ping_channel.remote_sequence.store(seq, literustlib::serdes::ATOMIC_ORDERING);
let channel = c_data.channel_by_prop(header.property).unwrap();
let moved_len = super::Channel::move_seen_window(&mut seen_lock);
let new_remote_window_start = (remote_window_start + moved_len as u16) % 32768;
channel.remote_window_start.store(new_remote_window_start, literustlib::serdes::ATOMIC_ORDERING);
true
}
}
#[async_recursion::async_recursion]
async fn handle_packet(&mut self, slice: &[u8], from: core::net::SocketAddr) {
log::trace!("Got packet ({}) {:?}", slice.len(), slice);
match literustlib::packet::Packet::parse(slice) {
Ok(packet) => {
if packet.header.property == literustlib::packet::Property::Dummy {
return;
}
if let Some((c_data, u_data)) = self.connections.get(&from) {
c_data.last_seen_now().await;
match packet.header.property {
literustlib::packet::Property::Unreliable => {
Self::handle_user_packet(&self.handler, packet, c_data, u_data, &self.sender).await;
},
literustlib::packet::Property::Reliable => {
if packet.header.sequence.is_some() {
if Self::handle_reliable_tracking(c_data, &self.sender, &packet.header).await {
Self::handle_user_packet(&self.handler, packet, c_data, u_data, &self.sender).await;
}
} else {
log::warn!("Got Reliable packet without sequence number!");
}
},
literustlib::packet::Property::Sequenced => {
if let Some(seq) = packet.header.sequence {
c_data.ping_channel.remote_sequence.store(seq, literustlib::serdes::ATOMIC_ORDERING);
Self::handle_user_packet(&self.handler, packet, c_data, u_data, &self.sender).await;
} else {
log::warn!("Got Sequenced packet without sequence number!");
}
},
literustlib::packet::Property::ReliableOrdered => {
if packet.header.sequence.is_some() {
if Self::handle_reliable_tracking(c_data, &self.sender, &packet.header).await {
Self::handle_user_packet(&self.handler, packet, c_data, u_data, &self.sender).await;
}
} else {
log::warn!("Got ReliableOrdered packet without sequence number!");
}
},
literustlib::packet::Property::AckReliable => {
Self::handle_acks(c_data, &self.sender, &packet, self.conf.max_bulk_resends).await;
},
literustlib::packet::Property::AckReliableOrdered => {
Self::handle_acks(c_data, &self.sender, &packet, self.conf.max_bulk_resends).await;
},
literustlib::packet::Property::Ping => {
if let Some(remote_seq) = packet.header.sequence {
let old_seq = c_data.ping_channel.remote_sequence.swap(remote_seq, literustlib::serdes::ATOMIC_ORDERING);
let packet = literustlib::packet::Packet::without_data(literustlib::packet::Header {
property: literustlib::packet::Property::Pong,
sequence: Some(remote_seq),
fragment: None,
});
if let Err(e) = self.sender.raw_send_to(&packet, c_data).await {
log::error!("Failed to send pong: {}", e);
}
let local_seq = c_data.ping_channel.local_sequence.fetch_add(1, literustlib::serdes::ATOMIC_ORDERING);
let mut ping_time_lock = c_data.last_ping_time.lock().await;
if ping_time_lock.is_none() || old_seq.wrapping_sub(4) > local_seq {
let packet = literustlib::packet::Packet::without_data(literustlib::packet::Header {
property: literustlib::packet::Property::Ping,
sequence: Some(local_seq),
fragment: None,
});
if let Err(e) = self.sender.raw_send_to(&packet, c_data).await {
log::error!("Failed to send ping from server: {}", e);
} else {
*ping_time_lock = Some(chrono::Utc::now());
}
}
} else {
log::warn!("Got ping packet without sequence number!");
}
},
literustlib::packet::Property::Pong => {
if let Some(seq) = packet.header.sequence {
let expected_seq = c_data.ping_channel.local_sequence.load(literustlib::serdes::ATOMIC_ORDERING).wrapping_sub(1);
if expected_seq != seq {
log::warn!("Round trip time for id:{} may be unreliable, missed a ping", c_data.id);
} else {
let mut ping_time_lock = c_data.last_ping_time.lock().await;
if let Some(last_ping_time) = ping_time_lock.take() {
let now = chrono::Utc::now();
let ping_dur = now.signed_duration_since(last_ping_time);
let nanos_dur = ping_dur.num_nanoseconds().unwrap_or_default() as u64;
c_data.round_trip.store(nanos_dur, literustlib::serdes::ATOMIC_ORDERING);
log::trace!("Round trip time for id:{}: {}ns", c_data.id, nanos_dur);
}
}
} else {
log::warn!("Got pong packet without sequence number!");
}
},
literustlib::packet::Property::Dummy => {
unreachable!("dummy packet should've been ignored already");
},
literustlib::packet::Property::Disconnect => {
log::info!("Disconnect from {}", c_data.id);
if let Some(old_conn) = self.connections.get(&from) {
old_conn.0.disconnect();
}
},
literustlib::packet::Property::Merged => {
log::trace!("Handling merged packet with data len {}", packet.data.len());
let mut start_i = 0;
while start_i + 1 < packet.data.len() {
let size = u16::from_le_bytes([packet.data[start_i], packet.data[start_i + 1]]);
start_i += 2;
let end = start_i + size as usize;
if end > packet.data.len() {
log::warn!("Merged packet tried to read more than the packet");
break;
}
self.handle_packet(&packet.data[start_i..end], from).await;
start_i = end;
}
},
literustlib::packet::Property::StateUpdate => {
if let Some(seq) = packet.header.sequence {
c_data.ping_channel.remote_sequence.store(seq, literustlib::serdes::ATOMIC_ORDERING);
log::warn!("Got unsupported StateUpdate packet");
} else {
log::warn!("Got StateUpdate packet without sequence number!");
}
},
literustlib::packet::Property::Auth => {
log::warn!("Ignoring Auth packet");
},
literustlib::packet::Property::AckAuth => {
if Self::handle_acks(c_data, &self.sender, &packet, self.conf.max_bulk_resends).await {
self.handler.on_connect_done(c_data, u_data, &self.sender).await;
} else {
}
},
literustlib::packet::Property::ConnectRequest => {
log::trace!("Got duplicate ConnectRequest from {}", c_data.id);
},
prop => {
#[cfg(debug_assertions)]
panic!("Unsupported packet property {:?}", prop);
#[cfg(not(debug_assertions))]
{
log::warn!("Unsupported packet property {:?}, disconnecting", prop);
let packet = literustlib::packet::Packet::without_data(
literustlib::packet::Header::with_prop(literustlib::packet::Property::Disconnect)
);
if let Err(e) = self.sender.raw_send_to(&packet, c_data).await {
log::error!("Failed to send disconnect packet for unsupported packet property: {}", e);
}
}
}
}
} else {
match packet.header.property {
literustlib::packet::Property::ConnectRequest => {
self.on_connect_start_wrapper(&from, &packet.data).await;
},
literustlib::packet::Property::Pong | literustlib::packet::Property::Ping => {
log::debug!("Ignoring unconnected packet {:?}, dropping connection", packet.header.property);
let packet = literustlib::packet::Packet::without_data(
literustlib::packet::Header::with_prop(literustlib::packet::Property::Disconnect)
);
if let Err(e) = self.send_unconnected(&packet, from).await {
log::error!("Failed to send disconnect packet for invalid first packet: {}", e);
}
},
_ => {
log::debug!("Got incorrect first/unconnected packet {:?}, dropping connection", packet.header.property);
self.filter.strike(from);
let packet = literustlib::packet::Packet::without_data(
literustlib::packet::Header::with_prop(literustlib::packet::Property::Disconnect)
);
if let Err(e) = self.send_unconnected(&packet, from).await {
log::error!("Failed to send disconnect packet for invalid first packet: {}", e);
}
}
}
}
},
Err(e) => {
self.filter.strike(from);
log::error!("Failed to parse packet: {}", e);
},
}
}
async fn handle_user_packet(handler: &H, packet: literustlib::packet::Packet, c_data: &std::sync::Arc<super::Connection<D>>, u_data: &U, sender: &std::sync::Arc<super::DataSender<D>>) {
let mut header = packet.header.clone();
match c_data.pool.lock().await.handle_packet(packet) {
Ok(Some(packet_data)) => {
header.fragment = None;
Self::on_receive_wrapper(handler, packet_data, &header, c_data, u_data, sender).await;
},
Ok(None) => {},
Err(e) => log::error!("Failed to pool packet: {}", e),
}
}
async fn on_receive_wrapper(handler: &H, p_data: D, header: &literustlib::packet::Header, c_data: &std::sync::Arc<super::Connection<D>>, u_data: &U, sender: &std::sync::Arc<super::DataSender<D>>) {
handler.on_receive(p_data, header, c_data, u_data, sender).await;
}
async fn on_connect_start_wrapper(&mut self, addr: &core::net::SocketAddr, data: &bytes::Bytes) {
let connect_id = if data.len() > 12 {
i64::from_be_bytes([
data[4],
data[5],
data[6],
data[7],
data[8],
data[9],
data[10],
data[11],
])
} else {
self.filter.strike(*addr);
return;
};
let connect_key = if data.len() > 16 {
match String::from_utf8(Vec::from(&data[16..])) {
Ok(key) => key,
Err(_e) => {
self.filter.strike(*addr);
return;
}
}
} else {
self.filter.strike(*addr);
return;
};
let c_data = std::sync::Arc::new(super::Connection::new(connect_id, *addr, self.socket.clone(), self.window));
if let Some(u_data) = self.handler.on_connect_start(addr, connect_key, &c_data).await {
let u_data = std::sync::Arc::new(u_data);
self.connections.insert(addr.to_owned(), (c_data.clone(), u_data));
if let Err(e) = self.accept_connection(addr).await {
log::error!("Failed to accept connection: {}", e);
} else {
log::debug!("Hello {} ({})", addr, connect_id);
tokio::spawn(background_process(c_data, self.sender.clone(), self.conf.clone()));
}
} else {
self.filter.strike(*addr);
}
}
async fn accept_connection(&mut self, addr: &core::net::SocketAddr) -> std::io::Result<()> {
let (c_data, _) = self.connections.get(addr).unwrap();
let connect_id = c_data.id.to_be_bytes();
let payload = [
1,
connect_id[0],
connect_id[1],
connect_id[2],
connect_id[3],
connect_id[4],
connect_id[5],
connect_id[6],
connect_id[7],
];
self.sender.send_to(bytes::Bytes::copy_from_slice(&payload), literustlib::packet::Property::Auth, c_data).await?;
Ok(())
}
async fn send_unconnected(&self, packet: &literustlib::packet::Packet, addr: core::net::SocketAddr) -> std::io::Result<()> {
let mut packet_bytes = Vec::new();
packet.dump(&mut std::io::Cursor::new(&mut packet_bytes))?;
self.socket.send_to(&packet_bytes, addr).await?;
Ok(())
}
}
async fn send_oldest_unacked_packets<D: literustlib::packet::PacketData + 'static>(c_data: &std::sync::Arc<super::Connection<D>>, sender: &std::sync::Arc<super::DataSender<D>>, property: literustlib::packet::Property, max_bulk_resends: usize) {
let channel = c_data.channel_by_prop(property).unwrap();
let pending_lock = channel.pending.lock().await;
let window_start = channel.window_start.load(literustlib::serdes::ATOMIC_ORDERING);
let mut resend_count = 0;
for packet in pending_lock.iter() {
let p_seq = packet.header.sequence.unwrap();
let window_offset = if p_seq < window_start {
(32768 - window_start) + p_seq
} else {
p_seq - window_start
};
if window_offset < channel.window_size {
if let Err(e) = sender.raw_send_to(&packet, c_data).await {
log::error!("Failed to resend packet #{} {}", p_seq, e);
} else {
resend_count += 1;
if resend_count > max_bulk_resends { break; }
}
}
}
}
async fn background_process<D: literustlib::packet::PacketData + 'static>(connection: std::sync::Arc<super::Connection<D>>, sender: std::sync::Arc<super::DataSender<D>>, conf: std::sync::Arc<super::ServerConfig>) {
const RELIABLE_PROPS: [literustlib::packet::Property; 3] = [
literustlib::packet::Property::Reliable,
literustlib::packet::Property::ReliableOrdered,
literustlib::packet::Property::Auth,
];
let max_bulk_resends = conf.max_bulk_resends;
let timeout_dur = chrono::TimeDelta::from_std(conf.timeout)
.unwrap_or_else(|_| {
log::warn!("Failed to convert ServerConfig.timeout's std::time::Duration to chrono::TimeDelta");
chrono::TimeDelta::zero()
});
let sleep_dur_mult = conf.resend_delay_rtt_mult;
let sleep_dur_base = std::time::Duration::from_secs_f64(conf.resend_delay_base);
log::debug!("Started background packet task for {}", connection.id);
while connection.is_connected.load(literustlib::serdes::ATOMIC_ORDERING) {
for prop in RELIABLE_PROPS.iter() {
send_oldest_unacked_packets(&connection, &sender, *prop, max_bulk_resends).await;
}
if !timeout_dur.is_zero() {
let last_seen = connection.last_seen_delta().await;
if last_seen > timeout_dur {
connection.goodbye(&sender).await;
}
}
let rtt = connection.round_trip.load(literustlib::serdes::ATOMIC_ORDERING);
let rtt_duration = std::time::Duration::from_nanos((rtt as f64 * sleep_dur_mult).ceil() as u64);
let sleep_dur = sleep_dur_base + rtt_duration;
tokio::time::sleep(sleep_dur).await;
}
let packet = literustlib::packet::Packet::without_data(
literustlib::packet::Header::with_prop(literustlib::packet::Property::Disconnect)
);
if let Err(e) = sender.raw_send_to(&packet, &connection).await {
log::error!("Failed send disconnect packet for {}: {}", connection.id, e);
}
}