use bevy::prelude::Res;
use std::fmt::Debug;
use std::io;
use std::io::Write;
use std::marker::PhantomData;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use crossbeam_channel::{Receiver, Sender};
use crate::packet_length_serializer::{PacketLengthDeserializationError, PacketLengthSerializer};
use crate::protocol::NetworkStream;
use crate::serializer::Serializer;
pub struct EcsConnection<SendingPacket>
where
SendingPacket: Send + Sync + Debug + 'static,
{
pub(crate) id: ConnectionId,
pub(crate) packet_tx: Sender<SendingPacket>,
pub(crate) local_addr: SocketAddr,
pub(crate) peer_addr: SocketAddr,
}
impl<SendingPacket> Clone for EcsConnection<SendingPacket>
where
SendingPacket: Send + Sync + Debug + 'static,
{
fn clone(&self) -> Self {
EcsConnection {
id: self.id,
packet_tx: self.packet_tx.clone(),
local_addr: self.local_addr,
peer_addr: self.peer_addr,
}
}
}
impl<SendingPacket> EcsConnection<SendingPacket>
where
SendingPacket: Send + Sync + Debug + 'static,
{
pub fn id(&self) -> ConnectionId {
self.id
}
pub fn peer_addr(&self) -> SocketAddr {
self.peer_addr
}
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub fn send(&self, packet: SendingPacket) {
self.packet_tx.send(packet).unwrap();
}
}
pub(crate) struct RawConnection<ReceivingPacket, SendingPacket, NS, S, LS>
where
ReceivingPacket: Send + Sync + Debug + 'static,
SendingPacket: Send + Sync + Debug + 'static,
NS: NetworkStream,
S: Serializer<ReceivingPacket, SendingPacket>,
LS: PacketLengthSerializer,
{
pub stream: NS,
pub serializer: Arc<S>,
pub packet_length_serializer: Arc<LS>,
pub packets_rx: Receiver<SendingPacket>,
pub id: ConnectionId,
pub _receive_packet: PhantomData<ReceivingPacket>,
pub _send_packet: PhantomData<SendingPacket>,
}
#[derive(
Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Debug, Hash, bevy::ecs::component::Component,
)]
pub struct ConnectionId(usize);
impl ConnectionId {
pub fn next() -> ConnectionId {
static CONNECTION_ID: AtomicUsize = AtomicUsize::new(0);
ConnectionId(CONNECTION_ID.fetch_add(1, Ordering::Relaxed))
}
}
pub(crate) static MAX_PACKET_SIZE: AtomicUsize = AtomicUsize::new(usize::MAX);
#[derive(Copy, Clone)]
pub struct MaxPacketSize(pub usize);
impl<ReceivingPacket, SendingPacket, NS, S, LS>
RawConnection<ReceivingPacket, SendingPacket, NS, S, LS>
where
ReceivingPacket: Send + Sync + Debug + 'static,
SendingPacket: Send + Sync + Debug + 'static,
NS: NetworkStream,
S: Serializer<ReceivingPacket, SendingPacket>,
LS: PacketLengthSerializer,
{
pub fn new(
stream: NS,
serializer: S,
packet_length_serializer: LS,
packets_rx: Receiver<SendingPacket>,
) -> Self {
stream.set_nonblocking();
Self {
stream,
serializer: Arc::new(serializer),
packet_length_serializer: Arc::new(packet_length_serializer),
packets_rx,
id: ConnectionId::next(),
_receive_packet: PhantomData,
_send_packet: PhantomData,
}
}
pub fn send(&mut self, packet: SendingPacket) -> io::Result<()> {
let buf = self.serialize_packet(packet)?;
self.stream.write_all(&buf)?;
Ok(())
}
pub fn serialize_packet(&self, packet: SendingPacket) -> io::Result<Vec<u8>> {
let serialized = self
.serializer
.serialize(packet)
.expect("Error serializing packet");
let mut buf = self
.packet_length_serializer
.serialize_packet_length(serialized.len())
.expect("Error serializing packet length");
buf.write_all(&serialized)?;
Ok(buf)
}
pub fn receive(
&mut self,
) -> Result<ReceivingPacket, ReceiveError<ReceivingPacket, SendingPacket, S, LS>> {
let mut size = 0;
let mut length = Err(PacketLengthDeserializationError::NeedMoreBytes(LS::SIZE));
while let Err(PacketLengthDeserializationError::NeedMoreBytes(amt)) = length {
size += amt;
let mut buf = vec![0; size];
self.stream
.try_peek_exact(&mut buf)
.map_err(ReceiveError::Io)?;
length = self
.packet_length_serializer
.deserialize_packet_length(&buf);
}
match length {
Ok(length) => {
if length > MAX_PACKET_SIZE.load(Ordering::Relaxed) {
Err(ReceiveError::PacketTooBig)
} else {
let mut buf = vec![0; length + size];
self.stream.read_exact(&mut buf).map_err(ReceiveError::Io)?;
Ok(self
.serializer
.deserialize(&buf[size..])
.map_err(ReceiveError::Deserialization)?)
}
}
Err(PacketLengthDeserializationError::Err(err)) => {
Err(ReceiveError::LengthDeserialization(err))
}
Err(PacketLengthDeserializationError::NeedMoreBytes(_)) => unreachable!(),
}
}
pub fn id(&self) -> ConnectionId {
self.id
}
pub fn local_addr(&self) -> SocketAddr {
self.stream.local_addr()
}
pub fn peer_addr(&self) -> SocketAddr {
self.stream.peer_addr()
}
}
pub(crate) enum ReceiveError<ReceivingPacket, SendingPacket, S, LS>
where
ReceivingPacket: Send + Sync + Debug + 'static,
SendingPacket: Send + Sync + Debug + 'static,
S: Serializer<ReceivingPacket, SendingPacket>,
LS: PacketLengthSerializer,
{
Io(std::io::Error),
Deserialization(S::Error),
LengthDeserialization(LS::Error),
PacketTooBig,
}
pub(crate) fn max_packet_size_system(max_packet_size: Option<Res<MaxPacketSize>>) {
match max_packet_size {
Some(res) if res.is_changed() => {
MAX_PACKET_SIZE.store(res.0, Ordering::Relaxed);
}
_ => (),
}
}