use std::{io::ErrorKind, net::UdpSocket};
use aeronet_io::connection::{LocalAddr, PeerAddr};
use bevy_app::prelude::*;
use bevy_ecs::prelude::*;
use bytes::{BufMut, BytesMut};
use lightyear_core::time::Instant;
use lightyear_link::{
Link, LinkPlugin, LinkReceiveSystems, LinkStart, LinkSystems, Linked, Linking, Unlink, Unlinked,
};
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;
#[derive(Component)]
#[require(Link)]
pub struct UdpIo {
socket: Option<UdpSocket>,
buffer: BytesMut,
}
impl Default for UdpIo {
fn default() -> Self {
UdpIo {
socket: None,
buffer: BytesMut::with_capacity(MTU),
}
}
}
#[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>>) {
query
.par_iter_mut()
.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>>) {
query.par_iter_mut().for_each(|(mut link, mut udp_io)| {
let udp_io = &mut *udp_io;
loop {
udp_io.buffer.reserve(MTU);
let capacity = udp_io.buffer.capacity();
let current_len = udp_io.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 = udp_io.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 {
udp_io.buffer.advance_mut(recv_len);
}
let payload = udp_io.buffer.split_to(recv_len);
link.recv.push(payload.freeze(), Instant::now());
}
Err(ref e) if e.kind() == ErrorKind::WouldBlock => return,
Err(e) => {
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(PreUpdate, Self::send.in_set(LinkSystems::Send));
}
}