use parking_lot::RwLock;
use xs_foundation::collections::thin_map::ThinMapU32;
#[derive(Debug)]
struct Edge {
label: Box<[u8]>,
child: u32,
}
#[derive(Debug, Default)]
struct Node<T> {
children: Vec<Edge>,
payload: T,
}
#[derive(Debug)]
struct Arena<T> {
nodes: Vec<Node<T>>,
}
impl<T: Default> Arena<T> {
fn new() -> Self {
Self {
nodes: vec![Node::default()],
}
}
fn insert(&mut self, topic: &[u8]) -> &mut T {
let node = self.insert_node(topic);
&mut self.nodes[node].payload
}
fn insert_node(&mut self, topic: &[u8]) -> usize {
let mut node = 0usize;
let mut rest = topic;
loop {
if rest.is_empty() {
return node;
}
match child_index(&self.nodes[node].children, rest[0]) {
Ok(ei) => {
let label = self.nodes[node].children[ei].label.clone();
let cpl = common_prefix_len(&label, rest);
if cpl == label.len() {
node = self.nodes[node].children[ei].child as usize;
rest = &rest[cpl..];
continue;
}
let old_child = self.nodes[node].children[ei].child;
let mid = alloc_node(&mut self.nodes);
self.nodes[mid].children.push(Edge {
label: Box::from(&label[cpl..]),
child: old_child,
});
self.nodes[node].children[ei] = Edge {
label: Box::from(&label[..cpl]),
child: mid as u32,
};
if cpl == rest.len() {
return mid;
} else {
let leaf = alloc_node(&mut self.nodes);
insert_child(
&mut self.nodes[mid].children,
Edge {
label: Box::from(&rest[cpl..]),
child: leaf as u32,
},
);
return leaf;
}
}
Err(pos) => {
let leaf = alloc_node(&mut self.nodes);
self.nodes[node].children.insert(
pos,
Edge {
label: Box::from(rest),
child: leaf as u32,
},
);
return leaf;
}
}
}
}
fn get_terminal_mut(&mut self, topic: &[u8]) -> Option<&mut T> {
let node = self.find_terminal(topic)?;
Some(&mut self.nodes[node].payload)
}
fn find_terminal(&self, topic: &[u8]) -> Option<usize> {
let mut node = 0usize;
let mut rest = topic;
loop {
if rest.is_empty() {
return Some(node);
}
match child_index(&self.nodes[node].children, rest[0]) {
Ok(ei) => {
let edge = &self.nodes[node].children[ei];
let l = edge.label.len();
if rest.len() >= l && rest[..l] == *edge.label {
node = edge.child as usize;
rest = &rest[l..];
} else {
return None;
}
}
Err(_) => return None,
}
}
}
fn walk<F: FnMut(&T) -> bool>(&self, topic: &[u8], mut visit: F) {
let mut node = 0usize;
let mut rest = topic;
loop {
if !visit(&self.nodes[node].payload) {
return;
}
if rest.is_empty() {
return;
}
match child_index(&self.nodes[node].children, rest[0]) {
Ok(ei) => {
let edge = &self.nodes[node].children[ei];
let l = edge.label.len();
if rest.len() >= l && rest[..l] == *edge.label {
node = edge.child as usize;
rest = &rest[l..];
} else {
return;
}
}
Err(_) => return,
}
}
}
fn payloads_mut(&mut self) -> impl Iterator<Item = &mut T> + '_ {
self.nodes.iter_mut().map(|n| &mut n.payload)
}
fn collect_terminals<F: Fn(&T) -> bool>(&self, is_terminal: F) -> Vec<Vec<u8>> {
let mut out = Vec::new();
let mut stack: Vec<(usize, Vec<u8>)> = vec![(0, Vec::new())];
while let Some((node, prefix)) = stack.pop() {
if is_terminal(&self.nodes[node].payload) {
out.push(prefix.clone());
}
for edge in &self.nodes[node].children {
let mut child_prefix = prefix.clone();
child_prefix.extend_from_slice(&edge.label);
stack.push((edge.child as usize, child_prefix));
}
}
out
}
}
#[inline]
fn alloc_node<T: Default>(nodes: &mut Vec<Node<T>>) -> usize {
nodes.push(Node::default());
nodes.len() - 1
}
#[inline]
fn child_index(children: &[Edge], first_byte: u8) -> Result<usize, usize> {
children.binary_search_by(|e| e.label[0].cmp(&first_byte))
}
#[inline]
fn insert_child(children: &mut Vec<Edge>, edge: Edge) {
let pos = child_index(children, edge.label[0]).unwrap_or_else(|p| p);
children.insert(pos, edge);
}
#[inline]
fn common_prefix_len(a: &[u8], b: &[u8]) -> usize {
a.iter().zip(b.iter()).take_while(|(x, y)| x == y).count()
}
#[derive(Debug, Default)]
struct PeerSet(ThinMapU32<u32>);
unsafe impl Send for PeerSet {}
unsafe impl Sync for PeerSet {}
impl PeerSet {
#[inline]
fn add(&mut self, peer_idx: u32) {
match self.0.get_mut(&peer_idx) {
Some(rc) => *rc += 1,
None => {
self.0.insert(peer_idx, 1);
}
}
}
#[inline]
fn remove_one(&mut self, peer_idx: u32) -> bool {
match self.0.get_mut(&peer_idx) {
Some(rc) if *rc > 1 => {
*rc -= 1;
false
}
Some(_) => {
self.0.remove(&peer_idx);
true
}
None => false,
}
}
#[inline]
fn remove_all(&mut self, peer_idx: u32) {
self.0.remove(&peer_idx);
}
#[inline]
fn for_each(&self, mut f: impl FnMut(u32)) {
for kv in self.0.iter() {
f(kv.key);
}
}
}
#[derive(Debug)]
pub(crate) struct SubscriptionMatcher {
inner: RwLock<Arena<PeerSet>>,
}
impl Default for SubscriptionMatcher {
fn default() -> Self {
Self::new()
}
}
impl SubscriptionMatcher {
pub fn new() -> Self {
Self {
inner: RwLock::new(Arena::new()),
}
}
pub fn subscribe(&self, peer_idx: u32, topic: &[u8]) {
self.inner.write().insert(topic).add(peer_idx);
}
pub fn unsubscribe(&self, peer_idx: u32, topic: &[u8]) -> bool {
match self.inner.write().get_terminal_mut(topic) {
Some(set) => set.remove_one(peer_idx),
None => false,
}
}
pub fn remove_peer(&self, peer_idx: u32) {
let mut inner = self.inner.write();
for payload in inner.payloads_mut() {
payload.remove_all(peer_idx);
}
}
pub fn for_each_match(&self, topic: &[u8], mut visit: impl FnMut(u32)) {
self.inner.read().walk(topic, |set| {
set.for_each(&mut visit);
true
});
}
}
#[derive(Debug)]
pub(crate) struct PrefixMatcher {
inner: RwLock<Arena<u32>>,
}
impl Default for PrefixMatcher {
fn default() -> Self {
Self::new()
}
}
impl PrefixMatcher {
pub fn new() -> Self {
Self {
inner: RwLock::new(Arena::new()),
}
}
pub fn subscribe(&self, topic: &[u8]) {
*self.inner.write().insert(topic) += 1;
}
pub fn unsubscribe(&self, topic: &[u8]) -> bool {
match self.inner.write().get_terminal_mut(topic) {
Some(count) if *count > 0 => {
*count -= 1;
*count == 0
}
_ => false,
}
}
pub fn matches(&self, message_topic: &[u8]) -> bool {
let mut found = false;
self.inner.read().walk(message_topic, |&count| {
if count > 0 {
found = true;
false } else {
true
}
});
found
}
pub fn get_all_topics(&self) -> Vec<Vec<u8>> {
self.inner.read().collect_terminals(|&count| count > 0)
}
}
#[cfg(test)]
mod matcher_tests {
use super::*;
fn matches(m: &SubscriptionMatcher, topic: &[u8]) -> Vec<u32> {
let mut out = Vec::new();
m.for_each_match(topic, |idx| out.push(idx));
out.sort_unstable();
out.dedup();
out
}
#[test]
fn empty_subscription_matches_everything() {
let m = SubscriptionMatcher::new();
m.subscribe(7, b"");
assert_eq!(matches(&m, b"anything"), vec![7]);
assert_eq!(matches(&m, b""), vec![7]);
}
#[test]
fn prefix_match_semantics() {
let m = SubscriptionMatcher::new();
m.subscribe(1, b"news");
assert_eq!(matches(&m, b"news/weather"), vec![1]);
assert_eq!(matches(&m, b"news"), vec![1]);
assert_eq!(matches(&m, b"new"), Vec::<u32>::new());
assert_eq!(matches(&m, b"sports"), Vec::<u32>::new());
}
#[test]
fn selective_delivery_across_peers() {
let m = SubscriptionMatcher::new();
m.subscribe(10, b"sports");
m.subscribe(20, b"news");
assert_eq!(matches(&m, b"sports/nba"), vec![10]);
assert_eq!(matches(&m, b"news/weather"), vec![20]);
assert_eq!(matches(&m, b"other"), Vec::<u32>::new());
}
#[test]
fn overlapping_prefixes_dedup_to_single_peer() {
let m = SubscriptionMatcher::new();
m.subscribe(5, b"a");
m.subscribe(5, b"ab");
let mut visits = Vec::new();
m.for_each_match(b"abc", |idx| visits.push(idx));
assert_eq!(visits, vec![5, 5], "expected one visit per matching prefix");
assert_eq!(matches(&m, b"abc"), vec![5]);
}
#[test]
fn edge_split_preserves_existing_subscriptions() {
let m = SubscriptionMatcher::new();
m.subscribe(1, b"foobar");
m.subscribe(2, b"foo");
assert_eq!(matches(&m, b"foobar"), vec![1, 2]);
assert_eq!(matches(&m, b"foobaz"), vec![2]);
assert_eq!(matches(&m, b"foo"), vec![2]);
m.subscribe(3, b"food");
assert_eq!(matches(&m, b"food"), vec![2, 3]);
assert_eq!(matches(&m, b"foobar"), vec![1, 2]);
}
#[test]
fn refcount_requires_matching_unsubscribes() {
let m = SubscriptionMatcher::new();
m.subscribe(9, b"news");
m.subscribe(9, b"news");
assert!(!m.unsubscribe(9, b"news"), "still one ref left");
assert_eq!(matches(&m, b"news"), vec![9]);
assert!(m.unsubscribe(9, b"news"), "last ref removed");
assert_eq!(matches(&m, b"news"), Vec::<u32>::new());
assert!(!m.unsubscribe(9, b"news"));
}
#[test]
fn remove_peer_purges_all_subscriptions() {
let m = SubscriptionMatcher::new();
m.subscribe(1, b"");
m.subscribe(1, b"a/b/c");
m.subscribe(2, b"a/b/c");
m.remove_peer(1);
assert_eq!(matches(&m, b"a/b/c/d"), vec![2]);
assert_eq!(matches(&m, b"unrelated"), Vec::<u32>::new());
m.subscribe(1, b"fresh");
assert_eq!(matches(&m, b"a/b/c"), vec![2]);
assert_eq!(matches(&m, b"fresh/topic"), vec![1]);
}
}
#[cfg(test)]
mod prefix_matcher_tests {
use super::*;
#[test]
fn empty_subscription_matches_all() {
let m = PrefixMatcher::new();
m.subscribe(b"");
assert!(m.matches(b"sports/football"));
assert!(m.matches(b"news/weather"));
assert!(m.matches(b""));
}
#[test]
fn prefix_and_divergent() {
let m = PrefixMatcher::new();
m.subscribe(b"news");
assert!(m.matches(b"news"));
assert!(m.matches(b"news/weather"));
assert!(!m.matches(b"new"));
assert!(!m.matches(b"sports"));
}
#[test]
fn overlapping_refcount_and_unsubscribe() {
let m = PrefixMatcher::new();
m.subscribe(b"news");
m.subscribe(b"news");
assert!(!m.unsubscribe(b"news"), "count still 1");
assert!(m.matches(b"news"));
assert!(m.unsubscribe(b"news"), "count reached zero");
assert!(!m.matches(b"news"));
assert!(!m.unsubscribe(b"news"));
}
#[test]
fn get_all_topics_roundtrip() {
let m = PrefixMatcher::new();
m.subscribe(b"sports");
m.subscribe(b"sports/football");
m.subscribe(b"news/weather");
let mut topics = m.get_all_topics();
topics.sort();
let mut expected = vec![
b"news/weather".to_vec(),
b"sports".to_vec(),
b"sports/football".to_vec(),
];
expected.sort();
assert_eq!(topics, expected);
}
#[test]
fn empty_subscription_appears_in_topics() {
let m = PrefixMatcher::new();
m.subscribe(b"");
assert_eq!(m.get_all_topics(), vec![Vec::<u8>::new()]);
}
}