use crate::{Identity, Tick};
use bevy::{
prelude::{Deref, DerefMut, Resource},
utils::HashMap,
};
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum SendRule {
All,
Except(u32),
Only(u32),
}
impl SendRule {
pub fn includes(&self, ident: Identity) -> bool {
match self {
Self::All => true,
Self::Except(client_id) => ident != Identity::Client(*client_id),
Self::Only(client_id) => ident == Identity::Client(*client_id),
}
}
}
const THRESHOLD: usize = 1100;
const CAP: usize = 1198;
const PREALLOC: usize = 1500;
#[derive(Clone, Copy, Hash, PartialEq, Eq)]
pub struct BufferKey {
pub recipient: Identity,
pub channel: u8,
}
#[derive(Resource, Deref, DerefMut)]
pub struct WriteBuffer(Vec<u8>);
impl Default for WriteBuffer {
fn default() -> Self {
Self(Vec::with_capacity(4096))
}
}
impl std::io::Write for WriteBuffer {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.write(buf)
}
fn flush(&mut self) -> std::io::Result<()> {
self.0.flush()
}
}
#[derive(Resource, Default)]
pub struct Buffers {
current: HashMap<BufferKey, Vec<u8>>,
filled: Vec<(BufferKey, Vec<u8>)>,
taken_cache: Option<Vec<TakenBuffer>>,
}
impl Buffers {
pub fn remove(&mut self, ident: Identity) {
self.current.retain(|key, _| key.recipient != ident);
}
pub fn take(
&'_ mut self,
tick: Tick,
channel: u8,
this_run: bevy::ecs::component::Tick,
targets: impl Iterator<Item = (impl Into<Identity>, impl Into<RecipientData>)>,
) -> TakenBuffers<'_> {
let mut taken = self.taken_cache.take().unwrap_or_default();
for (recipient, info) in targets {
let recipient = recipient.into();
let buffer = self
.current
.remove(&BufferKey { recipient, channel })
.unwrap_or(Vec::with_capacity(PREALLOC));
taken.push(TakenBuffer {
recipient,
info: info.into(),
buffer,
last_fragment: 0,
});
}
TakenBuffers {
this_run,
tick: tick.to_le_bytes(),
channel,
buffers: self,
taken,
overhead: 0,
}
}
pub fn drain(&mut self, tick: Tick) -> impl Iterator<Item = (BufferKey, Vec<u8>)> + '_ {
let tick = tick.to_le_bytes();
self.current
.iter_mut()
.filter_map(move |(key, buf)| {
if buf.is_empty() {
None
} else {
let mut packet = Vec::with_capacity(buf.len() + 4);
packet.extend(tick);
packet.append(buf);
Some((*key, packet))
}
})
.chain(self.filled.drain(..))
}
}
#[derive(Default)]
pub struct RecipientData {
pub last_ack: Option<bevy::ecs::component::Tick>,
}
pub struct WriteFilters {
pub rule: SendRule,
pub changed: bevy::ecs::component::Tick,
}
pub struct TakenBuffer {
recipient: Identity,
info: RecipientData,
buffer: Vec<u8>,
last_fragment: usize,
}
pub struct TakenBuffers<'a> {
this_run: bevy::ecs::component::Tick,
tick: [u8; 4],
channel: u8,
buffers: &'a mut Buffers,
taken: Vec<TakenBuffer>,
overhead: u8,
}
impl<'a> Drop for TakenBuffers<'a> {
fn drop(&mut self) {
let mut taken = std::mem::take(&mut self.taken);
for taken in taken.drain(..) {
self.buffers.current.insert(
BufferKey {
recipient: taken.recipient,
channel: self.channel,
},
taken.buffer,
);
}
self.buffers.taken_cache = Some(taken);
}
}
impl<'a> TakenBuffers<'a> {
pub fn overhead(&mut self, overhead: u8) {
self.overhead += overhead;
}
pub fn send(&mut self, rule: SendRule, buf: &mut WriteBuffer) {
self.send_filtered(
WriteFilters {
rule,
changed: self.this_run,
},
buf,
);
}
pub fn send_filtered(&mut self, filter: WriteFilters, buf: &mut WriteBuffer) {
for taken in &mut self.taken {
if filter.rule.includes(taken.recipient)
&& (taken.info.last_ack.is_none()
|| filter
.changed
.is_newer_than(taken.info.last_ack.unwrap(), self.this_run))
{
taken.buffer.extend(buf.iter());
}
}
buf.clear();
}
pub fn fragment(&mut self) {
for taken in &mut self.taken {
let mut len = taken.buffer.len();
if len <= self.overhead as usize {
taken.buffer.drain(taken.last_fragment..);
continue;
}
if len < THRESHOLD {
taken.last_fragment = len;
continue;
}
while len > THRESHOLD {
let end = if len < CAP || taken.last_fragment == 0 {
taken.buffer.len()
} else {
taken.last_fragment
};
let mut packet = Vec::with_capacity(end + 4);
packet.extend(self.tick);
packet.extend(taken.buffer.drain(..end));
self.buffers.filled.push((
BufferKey {
recipient: taken.recipient,
channel: self.channel,
},
packet,
));
len = taken.buffer.len();
taken.last_fragment = 0;
}
}
self.overhead = 0;
}
}