use crate::{Result, resp::RespResponse};
use crossbeam_queue::SegQueue;
use futures_util::{Stream, task::AtomicWaker};
use std::{
pin::Pin,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll},
};
type Item = Result<RespResponse>;
struct Shared {
queue: SegQueue<(Item, usize)>,
bytes: AtomicUsize,
waker: AtomicWaker,
dropped: AtomicUsize,
senders: AtomicUsize,
receiver_alive: AtomicUsize,
max_bytes: usize,
}
#[derive(Debug)]
pub(crate) struct SendError(Item);
impl SendError {
pub(crate) fn into_inner(self) -> Item {
self.0
}
}
impl std::fmt::Display for SendError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("the subscriber is gone")
}
}
pub(crate) struct BoundedSender {
shared: Arc<Shared>,
}
pub(crate) struct BoundedReceiver {
shared: Arc<Shared>,
}
pub(crate) fn bounded_channel(max_bytes: usize) -> (BoundedSender, BoundedReceiver) {
let shared = Arc::new(Shared {
queue: SegQueue::new(),
bytes: AtomicUsize::new(0),
waker: AtomicWaker::new(),
dropped: AtomicUsize::new(0),
senders: AtomicUsize::new(1),
receiver_alive: AtomicUsize::new(1),
max_bytes,
});
(
BoundedSender {
shared: Arc::clone(&shared),
},
BoundedReceiver { shared },
)
}
impl BoundedSender {
pub(crate) fn send(&self, item: Item) -> std::result::Result<(), SendError> {
if self.shared.receiver_alive.load(Ordering::Acquire) == 0 {
return Err(SendError(item));
}
let cost = item
.as_ref()
.map(|response| response.retained_bytes())
.unwrap_or(0);
self.shared.queue.push((item, cost));
self.shared.bytes.fetch_add(cost, Ordering::AcqRel);
if self.shared.max_bytes != 0 {
let mut evicted = 0usize;
#[expect(
clippy::arithmetic_side_effects,
reason = "one increment per queue entry actually popped, so the \
count is bounded by the queue length."
)]
while self.shared.bytes.load(Ordering::Acquire) > self.shared.max_bytes
&& self.shared.queue.len() > 1
{
let Some((_, cost)) = self.shared.queue.pop() else {
break;
};
self.shared.bytes.fetch_sub(cost, Ordering::AcqRel);
evicted += 1;
}
if evicted > 0 {
self.shared.dropped.fetch_add(evicted, Ordering::Relaxed);
}
}
self.shared.waker.wake();
Ok(())
}
}
impl Clone for BoundedSender {
fn clone(&self) -> Self {
self.shared.senders.fetch_add(1, Ordering::Relaxed);
Self {
shared: Arc::clone(&self.shared),
}
}
}
impl Drop for BoundedSender {
fn drop(&mut self) {
if self.shared.senders.fetch_sub(1, Ordering::AcqRel) == 1 {
self.shared.waker.wake();
}
}
}
impl std::fmt::Debug for BoundedSender {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("BoundedSender")
}
}
impl BoundedReceiver {
pub(crate) fn dropped_messages(&self) -> usize {
self.shared.dropped.load(Ordering::Relaxed)
}
}
impl Stream for BoundedReceiver {
type Item = Item;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let shared = &self.shared;
shared.waker.register(cx.waker());
if let Some((item, cost)) = shared.queue.pop() {
shared.bytes.fetch_sub(cost, Ordering::AcqRel);
return Poll::Ready(Some(item));
}
if shared.senders.load(Ordering::Acquire) == 0 {
return Poll::Ready(None);
}
Poll::Pending
}
}
impl Drop for BoundedReceiver {
fn drop(&mut self) {
self.shared.receiver_alive.store(0, Ordering::Release);
while self.shared.queue.pop().is_some() {}
self.shared.bytes.store(0, Ordering::Release);
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::panic,
clippy::arithmetic_side_effects,
reason = "test code: a panic is how a test reports failure"
)]
mod tests {
use super::*;
use crate::resp::{RespFrameParser, RespResponse, RespTapeMut};
use bytes::Bytes;
use futures_util::{FutureExt, StreamExt};
fn response(id: usize, payload: usize) -> Item {
let body = format!("+{id:0>width$}\r\n", width = payload - 3);
let bytes = Bytes::from(body);
let mut tape = RespTapeMut::default();
let mut parser = RespFrameParser::new(&bytes, &mut tape);
let (frame, _) = parser.parse().unwrap();
Ok(RespResponse::new(bytes.into(), frame))
}
fn id_of(item: &Item) -> usize {
let text: String = item.as_ref().unwrap().to().unwrap();
text.trim_start_matches('0').parse().unwrap_or(0)
}
#[test]
fn the_oldest_messages_are_the_ones_dropped() {
const PAYLOAD: usize = 100;
let (sender, mut receiver) = bounded_channel(3 * PAYLOAD);
for id in 0..10 {
sender.send(response(id, PAYLOAD)).unwrap();
}
let mut received = Vec::new();
while let Some(Some(item)) = receiver.next().now_or_never() {
received.push(id_of(&item));
}
assert_eq!(
vec![7, 8, 9],
received,
"the newest messages must be the ones kept"
);
assert_eq!(
7,
receiver.dropped_messages(),
"every evicted message must be counted"
);
}
#[test]
fn a_consumer_that_keeps_up_loses_nothing() {
const PAYLOAD: usize = 100;
let (sender, mut receiver) = bounded_channel(3 * PAYLOAD);
for id in 0..10 {
sender.send(response(id, PAYLOAD)).unwrap();
let item = receiver.next().now_or_never().unwrap().unwrap();
assert_eq!(id, id_of(&item));
}
assert_eq!(0, receiver.dropped_messages());
}
#[test]
fn a_zero_budget_keeps_everything() {
const PAYLOAD: usize = 100;
let (sender, mut receiver) = bounded_channel(0);
for id in 0..50 {
sender.send(response(id, PAYLOAD)).unwrap();
}
let mut count = 0;
while receiver.next().now_or_never().flatten().is_some() {
count += 1;
}
assert_eq!(50, count);
assert_eq!(0, receiver.dropped_messages());
}
#[test]
fn the_stream_ends_when_the_last_sender_goes() {
let (sender, mut receiver) = bounded_channel(0);
sender.send(response(1, 100)).unwrap();
drop(sender);
assert!(receiver.next().now_or_never().flatten().is_some());
assert!(
receiver.next().now_or_never().unwrap().is_none(),
"the stream must end, not stay pending"
);
}
#[test]
fn sending_to_a_departed_subscriber_returns_the_message() {
let (sender, receiver) = bounded_channel(0);
drop(receiver);
let error = sender.send(response(1, 100)).unwrap_err();
assert_eq!(1, id_of(&error.into_inner()));
}
}