use std::{
collections::VecDeque,
mem,
sync::{
Arc,
atomic::{
AtomicBool, AtomicU64,
Ordering::{Acquire, Relaxed, Release},
},
},
};
use parking_lot::Mutex;
use whasher::{GxPapayaMap, new_papaya_map};
use super::{
pattern_subscription_entry::{PatternSubscriberSet, PatternSubscriptionEntry},
subscriber::PubSubSink,
};
use crate::objects::sortedset::sorted_set_object::glob_match;
type ChannelSubscriptions = GxPapayaMap<Box<[u8]>, PatternSubscriberSet>;
type PatternSubscriptions = GxPapayaMap<Box<[u8]>, PatternSubscriptionEntry>;
type PendingEntry = (Box<[u8]>, Box<[u8]>);
pub struct SubscribeBroker {
subscriptions: ChannelSubscriptions,
pattern_subscriptions: PatternSubscriptions,
pending: Mutex<VecDeque<PendingEntry>>,
previous_address: AtomicU64,
page_size_bits: u32,
disposed: AtomicBool,
}
impl SubscribeBroker {
pub fn new(page_size_bytes: usize) -> Self {
Self {
subscriptions: new_papaya_map(),
pattern_subscriptions: new_papaya_map(),
pending: Mutex::new(VecDeque::new()),
previous_address: AtomicU64::new(0),
page_size_bits: page_size_bytes.max(2).ilog2(),
disposed: AtomicBool::new(false),
}
}
pub fn remove_subscription(&self, subscriber: u64) {
let subscriptions = self.subscriptions.pin();
for (_, set) in subscriptions.iter() {
set.pin().remove(&subscriber);
}
subscriptions.retain(|_, set| !set.is_empty());
let patterns = self.pattern_subscriptions.pin();
for (_, entry) in patterns.iter() {
entry.subscriptions.pin().remove(&subscriber);
}
patterns.retain(|_, entry| !entry.subscriptions.is_empty());
}
fn broadcast(&self, key: &[u8], value: &[u8]) -> usize {
let mut num_subscribers = 0;
let subscriptions = self.subscriptions.pin();
if let Some(sessions) = subscriptions.get(key) {
for (_, session) in sessions.pin().iter() {
session.publish(key, value);
num_subscribers += 1;
}
}
for (_, entry) in self.pattern_subscriptions.pin().iter() {
if glob_match(&entry.pattern, key) {
for (_, session) in entry.subscriptions.pin().iter() {
session.pattern_publish(&entry.pattern, key, value);
num_subscribers += 1;
}
}
}
num_subscribers
}
pub fn consume(&self, payload: &[u8], current_address: u64, next_address: u64) -> usize {
if self.disposed.load(Acquire) {
return 0;
}
let previous_address = self.previous_address.load(Relaxed);
if previous_address > 0 && current_address > previous_address {
let page_mask = (1usize << self.page_size_bits) as u64 - 1;
let payload_len = payload.len() as u64;
if (current_address & page_mask) != 0 || current_address >= previous_address + payload_len {
log::warn!("SubscribeBroker: Skipping from {previous_address} to {current_address}");
}
}
let Some((key, value)) = decode_payload(payload) else {
log::warn!("SubscribeBroker.Consume: malformed payload at {current_address}");
return 0;
};
let notified = self.broadcast(&key, &value);
self.previous_address.store(next_address, Relaxed);
notified
}
pub fn consume_pending(&self) -> usize {
if self.disposed.load(Acquire) {
return 0;
}
let pending = mem::take(&mut *self.pending.lock());
pending
.iter()
.fold(0, |n, (key, value)| n + self.broadcast(key, value))
}
pub fn subscribe(&self, channel: &[u8], subscriber: u64, sink: Arc<dyn PubSubSink>) -> bool {
if self.disposed.load(Acquire) {
return false;
}
let subscriptions = self.subscriptions.pin();
let sessions = subscriptions.get_or_insert_with(channel.into(), new_papaya_map);
sessions.pin().insert(subscriber, sink).is_none()
}
pub fn pattern_subscribe(
&self,
pattern: &[u8],
subscriber: u64,
sink: Arc<dyn PubSubSink>,
) -> bool {
if self.disposed.load(Acquire) {
return false;
}
let patterns = self.pattern_subscriptions.pin();
let entry = patterns.get_or_insert_with(pattern.into(), || {
PatternSubscriptionEntry::new(pattern.into())
});
entry.subscriptions.pin().insert(subscriber, sink).is_none()
}
pub fn unsubscribe(&self, channel: &[u8], subscriber: u64) -> bool {
let subscriptions = self.subscriptions.pin();
let removed = subscriptions
.get(channel)
.is_some_and(|sessions| sessions.pin().remove(&subscriber).is_some());
if removed {
subscriptions.retain(|_, set| !set.is_empty());
}
removed
}
pub fn pattern_unsubscribe(&self, pattern: &[u8], subscriber: u64) -> bool {
let patterns = self.pattern_subscriptions.pin();
let Some(entry) = patterns.get(pattern) else {
return false;
};
let removed = entry.subscriptions.pin().remove(&subscriber).is_some();
if removed && entry.subscriptions.is_empty() {
patterns.retain(|_, entry| !entry.subscriptions.is_empty());
}
removed
}
pub fn list_all_subscriptions(&self) -> Vec<Vec<u8>> {
self
.subscriptions
.pin()
.iter()
.filter(|(_, set)| !set.is_empty())
.map(|(channel, _)| channel.to_vec())
.collect()
}
pub fn list_all_pattern_subscriptions(&self) -> Vec<Vec<u8>> {
self
.pattern_subscriptions
.pin()
.iter()
.filter(|(_, entry)| !entry.subscriptions.is_empty())
.map(|(pattern, _)| pattern.to_vec())
.collect()
}
pub fn publish_now(&self, key: &[u8], value: &[u8]) -> usize {
if self.is_idle() {
return 0;
}
self.broadcast(key, value)
}
pub fn publish(&self, key: &[u8], value: &[u8]) {
if self.is_idle() {
return;
}
self.pending.lock().push_back((key.into(), value.into()));
}
pub fn get_channels(&self) -> Vec<Vec<u8>> {
self.list_all_subscriptions()
}
pub fn get_channels_matching(&self, pattern: &[u8]) -> Vec<Vec<u8>> {
self
.subscriptions
.pin()
.iter()
.filter(|(channel, set)| !set.is_empty() && glob_match(pattern, channel))
.map(|(channel, _)| channel.to_vec())
.collect()
}
pub fn num_pattern_subscriptions(&self) -> usize {
self
.pattern_subscriptions
.pin()
.iter()
.filter(|(_, entry)| !entry.subscriptions.is_empty())
.count()
}
pub fn num_subscriptions(&self, channel: &[u8]) -> usize {
self
.subscriptions
.pin()
.get(channel)
.map_or(0, PatternSubscriberSet::len)
}
pub fn dispose(&self) {
self.disposed.store(true, Release);
self.pending.lock().clear();
self.subscriptions.pin().clear();
self.pattern_subscriptions.pin().clear();
}
fn is_idle(&self) -> bool {
self.subscriptions.is_empty() && self.pattern_subscriptions.is_empty()
}
}
fn decode_payload(payload: &[u8]) -> Option<PendingEntry> {
let mut cursor = 0usize;
let read_len = |payload: &[u8], cursor: &mut usize| -> Option<usize> {
let head = payload.get(*cursor..cursor.checked_add(4)?)?;
*cursor += 4;
let len = i32::from_le_bytes(head.try_into().ok()?);
usize::try_from(len).ok()
};
let key_len = read_len(payload, &mut cursor)?;
let key = payload.get(cursor..cursor.checked_add(key_len)?)?;
cursor += key_len;
let value_len = read_len(payload, &mut cursor)?;
let value = payload.get(cursor..cursor.checked_add(value_len)?)?;
Some((key.into(), value.into()))
}
#[cfg(test)]
mod tests {
use super::{super::subscriber::PubSubMailbox, *};
struct Fixture {
broker: SubscribeBroker,
mailbox: Arc<PubSubMailbox>,
}
fn fixture(page_size: usize) -> Fixture {
let broker = SubscribeBroker::new(page_size);
let mailbox = Arc::new(PubSubMailbox::new(16));
Fixture { broker, mailbox }
}
#[test]
fn subscribe_unsubscribe_channel_lifecycle() {
let f = fixture(4096);
assert!(f.broker.subscribe(b"news", 1, f.mailbox.clone()));
assert!(!f.broker.subscribe(b"news", 1, f.mailbox.clone()));
assert_eq!(f.broker.num_subscriptions(b"news"), 1);
assert_eq!(f.broker.get_channels(), vec![b"news".to_vec()]);
assert!(f.broker.unsubscribe(b"news", 1));
assert!(!f.broker.unsubscribe(b"news", 1));
assert_eq!(f.broker.num_subscriptions(b"news"), 0);
assert!(f.broker.get_channels().is_empty());
}
#[test]
fn pattern_subscribe_and_broadcast_match() {
let f = fixture(4096);
assert!(f.broker.pattern_subscribe(b"news.*", 1, f.mailbox.clone()));
assert_eq!(f.broker.num_pattern_subscriptions(), 1);
assert_eq!(
f.broker.list_all_pattern_subscriptions(),
vec![b"news.*".to_vec()]
);
let notified = f.broker.publish_now(b"news.tech", b"hello");
assert_eq!(notified, 1);
let messages = f.mailbox.drain();
assert_eq!(messages.len(), 1);
assert_eq!(messages[0].channel.as_ref(), b"news.tech");
assert_eq!(f.broker.publish_now(b"other", b"x"), 0);
assert!(f.broker.pattern_unsubscribe(b"news.*", 1));
assert_eq!(f.broker.num_pattern_subscriptions(), 0);
}
#[test]
fn publish_now_reaches_channel_subscribers() {
let f = fixture(4096);
assert_eq!(f.broker.publish_now(b"ch", b"v"), 0);
f.broker.subscribe(b"ch", 7, f.mailbox.clone());
assert_eq!(f.broker.publish_now(b"ch", b"v"), 1);
let messages = f.mailbox.drain();
assert_eq!(messages[0].value.as_ref(), b"v");
}
#[test]
fn publish_enqueue_then_consume_pending_broadcasts() {
let f = fixture(4096);
f.broker.subscribe(b"ch", 1, f.mailbox.clone());
f.broker.publish(b"ch", b"queued");
assert!(f.mailbox.is_empty());
assert_eq!(f.broker.consume_pending(), 1);
assert_eq!(f.mailbox.len(), 1);
assert_eq!(f.broker.consume_pending(), 0);
}
#[test]
fn consume_decodes_length_prefixed_payload() {
let f = fixture(4096);
f.broker.subscribe(b"k", 1, f.mailbox.clone());
let mut payload = Vec::new();
payload.extend_from_slice(&1i32.to_le_bytes());
payload.push(b'k');
payload.extend_from_slice(&2i32.to_le_bytes());
payload.extend_from_slice(b"v1");
assert_eq!(f.broker.consume(&payload, 8, 24), 1);
assert_eq!(f.mailbox.drain()[0].value.as_ref(), b"v1");
assert_eq!(f.broker.consume(&payload, 100, 132), 1);
assert_eq!(f.broker.consume(&[1, 2, 3], 200, 208), 0);
}
#[test]
fn get_channels_matching_filters_by_glob() {
let f = fixture(4096);
f.broker.subscribe(b"apple", 1, f.mailbox.clone());
f.broker.subscribe(b"banana", 2, f.mailbox.clone());
assert_eq!(
f.broker.get_channels_matching(b"a*"),
vec![b"apple".to_vec()]
);
assert_eq!(f.broker.get_channels_matching(b"*").len(), 2);
}
#[test]
fn remove_subscription_clears_all_kinds() {
let f = fixture(4096);
f.broker.subscribe(b"ch", 1, f.mailbox.clone());
f.broker.pattern_subscribe(b"p*", 1, f.mailbox.clone());
f.broker.remove_subscription(1);
assert!(f.broker.get_channels().is_empty());
assert_eq!(f.broker.num_pattern_subscriptions(), 0);
}
#[test]
fn dispose_rejects_further_operations() {
let f = fixture(4096);
f.broker.subscribe(b"ch", 1, f.mailbox.clone());
f.broker.dispose();
assert!(!f.broker.subscribe(b"ch2", 2, f.mailbox.clone()));
assert_eq!(f.broker.publish_now(b"ch", b"v"), 0);
f.broker.publish(b"ch", b"v");
assert_eq!(f.broker.consume_pending(), 0);
assert!(f.broker.get_channels().is_empty());
}
#[test]
fn equals_on_pattern_entry() {
use super::super::pattern_subscription_entry::PatternSubscriptionEntry;
let a = PatternSubscriptionEntry::new(Box::from(b"ab*".as_slice()));
let b = PatternSubscriptionEntry::new(Box::from(b"ab*".as_slice()));
let c = PatternSubscriptionEntry::new(Box::from(b"ba*".as_slice()));
assert!(a.equals(&b));
assert!(!a.equals(&c));
}
}