use std::any::TypeId;
use std::collections::{HashMap, HashSet};
use std::sync::atomic::{AtomicUsize, Ordering};
use dashmap::DashMap;
use parking_lot::RwLock;
use tokio::sync::mpsc;
use tracing::{debug, trace, warn};
use super::types::IpcPushNotification;
fn matches_pattern(message_type: &str, pattern: &str) -> bool {
if !pattern.contains('*') {
return message_type == pattern;
}
let parts: Vec<&str> = pattern.split('*').collect();
match parts.as_slice() {
["", ""] => true, ["", suffix] => message_type.ends_with(suffix), [prefix, ""] => message_type.starts_with(prefix), [prefix, suffix] => {
message_type.starts_with(prefix)
&& message_type.ends_with(suffix)
&& message_type.len() >= prefix.len() + suffix.len()
}
_ => false, }
}
pub type ConnectionId = usize;
pub type PushSender = mpsc::Sender<IpcPushNotification>;
#[derive(Debug, Default)]
pub struct SubscriptionStats {
pub subscriptions_added: AtomicUsize,
pub subscriptions_removed: AtomicUsize,
pub push_notifications_sent: AtomicUsize,
pub push_notifications_dropped: AtomicUsize,
}
impl SubscriptionStats {
#[must_use]
pub fn subscriptions_added(&self) -> usize {
self.subscriptions_added.load(Ordering::Relaxed)
}
#[must_use]
pub fn subscriptions_removed(&self) -> usize {
self.subscriptions_removed.load(Ordering::Relaxed)
}
#[must_use]
pub fn push_notifications_sent(&self) -> usize {
self.push_notifications_sent.load(Ordering::Relaxed)
}
#[must_use]
pub fn push_notifications_dropped(&self) -> usize {
self.push_notifications_dropped.load(Ordering::Relaxed)
}
}
struct ConnectionInfo {
push_sender: PushSender,
subscribed_types: HashSet<String>,
}
pub struct SubscriptionManager {
connections: DashMap<ConnectionId, ConnectionInfo>,
type_to_connections: DashMap<String, HashSet<ConnectionId>>,
type_id_to_name: RwLock<HashMap<TypeId, String>>,
stats: SubscriptionStats,
}
impl Default for SubscriptionManager {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for SubscriptionManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SubscriptionManager")
.field("connection_count", &self.connections.len())
.field("subscribed_types_count", &self.type_to_connections.len())
.field("type_id_mappings", &self.type_id_to_name.read().len())
.field("stats", &self.stats)
.finish()
}
}
impl SubscriptionManager {
#[must_use]
pub fn new() -> Self {
Self {
connections: DashMap::new(),
type_to_connections: DashMap::new(),
type_id_to_name: RwLock::new(HashMap::new()),
stats: SubscriptionStats::default(),
}
}
#[must_use]
pub const fn stats(&self) -> &SubscriptionStats {
&self.stats
}
pub fn register_connection(&self, conn_id: ConnectionId, push_sender: PushSender) {
trace!(conn_id, "Registering connection for subscriptions");
self.connections.insert(
conn_id,
ConnectionInfo {
push_sender,
subscribed_types: HashSet::new(),
},
);
}
pub fn unregister_connection(&self, conn_id: ConnectionId) {
if let Some((_, info)) = self.connections.remove(&conn_id) {
for type_name in &info.subscribed_types {
if let Some(mut entry) = self.type_to_connections.get_mut(type_name) {
entry.remove(&conn_id);
if entry.is_empty() {
drop(entry);
self.type_to_connections.remove(type_name);
}
}
self.stats
.subscriptions_removed
.fetch_add(1, Ordering::Relaxed);
}
debug!(
conn_id,
removed_subscriptions = info.subscribed_types.len(),
"Unregistered connection and removed subscriptions"
);
}
}
pub fn subscribe(&self, conn_id: ConnectionId, message_types: &[String]) -> Vec<String> {
let Some(mut conn_entry) = self.connections.get_mut(&conn_id) else {
warn!(conn_id, "Cannot subscribe: connection not registered");
return Vec::new();
};
for type_name in message_types {
if conn_entry.subscribed_types.insert(type_name.clone()) {
self.type_to_connections
.entry(type_name.clone())
.or_default()
.insert(conn_id);
self.stats
.subscriptions_added
.fetch_add(1, Ordering::Relaxed);
trace!(conn_id, message_type = %type_name, "Added subscription");
}
}
conn_entry.subscribed_types.iter().cloned().collect()
}
pub fn unsubscribe(&self, conn_id: ConnectionId, message_types: &[String]) -> Vec<String> {
let Some(mut conn_entry) = self.connections.get_mut(&conn_id) else {
warn!(conn_id, "Cannot unsubscribe: connection not registered");
return Vec::new();
};
if message_types.is_empty() {
let types_to_remove: Vec<_> = conn_entry.subscribed_types.drain().collect();
for type_name in &types_to_remove {
if let Some(mut entry) = self.type_to_connections.get_mut(type_name) {
entry.remove(&conn_id);
if entry.is_empty() {
drop(entry);
self.type_to_connections.remove(type_name);
}
}
self.stats
.subscriptions_removed
.fetch_add(1, Ordering::Relaxed);
}
trace!(
conn_id,
count = types_to_remove.len(),
"Unsubscribed from all types"
);
return Vec::new();
}
for type_name in message_types {
if conn_entry.subscribed_types.remove(type_name) {
if let Some(mut entry) = self.type_to_connections.get_mut(type_name) {
entry.remove(&conn_id);
if entry.is_empty() {
drop(entry);
self.type_to_connections.remove(type_name);
}
}
self.stats
.subscriptions_removed
.fetch_add(1, Ordering::Relaxed);
trace!(conn_id, message_type = %type_name, "Removed subscription");
}
}
conn_entry.subscribed_types.iter().cloned().collect()
}
#[must_use]
pub fn get_subscriptions(&self, conn_id: ConnectionId) -> Vec<String> {
self.connections
.get(&conn_id)
.map(|entry| entry.subscribed_types.iter().cloned().collect())
.unwrap_or_default()
}
pub fn register_type_mapping(&self, type_id: TypeId, type_name: String) {
let mut map = self.type_id_to_name.write();
map.insert(type_id, type_name);
}
#[must_use]
pub fn get_type_name(&self, type_id: &TypeId) -> Option<String> {
let map = self.type_id_to_name.read();
map.get(type_id).cloned()
}
pub fn forward_to_subscribers(&self, notification: &IpcPushNotification) {
let message_type = ¬ification.message_type;
let mut matched_connections: HashSet<ConnectionId> = HashSet::new();
if let Some(connections_entry) = self.type_to_connections.get(message_type) {
matched_connections.extend(connections_entry.iter().copied());
}
for entry in self.type_to_connections.iter() {
let pattern = entry.key();
if pattern.contains('*') && matches_pattern(message_type, pattern) {
matched_connections.extend(entry.value().iter().copied());
}
}
if matched_connections.is_empty() {
trace!(message_type, "No subscribers for message type");
return;
}
for conn_id in matched_connections {
if let Some(conn_info) = self.connections.get(&conn_id) {
let notification_clone = notification.clone();
match conn_info.push_sender.try_send(notification_clone) {
Ok(()) => {
self.stats
.push_notifications_sent
.fetch_add(1, Ordering::Relaxed);
trace!(conn_id, message_type, "Forwarded push notification");
}
Err(mpsc::error::TrySendError::Full(_)) => {
self.stats
.push_notifications_dropped
.fetch_add(1, Ordering::Relaxed);
warn!(
conn_id,
message_type, "Push channel full, dropping notification"
);
}
Err(mpsc::error::TrySendError::Closed(_)) => {
self.stats
.push_notifications_dropped
.fetch_add(1, Ordering::Relaxed);
trace!(conn_id, message_type, "Push channel closed");
}
}
}
}
}
#[must_use]
pub fn connection_count(&self) -> usize {
self.connections.len()
}
#[must_use]
pub fn subscribed_types_count(&self) -> usize {
self.type_to_connections.len()
}
#[must_use]
pub fn total_subscriptions(&self) -> usize {
self.type_to_connections
.iter()
.map(|entry| entry.value().len())
.sum()
}
}
pub struct PushReceiver {
#[allow(dead_code)]
pub conn_id: ConnectionId,
pub receiver: mpsc::Receiver<IpcPushNotification>,
}
#[must_use]
pub fn create_push_channel(
conn_id: ConnectionId,
buffer_size: usize,
) -> (PushSender, PushReceiver) {
let (sender, receiver) = mpsc::channel(buffer_size);
(sender, PushReceiver { conn_id, receiver })
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[test]
fn test_subscription_manager_new() {
let manager = SubscriptionManager::new();
assert_eq!(manager.connection_count(), 0);
assert_eq!(manager.subscribed_types_count(), 0);
}
#[test]
fn test_register_unregister_connection() {
let manager = SubscriptionManager::new();
let (sender, _receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
assert_eq!(manager.connection_count(), 1);
manager.unregister_connection(1);
assert_eq!(manager.connection_count(), 0);
}
#[test]
fn test_subscribe_unsubscribe() {
let manager = SubscriptionManager::new();
let (sender, _receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
let subscribed = manager.subscribe(1, &["TypeA".to_string(), "TypeB".to_string()]);
assert_eq!(subscribed.len(), 2);
assert!(subscribed.contains(&"TypeA".to_string()));
assert!(subscribed.contains(&"TypeB".to_string()));
assert_eq!(manager.subscribed_types_count(), 2);
assert_eq!(manager.total_subscriptions(), 2);
let subscribed = manager.unsubscribe(1, &["TypeA".to_string()]);
assert_eq!(subscribed.len(), 1);
assert!(subscribed.contains(&"TypeB".to_string()));
assert_eq!(manager.subscribed_types_count(), 1);
assert_eq!(manager.total_subscriptions(), 1);
let subscribed = manager.unsubscribe(1, &[]);
assert!(subscribed.is_empty());
assert_eq!(manager.subscribed_types_count(), 0);
}
#[test]
fn test_unregister_cleans_subscriptions() {
let manager = SubscriptionManager::new();
let (sender, _receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(1, &["TypeA".to_string(), "TypeB".to_string()]);
assert_eq!(manager.subscribed_types_count(), 2);
manager.unregister_connection(1);
assert_eq!(manager.subscribed_types_count(), 0);
}
#[test]
fn test_multiple_connections_same_type() {
let manager = SubscriptionManager::new();
let (sender1, _receiver1) = mpsc::channel(10);
let (sender2, _receiver2) = mpsc::channel(10);
manager.register_connection(1, sender1);
manager.register_connection(2, sender2);
manager.subscribe(1, &["TypeA".to_string()]);
manager.subscribe(2, &["TypeA".to_string()]);
assert_eq!(manager.subscribed_types_count(), 1);
assert_eq!(manager.total_subscriptions(), 2);
manager.unregister_connection(1);
assert_eq!(manager.subscribed_types_count(), 1);
assert_eq!(manager.total_subscriptions(), 1);
manager.unregister_connection(2);
assert_eq!(manager.subscribed_types_count(), 0);
}
#[tokio::test]
async fn test_forward_to_subscribers() {
let manager = Arc::new(SubscriptionManager::new());
let (sender, mut receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(1, &["PriceUpdate".to_string()]);
let notification = IpcPushNotification::new(
"PriceUpdate",
Some("price_service".to_string()),
serde_json::json!({ "price": 100.0 }),
);
manager.forward_to_subscribers(¬ification);
let notification_out = receiver.try_recv().unwrap();
assert_eq!(notification_out.message_type, "PriceUpdate");
assert_eq!(manager.stats().push_notifications_sent(), 1);
}
#[test]
fn test_forward_no_subscribers() {
let manager = Arc::new(SubscriptionManager::new());
let notification =
IpcPushNotification::new("UnsubscribedType", None, serde_json::json!({}));
manager.forward_to_subscribers(¬ification);
assert_eq!(manager.stats().push_notifications_sent(), 0);
}
#[test]
fn test_type_mapping() {
struct TestMessage;
let manager = SubscriptionManager::new();
let type_id = TypeId::of::<TestMessage>();
manager.register_type_mapping(type_id, "TestMessage".to_string());
assert_eq!(
manager.get_type_name(&type_id),
Some("TestMessage".to_string())
);
}
#[tokio::test]
async fn test_create_push_channel() {
let conn_id = 42;
let buffer_size = 10;
let (sender, receiver) = create_push_channel(conn_id, buffer_size);
assert_eq!(receiver.conn_id, conn_id);
let notification = IpcPushNotification::new(
"TestMessage",
Some("test_actor".to_string()),
serde_json::json!({ "test": true }),
);
sender.send(notification.clone()).await.unwrap();
let mut channel = receiver.receiver;
let msg = channel.recv().await.unwrap();
assert_eq!(msg.message_type, "TestMessage");
}
#[test]
fn test_push_receiver_struct() {
let (_, receiver) = create_push_channel(123, 5);
assert_eq!(receiver.conn_id, 123);
}
#[test]
fn test_matches_pattern_exact() {
assert!(matches_pattern("timer.tick", "timer.tick"));
assert!(!matches_pattern("timer.tick", "timer.tock"));
assert!(!matches_pattern("timer.tick", "timer.tick.extra"));
}
#[test]
fn test_matches_pattern_prefix_wildcard() {
assert!(matches_pattern("system.started.timer", "system.started.*"));
assert!(matches_pattern("system.started.sink", "system.started.*"));
assert!(matches_pattern("system.started.", "system.started.*"));
assert!(!matches_pattern("system.stopped.timer", "system.started.*"));
assert!(!matches_pattern("system.start", "system.started.*"));
}
#[test]
fn test_matches_pattern_suffix_wildcard() {
assert!(matches_pattern("foo.bar.event", "*.event"));
assert!(matches_pattern(".event", "*.event")); assert!(!matches_pattern("event", "*.event")); assert!(!matches_pattern("foo.bar.message", "*.event"));
assert!(matches_pattern("event", "*event"));
assert!(matches_pattern("myevent", "*event"));
assert!(matches_pattern("foo.bar.event", "*event"));
}
#[test]
fn test_matches_pattern_match_all() {
assert!(matches_pattern("anything", "*"));
assert!(matches_pattern("", "*"));
assert!(matches_pattern("a.b.c.d.e", "*"));
}
#[test]
fn test_matches_pattern_prefix_and_suffix() {
assert!(matches_pattern("system.timer.event", "system.*.event"));
assert!(matches_pattern("system..event", "system.*.event"));
assert!(!matches_pattern("system.timer.message", "system.*.event"));
assert!(!matches_pattern("other.timer.event", "system.*.event"));
}
#[test]
fn test_matches_pattern_multiple_wildcards_not_supported() {
assert!(!matches_pattern("a.b.c", "a*b*c"));
assert!(!matches_pattern("abc", "*.*.*"));
}
#[tokio::test]
async fn test_forward_with_wildcard_subscription() {
let manager = Arc::new(SubscriptionManager::new());
let (sender, mut receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(1, &["system.started.*".to_string()]);
let notification = IpcPushNotification::new(
"system.started.timer",
Some("engine".to_string()),
serde_json::json!({ "name": "timer" }),
);
manager.forward_to_subscribers(¬ification);
let received = receiver.try_recv().unwrap();
assert_eq!(received.message_type, "system.started.timer");
assert_eq!(manager.stats().push_notifications_sent(), 1);
}
#[tokio::test]
async fn test_forward_wildcard_no_match() {
let manager = Arc::new(SubscriptionManager::new());
let (sender, mut receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(1, &["system.started.*".to_string()]);
let notification = IpcPushNotification::new(
"system.stopped.timer",
Some("engine".to_string()),
serde_json::json!({ "name": "timer" }),
);
manager.forward_to_subscribers(¬ification);
assert!(receiver.try_recv().is_err());
assert_eq!(manager.stats().push_notifications_sent(), 0);
}
#[tokio::test]
async fn test_forward_exact_and_wildcard_combined() {
let manager = Arc::new(SubscriptionManager::new());
let (sender, mut receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(
1,
&[
"timer.tick".to_string(),
"system.started.*".to_string(),
],
);
let notification1 = IpcPushNotification::new(
"timer.tick",
Some("timer".to_string()),
serde_json::json!({ "seq": 1 }),
);
manager.forward_to_subscribers(¬ification1);
let notification2 = IpcPushNotification::new(
"system.started.sink",
Some("engine".to_string()),
serde_json::json!({ "name": "sink" }),
);
manager.forward_to_subscribers(¬ification2);
let msg1 = receiver.try_recv().unwrap();
assert_eq!(msg1.message_type, "timer.tick");
let msg2 = receiver.try_recv().unwrap();
assert_eq!(msg2.message_type, "system.started.sink");
assert_eq!(manager.stats().push_notifications_sent(), 2);
}
#[tokio::test]
async fn test_forward_match_all_wildcard() {
let manager = Arc::new(SubscriptionManager::new());
let (sender, mut receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(1, &["*".to_string()]);
let notification = IpcPushNotification::new(
"anything.at.all",
Some("source".to_string()),
serde_json::json!({}),
);
manager.forward_to_subscribers(¬ification);
let received = receiver.try_recv().unwrap();
assert_eq!(received.message_type, "anything.at.all");
}
#[tokio::test]
async fn test_forward_no_duplicate_delivery() {
let manager = Arc::new(SubscriptionManager::new());
let (sender, mut receiver) = mpsc::channel(10);
manager.register_connection(1, sender);
manager.subscribe(
1,
&[
"system.started.timer".to_string(), "system.started.*".to_string(), ],
);
let notification = IpcPushNotification::new(
"system.started.timer",
Some("engine".to_string()),
serde_json::json!({ "name": "timer" }),
);
manager.forward_to_subscribers(¬ification);
let _ = receiver.try_recv().unwrap();
assert!(receiver.try_recv().is_err());
assert_eq!(manager.stats().push_notifications_sent(), 1);
}
}