extern crate alloc;
use bevy_app::{App, Plugin, PostUpdate, PreUpdate};
use bevy_ecs::prelude::*;
use bevy_ecs::relationship::RelationshipTarget;
use bevy_ecs::system::ParallelCommands;
use tracing::{debug, error, info};
use crate::UdpError;
use aeronet_io::connection::{LocalAddr, PeerAddr};
use bevy_platform::collections::{HashMap, hash_map::Entry};
use bytes::BufMut;
use core::net::SocketAddr;
use lightyear_core::buffer_pool::BufferPool;
use lightyear_core::time::Instant;
use lightyear_link::prelude::{LinkOf, Server};
use lightyear_link::{Link, LinkPlugin, LinkStart, LinkSystems, Linked, Linking, Unlink, Unlinked};
pub(crate) const MTU: usize = 1472;
#[derive(Component)]
#[require(Server)]
pub struct ServerUdpIo {
socket: Option<std::net::UdpSocket>,
recv_buffers: BufferPool,
connected_addresses: HashMap<SocketAddr, LinkOfStatus>,
}
#[derive(Component)]
pub struct UdpLinkOfIO;
#[derive(Debug)]
enum LinkOfStatus {
Spawning(Entity),
Spawned(Entity),
}
impl Default for ServerUdpIo {
fn default() -> Self {
ServerUdpIo {
socket: None,
recv_buffers: crate::recv_buffer_pool(),
connected_addresses: HashMap::with_capacity(1),
}
}
}
impl ServerUdpIo {
#[cfg(feature = "test_utils")]
pub fn recv_buffer_pool_misses(&self) -> usize {
self.recv_buffers.misses()
}
}
pub struct ServerUdpPlugin;
impl ServerUdpPlugin {
fn link(
trigger: On<LinkStart>,
mut query: Query<
(&mut ServerUdpIo, Option<&mut LocalAddr>),
(Without<Linking>, Without<Linked>),
>,
mut commands: Commands,
) -> Result {
if let Ok((mut udp_io, local_addr)) = query.get_mut(trigger.entity) {
let mut local_addr = local_addr.ok_or(UdpError::LocalAddrMissing)?;
let socket = std::net::UdpSocket::bind(local_addr.0)?;
socket.set_nonblocking(true)?;
local_addr.0 = socket.local_addr()?;
info!("Server UDP socket bound to {}", local_addr.0);
udp_io.socket = Some(socket);
commands.entity(trigger.entity).insert(Linked);
}
Ok(())
}
fn unlink(trigger: On<Unlink>, mut query: Query<&mut ServerUdpIo, Without<Unlinked>>) {
if let Ok(mut udp_io) = query.get_mut(trigger.entity) {
info!("Server UDP socket closed");
udp_io.socket = None;
}
}
fn send(
mut server_query: Query<(&mut ServerUdpIo, &Server), With<Linked>>,
mut link_query: Query<(&mut Link, &PeerAddr), With<UdpLinkOfIO>>,
) {
server_query
.iter_mut()
.for_each(|(mut server_udp_io, server)| {
server.collection().iter().for_each(|client_entity| {
let Some((mut link, remote_addr)) = link_query.get_mut(*client_entity).ok()
else {
debug!("Client entity {} not found in link query", client_entity);
return;
};
link.send.drain().for_each(|send_payload| {
server_udp_io
.socket
.as_mut()
.unwrap()
.send_to(send_payload.as_ref(), remote_addr.0)
.inspect_err(|e| {
error!("Error sending UDP packet to {}: {}", remote_addr.0, e);
})
.ok();
});
});
});
}
fn receive(
commands: ParallelCommands,
mut server_query: Query<(Entity, &mut ServerUdpIo), With<Linked>>,
link_query: Query<Option<&mut Link>>,
) {
server_query
.iter_mut()
.for_each(|(server_entity, mut server_udp_io)| {
let mut link_query = unsafe { link_query.reborrow_unsafe() };
let server_udp_io = &mut *server_udp_io;
server_udp_io.recv_buffers.reclaim_pending();
loop {
let mut buffer = server_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, crate::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 server_udp_io.socket.as_mut().unwrap().recv_from(buf_slice) {
Ok((recv_len, address)) => {
unsafe {
buffer.advance_mut(recv_len);
}
let payload = server_udp_io.recv_buffers.split_for_handoff(buffer);
match server_udp_io.connected_addresses.entry(address) {
Entry::Occupied(mut entry) => {
match *entry.get_mut() {
LinkOfStatus::Spawning(_) => {
continue;
}
LinkOfStatus::Spawned(entity) => {
match link_query.get_mut(entity) {
Ok(mut link) => {
match link.as_mut() {
None => {
debug!("despawning entity {} because it has no udp link", entity);
entry.remove();
commands.command_scope(|mut c| {
if let Ok(mut e) = c.get_entity(entity) {
e.try_despawn();
}
});
}
Some(link) => {
link.recv.push(payload, Instant::now());
}
}
}
Err(_) => {
error!(
"Received UDP packet for unknown entity: {}",
entity
);
entry.remove();
continue;
}
}
}
}
}
Entry::Vacant(vacant) => {
let mut link = Link::default();
link.recv.push(payload, Instant::now());
commands.command_scope(|mut c| {
let entity = c
.spawn((
LinkOf {
server: server_entity,
},
link,
Linked,
PeerAddr(address),
UdpLinkOfIO,
))
.id();
info!(?entity, ?server_entity, "Received UDP packet from new address {address}, Spawn new LinkOf");
vacant.insert(LinkOfStatus::Spawning(entity));
});
continue;
}
};
}
Err(ref e) if e.kind() == std::io::ErrorKind::WouldBlock => {
server_udp_io.recv_buffers.recycle(buffer);
break;
}
Err(ref e) if e.kind() == std::io::ErrorKind::ConnectionReset => {
server_udp_io.recv_buffers.recycle(buffer);
continue;
}
Err(e) => {
server_udp_io.recv_buffers.recycle(buffer);
error!("Error receiving UDP packet: {}", e);
break;
}
}
}
server_udp_io.connected_addresses.iter_mut().for_each(|(addr, status)| {
if let LinkOfStatus::Spawning(entity) = status {
*status = LinkOfStatus::Spawned(*entity);
}
});
});
}
}
impl Plugin for ServerUdpPlugin {
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(LinkSystems::Receive));
app.add_systems(PostUpdate, Self::send.in_set(LinkSystems::Send));
}
}
#[cfg(test)]
mod tests {
use super::*;
use core::net::Ipv4Addr;
#[test]
fn link_updates_local_addr_with_os_assigned_port() {
let mut app = App::new();
app.add_plugins(ServerUdpPlugin);
let requested_addr = SocketAddr::new(Ipv4Addr::LOCALHOST.into(), 0);
let server = app
.world_mut()
.spawn((LocalAddr(requested_addr), ServerUdpIo::default()))
.id();
app.world_mut().trigger(LinkStart { entity: server });
app.world_mut().flush();
let bound_addr = app.world().get::<LocalAddr>(server).unwrap().0;
assert_eq!(bound_addr.ip(), requested_addr.ip());
assert_ne!(bound_addr.port(), 0);
assert!(app.world().get::<Linked>(server).is_some());
}
}