use std::{
ops::{Deref, DerefMut},
sync::{Arc, Weak},
};
use ruma::{OwnedEventId, OwnedRoomId};
use tokio::sync::{broadcast::Receiver, mpsc};
use tracing::{trace, warn};
#[derive(Default)]
pub struct SubscribersHandle(Arc<()>);
impl SubscribersHandle {
pub fn count(&self) -> usize {
Arc::weak_count(&self.0)
}
pub fn new_subscriber_handle(&self) -> SubscriberHandle {
SubscriberHandle(Arc::downgrade(&self.0))
}
}
pub struct SubscriberHandle(Weak<()>);
impl SubscriberHandle {
pub fn count(&self) -> usize {
Weak::weak_count(&self.0)
}
}
#[allow(missing_debug_implementations)]
pub struct Subscriber<T> {
subscriber_receiver: Receiver<T>,
auto_shrink_message: Option<AutoShrinkMessage>,
auto_shrink_sender: mpsc::Sender<AutoShrinkMessage>,
subscriber_handle: Option<SubscriberHandle>,
}
impl<T> Subscriber<T> {
pub(super) fn new(
subscriber_receiver: Receiver<T>,
auto_shrink_message: AutoShrinkMessage,
auto_shrink_sender: mpsc::Sender<AutoShrinkMessage>,
subscribers_handle: &SubscribersHandle,
) -> Self {
Self {
subscriber_receiver,
auto_shrink_message: Some(auto_shrink_message),
auto_shrink_sender,
subscriber_handle: Some(subscribers_handle.new_subscriber_handle()),
}
}
}
impl<T> Drop for Subscriber<T> {
fn drop(&mut self) {
let number_of_subscribers = self
.subscriber_handle
.take()
.expect("Unreachable: `subscriber_handle` must be `Some`")
.count();
trace!("dropping a room event cache subscriber; count: {number_of_subscribers}");
if number_of_subscribers == 1 {
let mut message = self
.auto_shrink_message
.take()
.expect("Unreachable: `auto_shrink_message` must be `Some`");
let mut num_attempts = 0;
while let Err(err) = self.auto_shrink_sender.try_send(message) {
num_attempts += 1;
if num_attempts > 1024 {
warn!(
"couldn't send notification to the auto-shrink channel \
after 1024 attempts; giving up"
);
return;
}
match err {
mpsc::error::TrySendError::Full(stolen_message) => {
message = stolen_message;
}
mpsc::error::TrySendError::Closed(_) => return,
}
}
trace!("sent notification to the parent channel that we were the last subscriber");
}
}
}
impl<T> Deref for Subscriber<T> {
type Target = Receiver<T>;
fn deref(&self) -> &Self::Target {
&self.subscriber_receiver
}
}
impl<T> DerefMut for Subscriber<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.subscriber_receiver
}
}
#[derive(Debug)]
pub enum AutoShrinkMessage {
Room {
room_id: OwnedRoomId,
},
Thread {
room_id: OwnedRoomId,
thread_id: OwnedEventId,
},
}
#[cfg(test)]
mod tests {
use assert_matches::assert_matches;
use ruma::owned_room_id;
use tokio::sync::{broadcast, mpsc};
use super::{AutoShrinkMessage, Subscriber, SubscribersHandle};
#[test]
fn test_subscribers_handle() {
let subscribers_handle = SubscribersHandle::default();
assert_eq!(subscribers_handle.count(), 0);
let handle0 = subscribers_handle.new_subscriber_handle();
assert_eq!(subscribers_handle.count(), 1);
assert_eq!(handle0.count(), 1);
let handle1 = subscribers_handle.new_subscriber_handle();
assert_eq!(subscribers_handle.count(), 2);
assert_eq!(handle0.count(), 2);
assert_eq!(handle1.count(), 2);
drop(handle0);
assert_eq!(subscribers_handle.count(), 1);
assert_eq!(handle1.count(), 1);
drop(handle1);
assert_eq!(subscribers_handle.count(), 0);
let handle2 = subscribers_handle.new_subscriber_handle();
drop(subscribers_handle);
assert_eq!(handle2.count(), 0);
}
#[test]
fn test_subscriber_t_derefs_to_t() {
let (auto_shrink_sender, _auto_shrink_receiver) = mpsc::channel(1);
let (subscriber_sender, subscriber_receiver) = broadcast::channel(1);
let subscribers_handle = SubscribersHandle::default();
let mut subscriber = Subscriber::new(
subscriber_receiver,
AutoShrinkMessage::Room { room_id: owned_room_id!("!r0") },
auto_shrink_sender,
&subscribers_handle,
);
subscriber_sender.send('a').unwrap();
assert_eq!(subscriber.try_recv().unwrap(), 'a');
assert!(subscriber.is_empty());
}
#[test]
fn test_subscriber_send_auto_shrink_message_on_last_drop() {
let (auto_shrink_sender, mut auto_shrink_receiver) = mpsc::channel(1);
let (_subscriber_sender, subscriber_receiver) = broadcast::channel::<()>(1);
let subscribers_handle = SubscribersHandle::default();
let room_id = owned_room_id!("!r0");
let auto_shrink_message = AutoShrinkMessage::Room { room_id: room_id.clone() };
let subscriber0 = Subscriber::new(
subscriber_receiver.resubscribe(),
AutoShrinkMessage::Room { room_id: room_id.clone() },
auto_shrink_sender.clone(),
&subscribers_handle,
);
let subscriber1 = Subscriber::new(
subscriber_receiver,
auto_shrink_message,
auto_shrink_sender,
&subscribers_handle,
);
drop(subscriber0);
assert!(auto_shrink_receiver.is_empty());
drop(subscriber1);
assert_matches!(
auto_shrink_receiver.try_recv().unwrap(),
AutoShrinkMessage::Room { room_id: expected_room_id } => {
assert_eq!(expected_room_id, room_id);
}
);
assert!(auto_shrink_receiver.is_empty());
}
#[test]
fn test_subscriber_send_auto_shrink_message_with_full_channel() {
let (auto_shrink_sender, mut auto_shrink_receiver) = mpsc::channel(1);
let (_subscriber_sender, subscriber_receiver) = broadcast::channel::<()>(1);
let subscribers_handle = SubscribersHandle::default();
let noisy_room_id = owned_room_id!("!r1");
let auto_shrink_noisy_message = AutoShrinkMessage::Room { room_id: noisy_room_id.clone() };
let room_id = owned_room_id!("!r0");
let auto_shrink_message = AutoShrinkMessage::Room { room_id };
auto_shrink_sender.try_send(auto_shrink_noisy_message).unwrap();
let subscriber = Subscriber::new(
subscriber_receiver,
auto_shrink_message,
auto_shrink_sender,
&subscribers_handle,
);
drop(subscriber);
assert_matches!(
auto_shrink_receiver.try_recv().unwrap(),
AutoShrinkMessage::Room { room_id: expected_room_id } => {
assert_eq!(expected_room_id, noisy_room_id);
}
);
assert!(auto_shrink_receiver.is_empty());
}
#[test]
fn test_subscriber_send_auto_shrink_message_with_closed_channel() {
let (auto_shrink_sender, auto_shrink_receiver) = mpsc::channel(1);
let (_subscriber_sender, subscriber_receiver) = broadcast::channel::<()>(1);
let subscribers_handle = SubscribersHandle::default();
let room_id = owned_room_id!("!r0");
let auto_shrink_message = AutoShrinkMessage::Room { room_id };
let subscriber = Subscriber::new(
subscriber_receiver,
auto_shrink_message,
auto_shrink_sender,
&subscribers_handle,
);
drop(auto_shrink_receiver);
drop(subscriber);
}
}