use {
crate::convert,
aeronet_io::{
Session,
server::{Server, ServerEndpoint},
},
aeronet_transport::{
AeronetTransportPlugin, Transport, TransportSystems,
sampling::{SessionSamplingPlugin, SessionStats, SessionStatsSampling},
},
bevy_app::prelude::*,
bevy_ecs::prelude::*,
bevy_platform::time::Instant,
bevy_reflect::Reflect,
bevy_replicon::{prelude::*, server::ServerSystems},
bevy_state::state::{NextState, State},
core::num::Saturating,
log::{trace, warn},
};
pub struct AeronetRepliconServerPlugin;
impl Plugin for AeronetRepliconServerPlugin {
fn build(&self, app: &mut App) {
if !app.is_plugin_added::<AeronetTransportPlugin>() {
app.add_plugins(AeronetTransportPlugin);
}
if !app.is_plugin_added::<SessionSamplingPlugin>() {
app.add_plugins(SessionSamplingPlugin);
}
app.configure_sets(
PreUpdate,
(
TransportSystems::Poll,
ServerTransportSystems::Poll,
ServerSystems::ReceivePackets,
)
.chain(),
)
.configure_sets(
PostUpdate,
(
ServerSystems::SendPackets,
ServerTransportSystems::Flush,
TransportSystems::Flush,
)
.chain(),
)
.add_systems(
PreUpdate,
(poll, update_state, update_client_data)
.chain()
.in_set(ServerTransportSystems::Poll)
.run_if(resource_exists::<ServerMessages>),
)
.add_systems(
PostUpdate,
(flush, handle_disconnect_requests)
.in_set(ServerTransportSystems::Flush)
.run_if(resource_exists::<ServerMessages>),
)
.add_observer(on_connected);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, SystemSet)]
pub enum ServerTransportSystems {
Poll,
Flush,
}
#[derive(Debug, Clone, Copy, Default, Component, Reflect)]
#[reflect(Component)]
pub struct AeronetRepliconServer;
type OpenedServer = (
With<ServerEndpoint>,
With<Server>,
With<AeronetRepliconServer>,
);
fn update_state(
server_state: Res<State<ServerState>>,
mut next_server_state: ResMut<NextState<ServerState>>,
open_servers: Query<(), OpenedServer>,
) {
let running = open_servers.iter().next().is_some();
let next_state = if running {
ServerState::Running
} else {
ServerState::Stopped
};
if *server_state.get() != next_state {
next_server_state.set(next_state);
}
}
fn on_connected(
trigger: On<Add, Session>,
sessions: Query<&Session>,
child_of: Query<&ChildOf>,
open_servers: Query<(), OpenedServer>,
channels: Res<RepliconChannels>,
mut commands: Commands,
) {
let client = trigger.event_target();
let session = sessions
.get(client)
.expect("we are adding this component to this entity");
let Ok(&ChildOf(server)) = child_of.get(client) else {
return;
};
if open_servers.get(server).is_err() {
return;
}
let rx_lanes = channels
.client_channels()
.iter()
.map(|channel| convert::to_lane_kind(*channel));
let tx_lanes = channels
.server_channels()
.iter()
.map(|channel| convert::to_lane_kind(*channel));
let transport = match Transport::new(session, rx_lanes, tx_lanes, Instant::now()) {
Ok(transport) => transport,
Err(err) => {
warn!("Failed to create transport for {client} connecting to {server}: {err:?}");
return;
}
};
log::info!("insert client");
commands.entity(client).insert((
ConnectedClient {
max_size: session.mtu(),
},
transport,
));
}
fn poll(
mut server_msgs: ResMut<ServerMessages>,
mut clients: Query<(Entity, &mut Transport, &ChildOf)>,
open_servers: Query<(), OpenedServer>,
) {
for (client, mut transport, &ChildOf(server)) in &mut clients {
if open_servers.get(server).is_err() {
continue;
}
let mut msgs_recv = Saturating(0usize);
let mut bytes_recv = Saturating(0usize);
for msg in transport.recv.msgs.drain() {
msgs_recv += 1;
bytes_recv += msg.payload.len();
let channel_id = convert::to_channel_id(msg.lane);
server_msgs.insert_received(client, channel_id, msg.payload);
}
if msgs_recv.0 > 0 || bytes_recv.0 > 0 {
trace!(
"Server {server} received {msgs_recv} messages ({bytes_recv} bytes) from client \
{client}"
);
}
for _ in transport.recv.acks.drain() {
}
}
}
fn update_client_data(
mut commands: Commands,
mut clients: Query<(
Entity,
&Session,
&SessionStats,
&ConnectedClient,
&mut ConnectedClientStats,
)>,
sampling: Res<SessionStatsSampling>,
) {
for (client, session, session_stats, connected_client, mut client_stats) in &mut clients {
let stats = session_stats.last().copied().unwrap_or_default();
let mtu = session.mtu();
if connected_client.max_size != mtu {
commands
.entity(client)
.insert(ConnectedClient { max_size: mtu });
}
client_stats.rtt = stats.msg_rtt.as_secs_f64();
client_stats.packet_loss = stats.loss;
#[expect(clippy::cast_precision_loss, reason = "precision loss is acceptable")]
{
client_stats.received_bps = stats.packets_delta.bytes_recv.0 as f64 * sampling.rate();
client_stats.sent_bps = stats.packets_delta.bytes_sent.0 as f64 * sampling.rate();
}
}
}
fn flush(mut server_msgs: ResMut<ServerMessages>, mut clients: Query<&mut Transport>) {
let now = Instant::now();
for (client, channel_id, msg) in server_msgs.drain_sent() {
let Ok(mut transport) = clients.get_mut(client) else {
warn!("Sending to non-existent client {client}");
continue;
};
let Some(lane_index) = convert::to_lane_index(channel_id) else {
warn!("Channel {channel_id} is too large to convert to a lane index");
continue;
};
trace!("Sent 1 message ({} bytes) to client {client}", msg.len());
_ = transport.send.push(lane_index, msg, now);
}
}
fn handle_disconnect_requests(
mut commands: Commands,
mut requests: MessageReader<DisconnectRequest>,
clients: Query<(), (With<Transport>, With<ChildOf>)>,
) {
for request in requests.read() {
let client = request.client;
if let Err(err) = clients.get(client) {
warn!("Requested to disconnect {client} which does not match query: {err:?}");
continue;
}
commands.entity(client).despawn();
}
}