#![allow(clippy::std_instead_of_core)]
use std::io::ErrorKind;
use std::net::UdpSocket;
use aeronet_io::connection::{LocalAddr, PeerAddr};
use bevy_app::prelude::*;
use bevy_ecs::prelude::*;
use bytes::BufMut;
use lightyear_core::buffer_pool::BufferPool;
use lightyear_core::time::Instant;
use lightyear_link::{
Link, LinkPlugin, LinkReceiveSystems, LinkStart, LinkSystems, Linked, Linking, Unlink, Unlinked,
};
use lightyear_utils::adaptive_for_each_mut;
use tracing::{error, info, trace};
#[cfg(feature = "server")]
pub mod server;
pub mod prelude {
pub use crate::UdpIo;
#[cfg(feature = "server")]
pub mod server {
pub use crate::server::ServerUdpIo;
}
}
pub(crate) const MTU: usize = 1472;
const MAX_RETAINED_RECV_BUFFERS: usize = 64;
fn recv_buffer_pool() -> BufferPool {
let mut pool = BufferPool::new(MTU, MAX_RETAINED_RECV_BUFFERS);
pool.preallocate(1);
pool
}
#[derive(Component)]
#[require(Link)]
pub struct UdpIo {
socket: Option<UdpSocket>,
recv_buffers: BufferPool,
}
impl Default for UdpIo {
fn default() -> Self {
Self {
socket: None,
recv_buffers: recv_buffer_pool(),
}
}
}
impl UdpIo {
#[cfg(feature = "test_utils")]
pub fn recv_buffer_pool_misses(&self) -> usize {
self.recv_buffers.misses()
}
}
#[derive(thiserror::Error, Debug)]
pub enum UdpError {
#[error("LocalAddr is required to start the UdpIo link")]
LocalAddrMissing,
}
pub struct UdpPlugin;
impl UdpPlugin {
fn link(
trigger: On<LinkStart>,
mut query: Query<(&mut UdpIo, Option<&LocalAddr>), (Without<Linking>, Without<Linked>)>,
mut commands: Commands,
) -> Result {
trace!("In LinkStart::UDP trigger");
if let Ok((mut udp_io, local_addr)) = query.get_mut(trigger.entity) {
let local_addr = local_addr.ok_or(UdpError::LocalAddrMissing)?.0;
let socket = UdpSocket::bind(local_addr)?;
info!("UDP socket bound to {}", local_addr);
socket.set_nonblocking(true)?;
udp_io.socket = Some(socket);
commands.entity(trigger.entity).insert(Linked);
}
Ok(())
}
fn unlink(trigger: On<Unlink>, mut query: Query<&mut UdpIo, Without<Unlinked>>) {
if let Ok(mut udp_io) = query.get_mut(trigger.entity) {
info!("UDP socket closed");
udp_io.socket = None;
}
}
fn send(mut query: Query<(&mut Link, &mut UdpIo, &PeerAddr), With<Linked>>) {
adaptive_for_each_mut!(query).for_each(|(mut link, mut udp_io, remote_addr)| {
link.send.drain().for_each(|payload| {
#[cfg(feature = "metrics")]
metrics::gauge!("udp/send").increment(payload.len() as f64);
udp_io
.socket
.as_mut()
.unwrap()
.send_to(payload.as_ref(), remote_addr.0)
.inspect_err(|e| error!("Error sending UDP packet: {}", e))
.ok();
});
})
}
fn receive(mut query: Query<(&mut Link, &mut UdpIo), With<Linked>>) {
adaptive_for_each_mut!(query).for_each(|(mut link, mut udp_io)| {
let udp_io = &mut *udp_io;
udp_io.recv_buffers.reclaim_pending();
loop {
let mut buffer = udp_io.recv_buffers.take();
let capacity = buffer.capacity();
let current_len = buffer.len();
assert_eq!(current_len, 0);
let available_uninit = capacity - current_len;
let max_recv_len = core::cmp::min(available_uninit, MTU);
let buf_slice: &mut [u8] = unsafe {
let ptr = buffer.as_mut_ptr().add(current_len);
core::slice::from_raw_parts_mut(ptr, max_recv_len)
};
match udp_io.socket.as_mut().unwrap().recv_from(buf_slice) {
Ok((recv_len, _)) => {
unsafe {
buffer.advance_mut(recv_len);
}
let payload = udp_io.recv_buffers.split_for_handoff(buffer);
link.recv.push(payload, Instant::now());
}
Err(ref e) if e.kind() == ErrorKind::WouldBlock => {
udp_io.recv_buffers.recycle(buffer);
return;
}
Err(ref e) if e.kind() == ErrorKind::ConnectionReset => {
udp_io.recv_buffers.recycle(buffer);
continue;
}
Err(e) => {
udp_io.recv_buffers.recycle(buffer);
error!("Error receiving UDP packet: {}", e);
return;
}
}
}
})
}
}
impl Plugin for UdpPlugin {
fn build(&self, app: &mut App) {
if !app.is_plugin_added::<LinkPlugin>() {
app.add_plugins(LinkPlugin);
}
app.add_observer(Self::link);
app.add_observer(Self::unlink);
app.add_systems(
PreUpdate,
Self::receive.in_set(LinkReceiveSystems::BufferToLink),
);
app.add_systems(PostUpdate, Self::send.in_set(LinkSystems::Send));
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn receive_buffer_pool_reclaims_a_published_datagram_after_drop() {
let mut pool = recv_buffer_pool();
let mut buffer = pool.take();
buffer.extend_from_slice(b"datagram");
let payload = pool.split_for_handoff(buffer);
let misses = pool.misses();
pool.reclaim_pending();
let in_flight_fallback = pool.take();
assert_eq!(pool.misses(), misses + 1);
pool.recycle(in_flight_fallback);
drop(payload);
pool.reclaim_pending();
let misses = pool.misses();
assert!(pool.take().capacity() >= MTU);
assert_eq!(pool.misses(), misses);
}
}