use bytes::Bytes;
use dashmap::DashMap;
use std::collections::VecDeque;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
struct BufferedMessage {
stream_id: Option<String>,
serial: u64,
message_bytes: Bytes,
timestamp: Instant,
}
struct ChannelBufferState {
messages: VecDeque<BufferedMessage>,
current_stream_id: Option<String>,
}
struct ChannelBuffer {
state: Mutex<ChannelBufferState>,
next_serial: AtomicU64,
}
pub enum ReplayLookup {
Recovered(Vec<Bytes>),
Expired,
StreamReset { current_stream_id: Option<String> },
}
pub struct ReplayBuffer {
buffers: DashMap<String, ChannelBuffer>,
max_buffer_size: usize,
buffer_ttl: Duration,
}
impl ReplayBuffer {
pub fn new(max_buffer_size: usize, buffer_ttl: Duration) -> Self {
Self {
buffers: DashMap::new(),
max_buffer_size,
buffer_ttl,
}
}
fn buffer_key(app_id: &str, channel: &str) -> String {
format!("{}\0{}", app_id, channel)
}
fn prune_expired_locked(
messages: &mut VecDeque<BufferedMessage>,
buffer_ttl: Duration,
now: Instant,
) {
while let Some(front) = messages.front() {
if now.duration_since(front.timestamp) >= buffer_ttl {
messages.pop_front();
} else {
break;
}
}
}
pub fn next_serial(&self, app_id: &str, channel: &str) -> u64 {
let key = Self::buffer_key(app_id, channel);
let entry = self.buffers.entry(key).or_insert_with(|| ChannelBuffer {
state: Mutex::new(ChannelBufferState {
messages: VecDeque::with_capacity(self.max_buffer_size),
current_stream_id: None,
}),
next_serial: AtomicU64::new(1),
});
entry.next_serial.fetch_add(1, Ordering::Relaxed)
}
pub fn store(
&self,
app_id: &str,
channel: &str,
stream_id: Option<&str>,
serial: u64,
message_bytes: Bytes,
) {
let key = Self::buffer_key(app_id, channel);
let entry = self.buffers.entry(key).or_insert_with(|| ChannelBuffer {
state: Mutex::new(ChannelBufferState {
messages: VecDeque::with_capacity(self.max_buffer_size),
current_stream_id: stream_id.map(ToString::to_string),
}),
next_serial: AtomicU64::new(serial + 1),
});
let mut state = entry.state.lock().unwrap();
state.current_stream_id = stream_id.map(ToString::to_string);
while state.messages.len() >= self.max_buffer_size {
state.messages.pop_front();
}
state.messages.push_back(BufferedMessage {
stream_id: stream_id.map(ToString::to_string),
serial,
message_bytes,
timestamp: Instant::now(),
});
}
pub fn get_messages_after(
&self,
app_id: &str,
channel: &str,
last_serial: u64,
) -> Option<Vec<Bytes>> {
match self.get_messages_after_position(app_id, channel, None, last_serial) {
ReplayLookup::Recovered(messages) => Some(messages),
ReplayLookup::Expired | ReplayLookup::StreamReset { .. } => None,
}
}
pub fn get_messages_after_position(
&self,
app_id: &str,
channel: &str,
stream_id: Option<&str>,
last_serial: u64,
) -> ReplayLookup {
let key = Self::buffer_key(app_id, channel);
let Some(entry) = self.buffers.get(&key) else {
return ReplayLookup::Expired;
};
let now = Instant::now();
let mut state = entry.state.lock().unwrap();
if let Some(expected_stream_id) = stream_id
&& state.current_stream_id.as_deref() != Some(expected_stream_id)
{
return ReplayLookup::StreamReset {
current_stream_id: state.current_stream_id.clone(),
};
}
Self::prune_expired_locked(&mut state.messages, self.buffer_ttl, now);
if state.messages.is_empty() {
if stream_id.is_some() {
return ReplayLookup::Expired;
}
let next = entry.next_serial.load(Ordering::Relaxed);
return if last_serial >= next.saturating_sub(1) {
ReplayLookup::Recovered(Vec::new())
} else {
ReplayLookup::Expired
};
}
let newest_serial = state.messages.back().map(|m| m.serial).unwrap_or(0);
if last_serial >= newest_serial {
return ReplayLookup::Recovered(Vec::new());
}
let oldest_serial = state.messages.front().map(|m| m.serial).unwrap_or(0);
if last_serial > 0 && last_serial < oldest_serial.saturating_sub(1) {
return ReplayLookup::Expired;
}
let contiguous = state.messages.make_contiguous();
let start_idx = contiguous.partition_point(|message| message.serial <= last_serial);
let mut result = Vec::with_capacity(contiguous.len().saturating_sub(start_idx));
for message in &contiguous[start_idx..] {
let _ = &message.stream_id;
result.push(message.message_bytes.clone());
}
ReplayLookup::Recovered(result)
}
pub fn evict_expired(&self) {
let now = Instant::now();
let mut empty_keys = Vec::new();
for entry in self.buffers.iter() {
let mut state = entry.value().state.lock().unwrap();
Self::prune_expired_locked(&mut state.messages, self.buffer_ttl, now);
if state.messages.is_empty() {
empty_keys.push(entry.key().clone());
}
}
for key in empty_keys {
self.buffers
.remove_if(&key, |_, v| v.state.lock().unwrap().messages.is_empty());
}
}
}