use anyhow::Result;
use serde::{Deserialize, Serialize};
use std::{collections::HashSet, fmt, ops::Deref, time::Instant};
use tokio::{io::AsyncWriteExt as _, net::tcp::OwnedWriteHalf};
use tracing::{debug, info};
use uuid::Uuid;
const DEFAULT_DECAY: u64 = 6;
#[derive(Debug, Default)]
pub struct MessageBuffer {
buffer: Vec<DecayWrapper<ProxiedMessage>>,
next_msg_id: u16,
}
impl MessageBuffer {
pub fn new() -> Self {
MessageBuffer {
buffer: Vec::new(),
next_msg_id: 0,
}
}
pub fn push(&mut self, msg: ProxiedMessage) {
self.buffer.push(DecayWrapper::new(msg));
}
pub fn retain<F>(&mut self, f: F)
where
F: FnMut(&DecayWrapper<ProxiedMessage>) -> bool,
{
self.buffer.retain(f);
}
#[allow(dead_code)]
pub fn len(&self) -> usize {
self.buffer.len()
}
pub fn is_empty(&self) -> bool {
self.buffer.is_empty()
}
pub fn iter(&self) -> std::slice::Iter<'_, DecayWrapper<ProxiedMessage>> {
self.buffer.iter()
}
pub async fn tick(&mut self, write: &mut OwnedWriteHalf) -> Result<bool> {
if self.is_empty() {
return Ok(false);
}
debug!("Messages in buffer:");
for msg in self.iter() {
debug!("{}", msg.inner());
}
let mut send_buffer = self
.iter()
.filter(|msg| msg.decayed() || msg.message_id() <= self.next_msg_id)
.map(|msg| msg.inner())
.collect::<Vec<&ProxiedMessage>>();
send_buffer.sort_by(|a, b| a.message_id.cmp(&b.message_id()));
if send_buffer.is_empty() {
debug!("send buf is empty");
return Ok(false);
}
let mut sent_messages = HashSet::new();
for msg in send_buffer {
match &msg.message() {
Payload::Data(data) => {
write.write_all(data).await?;
info!("Wrote message {} to stream", msg.message_id())
}
Payload::Close => {
return Ok(true);
}
}
sent_messages.insert(msg.message_id());
}
self.next_msg_id = sent_messages
.iter()
.max()
.expect("This is safe since we know we've set something")
+ 1;
self.retain(|msg| !sent_messages.contains(&msg.inner().message_id()));
info!("next_msg_id is: {}", self.next_msg_id.clone());
Ok(false)
}
}
#[derive(Debug)]
pub struct DecayWrapper<T> {
value: T,
start: Instant,
decay: u64,
}
impl<T> DecayWrapper<T> {
pub fn decayed(&self) -> bool {
debug!("Decayed: {:?}", self.start.elapsed().as_secs() > self.decay);
self.start.elapsed().as_secs() > self.decay
}
pub fn new(value: T) -> Self {
DecayWrapper {
value,
start: Instant::now(),
decay: DEFAULT_DECAY,
}
}
#[allow(dead_code)]
pub fn into_inner(self) -> T {
self.value
}
pub fn inner(&self) -> &T {
&self.value
}
}
impl<T> Deref for DecayWrapper<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.value
}
}
#[derive(Serialize, Deserialize, Clone, Debug)]
pub struct ProxiedMessage {
pub message: Payload,
pub session_id: Uuid,
pub message_id: u16,
}
impl ProxiedMessage {
pub fn new(message: Payload, session_id: Uuid, message_id: u16) -> Self {
ProxiedMessage {
message,
session_id,
message_id,
}
}
pub fn message(&self) -> &Payload {
&self.message
}
pub fn session_id(&self) -> Uuid {
self.session_id
}
pub fn message_id(&self) -> u16 {
self.message_id
}
}
#[derive(Serialize, Deserialize, Clone, Debug, PartialEq)]
pub enum Payload {
Data(Vec<u8>),
Close,
}
impl fmt::Display for ProxiedMessage {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let message = match self.message() {
Payload::Data(ref data) => format!("Data({})", data.len()),
Payload::Close => "Close".to_string(),
};
write!(
f,
"ProxiedMessage {{ message: {}, session_id: {}, message_id: {} }}",
message,
self.session_id(),
self.message_id()
)
}
}