use std::collections::{HashMap, HashSet};
use std::fmt;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use dashmap::{DashMap, DashSet};
use crate::ids::{ConnectionId, SessionId, StreamId};
#[derive(Default)]
pub struct SessionSubscriptions {
inner: DashMap<(ConnectionId, StreamId), Arc<DashSet<ConnectionId>>>,
}
impl SessionSubscriptions {
pub fn new() -> Self {
Self::default()
}
pub fn add(&self, publisher: ConnectionId, strm_id: StreamId, subscriber: ConnectionId) {
let entry = self
.inner
.entry((publisher, strm_id))
.or_insert_with(|| Arc::new(DashSet::new()))
.clone();
entry.insert(subscriber);
}
pub fn remove(
&self,
publisher: &ConnectionId,
strm_id: &StreamId,
subscriber: &ConnectionId,
) -> bool {
let key = (publisher.clone(), strm_id.clone());
let removed = if let Some(entry) = self.inner.get(&key) {
entry.remove(subscriber).is_some()
} else {
false
};
removed
}
pub fn subscribers_for(
&self,
publisher: &ConnectionId,
strm_id: &StreamId,
) -> Vec<ConnectionId> {
match self.inner.get(&(publisher.clone(), strm_id.clone())) {
Some(set) => set.iter().map(|e| e.clone()).collect(),
None => Vec::new(),
}
}
pub fn drop_connection(&self, connid: &ConnectionId) {
let keys: Vec<(ConnectionId, StreamId)> =
self.inner.iter().map(|e| e.key().clone()).collect();
for key in keys {
if key.0 == *connid {
self.inner.remove(&key);
continue;
}
if let Some(set) = self.inner.get(&key) {
set.remove(connid);
}
}
}
pub fn rows(&self) -> Vec<(ConnectionId, StreamId, Vec<ConnectionId>)> {
self.inner
.iter()
.map(|e| {
let (pub_id, strm) = e.key().clone();
let subs = e.value().iter().map(|s| s.clone()).collect();
(pub_id, strm, subs)
})
.collect()
}
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
pub fn drop_publisher_stream(&self, publisher: &ConnectionId, stream: &StreamId) {
self.inner.remove(&(publisher.clone(), stream.clone()));
}
}
pub struct SubscriptionRegistry {
sessions: DashMap<SessionId, Arc<SessionSubscriptions>>,
max_direct_subscribers: usize,
direct: Mutex<DirectSubscriptionState>,
}
#[derive(Clone, Eq, Hash, PartialEq)]
struct DirectSubscriptionRoute {
sid: SessionId,
publisher: ConnectionId,
stream: StreamId,
}
#[derive(Default)]
struct DirectSubscriptionState {
listeners: HashMap<ConnectionId, HashSet<DirectSubscriptionRoute>>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum DirectSubscriptionAdmissionError {
Disabled,
CapacityExhausted,
}
impl fmt::Display for DirectSubscriptionAdmissionError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(match self {
Self::Disabled => "direct subscriptions are disabled",
Self::CapacityExhausted => "direct subscriber capacity is exhausted",
})
}
}
impl std::error::Error for DirectSubscriptionAdmissionError {}
impl Default for SubscriptionRegistry {
fn default() -> Self {
Self::with_direct_listener_limit(1_000)
}
}
impl SubscriptionRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn with_direct_listener_limit(max_direct_subscribers: usize) -> Self {
Self {
sessions: DashMap::new(),
max_direct_subscribers,
direct: Mutex::new(DirectSubscriptionState::default()),
}
}
pub fn direct_listener_limit(&self) -> usize {
self.max_direct_subscribers
}
pub fn for_session(&self, sid: &SessionId) -> Arc<SessionSubscriptions> {
self.sessions
.entry(sid.clone())
.or_insert_with(|| Arc::new(SessionSubscriptions::new()))
.clone()
}
pub fn try_add_direct(
&self,
sid: &SessionId,
subscriber: &ConnectionId,
routes: &[(ConnectionId, StreamId)],
) -> Result<(), DirectSubscriptionAdmissionError> {
if routes.is_empty() {
return Ok(());
}
if self.max_direct_subscribers == 0 {
return Err(DirectSubscriptionAdmissionError::Disabled);
}
let mut direct = self
.direct
.lock()
.expect("direct subscription lock poisoned");
if !direct.listeners.contains_key(subscriber)
&& direct.listeners.len() >= self.max_direct_subscribers
{
return Err(DirectSubscriptionAdmissionError::CapacityExhausted);
}
let table = self.for_session(sid);
let subscriber_routes = direct.listeners.entry(subscriber.clone()).or_default();
for (publisher, stream) in routes {
table.add(publisher.clone(), stream.clone(), subscriber.clone());
subscriber_routes.insert(DirectSubscriptionRoute {
sid: sid.clone(),
publisher: publisher.clone(),
stream: stream.clone(),
});
}
Ok(())
}
pub fn remove_direct(
&self,
sid: &SessionId,
subscriber: &ConnectionId,
publisher: &ConnectionId,
stream: &StreamId,
) -> bool {
let mut direct = self
.direct
.lock()
.expect("direct subscription lock poisoned");
let removed = self.for_session(sid).remove(publisher, stream, subscriber);
if let Some(routes) = direct.listeners.get_mut(subscriber) {
routes.remove(&DirectSubscriptionRoute {
sid: sid.clone(),
publisher: publisher.clone(),
stream: stream.clone(),
});
if routes.is_empty() {
direct.listeners.remove(subscriber);
}
}
removed
}
pub fn active_direct_listener_count(&self) -> usize {
self.direct
.lock()
.expect("direct subscription lock poisoned")
.listeners
.len()
}
pub fn drop_publisher_stream(
&self,
sid: &SessionId,
publisher: &ConnectionId,
stream: &StreamId,
) {
let mut direct = self
.direct
.lock()
.expect("direct subscription lock poisoned");
direct.listeners.retain(|_, routes| {
routes.remove(&DirectSubscriptionRoute {
sid: sid.clone(),
publisher: publisher.clone(),
stream: stream.clone(),
});
!routes.is_empty()
});
self.for_session(sid)
.drop_publisher_stream(publisher, stream);
}
pub fn drop_session(&self, sid: &SessionId) {
let mut direct = self
.direct
.lock()
.expect("direct subscription lock poisoned");
direct.listeners.retain(|_, routes| {
routes.retain(|route| &route.sid != sid);
!routes.is_empty()
});
self.sessions.remove(sid);
}
pub fn drop_connection(&self, connid: &ConnectionId) {
let mut direct = self
.direct
.lock()
.expect("direct subscription lock poisoned");
direct.listeners.remove(connid);
direct.listeners.retain(|_, routes| {
routes.retain(|route| &route.publisher != connid);
!routes.is_empty()
});
let sids: Vec<SessionId> = self.sessions.iter().map(|e| e.key().clone()).collect();
for sid in sids {
if let Some(table_ref) = self.sessions.get(&sid) {
let table = Arc::clone(table_ref.value());
drop(table_ref);
table.drop_connection(connid);
}
}
}
}
#[derive(Clone)]
pub struct PublisherEntry {
pub connection: ConnectionId,
pub participant: String,
pub kind: String,
pub codec: Option<crate::capability::CodecInfo>,
}
impl fmt::Debug for PublisherEntry {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("PublisherEntry")
.field("connection", &self.connection)
.field("participant_present", &!self.participant.is_empty())
.field("participant_bytes", &self.participant.len())
.field("kind_present", &!self.kind.is_empty())
.field("kind_bytes", &self.kind.len())
.field("codec_present", &self.codec.is_some())
.finish()
}
}
pub struct PublisherRegistry {
inner: DashMap<(SessionId, String), PublisherRecord>,
by_participant: DashMap<(SessionId, String), Vec<String>>,
mutation_lock: Mutex<()>,
next_registration_id: AtomicU64,
}
#[derive(Clone, Debug)]
struct PublisherRecord {
entry: PublisherEntry,
registration_id: PublisherRegistrationId,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct PublisherRegistrationId(u64);
impl Default for PublisherRegistry {
fn default() -> Self {
Self {
inner: DashMap::new(),
by_participant: DashMap::new(),
mutation_lock: Mutex::new(()),
next_registration_id: AtomicU64::new(1),
}
}
}
impl PublisherRegistry {
pub fn new() -> Self {
Self::default()
}
pub fn register(&self, sid: SessionId, strm_id: String, entry: PublisherEntry) {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let registration_id = self.next_registration_id();
let participant_key = (sid.clone(), entry.participant.clone());
if let Some(previous) = self.inner.insert(
(sid.clone(), strm_id.clone()),
PublisherRecord {
entry,
registration_id,
},
) {
self.remove_participant_stream(&sid, &previous.entry.participant, &strm_id);
}
self.add_participant_stream(participant_key, strm_id);
}
pub(crate) fn register_managed(
&self,
sid: SessionId,
strm_id: String,
entry: PublisherEntry,
) -> std::result::Result<PublisherRegistrationId, PublisherEntry> {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let key = (sid.clone(), strm_id.clone());
if let Some(existing) = self.inner.get(&key) {
return Err(existing.entry.clone());
}
let registration_id = self.next_registration_id();
let participant_key = (sid, entry.participant.clone());
self.inner.insert(
key,
PublisherRecord {
entry,
registration_id,
},
);
self.add_participant_stream(participant_key, strm_id);
Ok(registration_id)
}
pub(crate) fn registration_is_current(
&self,
sid: &SessionId,
strm_id: &str,
registration_id: PublisherRegistrationId,
) -> bool {
self.inner
.get(&(sid.clone(), strm_id.to_string()))
.is_some_and(|record| record.registration_id == registration_id)
}
pub(crate) fn remove_registration(
&self,
sid: &SessionId,
strm_id: &str,
registration_id: PublisherRegistrationId,
) -> bool {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let key = (sid.clone(), strm_id.to_string());
let Some(record) = self.inner.get(&key) else {
return false;
};
if record.registration_id != registration_id {
return false;
}
drop(record);
let Some((_, removed)) = self.inner.remove(&key) else {
return false;
};
self.remove_participant_stream(sid, &removed.entry.participant, strm_id);
true
}
fn next_registration_id(&self) -> PublisherRegistrationId {
PublisherRegistrationId(self.next_registration_id.fetch_add(1, Ordering::Relaxed))
}
fn add_participant_stream(&self, participant_key: (SessionId, String), strm_id: String) {
self.by_participant
.entry(participant_key)
.and_modify(|v| {
if !v.iter().any(|s| s == &strm_id) {
v.push(strm_id.clone());
}
})
.or_insert_with(|| vec![strm_id]);
}
fn remove_participant_stream(&self, sid: &SessionId, participant: &str, strm_id: &str) {
let participant_key = (sid.clone(), participant.to_string());
if let Some(mut streams) = self.by_participant.get_mut(&participant_key) {
streams.retain(|stream| stream != strm_id);
}
let empty = self
.by_participant
.get(&participant_key)
.is_some_and(|streams| streams.is_empty());
if empty {
self.by_participant.remove(&participant_key);
}
}
pub fn publisher(&self, sid: &SessionId, strm_id: &str) -> Option<ConnectionId> {
self.inner
.get(&(sid.clone(), strm_id.to_string()))
.map(|record| record.entry.connection.clone())
}
pub fn entry(&self, sid: &SessionId, strm_id: &str) -> Option<PublisherEntry> {
self.inner
.get(&(sid.clone(), strm_id.to_string()))
.map(|record| record.entry.clone())
}
pub fn streams_for_participant(&self, sid: &SessionId, participant: &str) -> Vec<String> {
self.by_participant
.get(&(sid.clone(), participant.to_string()))
.map(|e| e.value().clone())
.unwrap_or_default()
}
pub fn with_current_routes<R>(
&self,
sid: &SessionId,
routes: &[(ConnectionId, StreamId)],
operation: impl FnOnce() -> R,
) -> Option<R> {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let all_current = routes.iter().all(|(publisher, stream)| {
self.inner
.get(&(sid.clone(), stream.as_str().to_string()))
.is_some_and(|record| &record.entry.connection == publisher)
});
all_current.then(operation)
}
pub fn remove_stream(&self, sid: &SessionId, strm_id: &str) {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let key = (sid.clone(), strm_id.to_string());
let Some((_, removed)) = self.inner.remove(&key) else {
return;
};
self.remove_participant_stream(sid, &removed.entry.participant, strm_id);
}
pub fn remove_stream_if_publisher(
&self,
sid: &SessionId,
strm_id: &str,
publisher: &ConnectionId,
) -> bool {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let key = (sid.clone(), strm_id.to_string());
let belongs_to_publisher = self
.inner
.get(&key)
.is_some_and(|record| &record.entry.connection == publisher);
if !belongs_to_publisher {
return false;
}
let Some((_, removed)) = self.inner.remove(&key) else {
return false;
};
self.remove_participant_stream(sid, &removed.entry.participant, strm_id);
true
}
pub fn drop_publisher(&self, connid: &ConnectionId) {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let to_remove: Vec<(SessionId, String, PublisherRegistrationId)> = self
.inner
.iter()
.filter(|record| &record.entry.connection == connid)
.map(|e| {
let (sid, strm) = e.key();
(sid.clone(), strm.clone(), e.registration_id)
})
.collect();
for (sid, strm, registration_id) in to_remove {
let key = (sid.clone(), strm.clone());
let is_same_registration = self
.inner
.get(&key)
.is_some_and(|record| record.registration_id == registration_id);
if is_same_registration {
if let Some((_, removed)) = self.inner.remove(&key) {
self.remove_participant_stream(&sid, &removed.entry.participant, &strm);
}
}
}
}
pub fn drop_session(&self, sid: &SessionId) {
let _guard = self
.mutation_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
self.inner.retain(|(s, _), _| s != sid);
self.by_participant.retain(|(s, _), _| s != sid);
}
}
#[cfg(test)]
mod publisher_registry_tests {
use super::*;
use std::time::Duration;
#[test]
fn publisher_entry_debug_redacts_participant_kind_and_codec_values() {
const CANARY: &str = "publisher-canary\r\nAuthorization: exposed";
let entry = PublisherEntry {
connection: ConnectionId::from_string(CANARY),
participant: CANARY.into(),
kind: CANARY.into(),
codec: Some(crate::capability::CodecInfo {
name: CANARY.into(),
clock_rate_hz: 48_000,
channels: 1,
fmtp: Some(CANARY.into()),
}),
};
let debug = format!("{entry:?}");
assert!(!debug.contains(CANARY));
assert!(debug.contains("codec_present: true"));
}
#[test]
fn remove_stream_updates_primary_and_participant_indexes() {
let registry = PublisherRegistry::new();
let sid = SessionId::new();
let connection = ConnectionId::new();
for stream in ["audio-main", "audio-backup"] {
registry.register(
sid.clone(),
stream.to_string(),
PublisherEntry {
connection: connection.clone(),
participant: "alice".to_string(),
kind: "audio".to_string(),
codec: None,
},
);
}
registry.remove_stream(&sid, "audio-main");
assert!(registry.entry(&sid, "audio-main").is_none());
assert_eq!(
registry.streams_for_participant(&sid, "alice"),
vec!["audio-backup".to_string()]
);
registry.remove_stream(&sid, "audio-backup");
assert!(registry.streams_for_participant(&sid, "alice").is_empty());
registry.remove_stream(&sid, "audio-backup");
}
#[test]
fn conditional_remove_cannot_delete_same_named_replacement() {
let registry = PublisherRegistry::new();
let sid = SessionId::new();
let old_publisher = ConnectionId::new();
let replacement = ConnectionId::new();
registry.register(
sid.clone(),
"audio-main".to_string(),
PublisherEntry {
connection: old_publisher.clone(),
participant: "old".to_string(),
kind: "audio".to_string(),
codec: None,
},
);
registry.register(
sid.clone(),
"audio-main".to_string(),
PublisherEntry {
connection: replacement.clone(),
participant: "replacement".to_string(),
kind: "audio".to_string(),
codec: None,
},
);
assert!(!registry.remove_stream_if_publisher(&sid, "audio-main", &old_publisher,));
assert_eq!(
registry.publisher(&sid, "audio-main"),
Some(replacement.clone())
);
assert!(registry.remove_stream_if_publisher(&sid, "audio-main", &replacement,));
assert!(registry.entry(&sid, "audio-main").is_none());
assert!(registry
.streams_for_participant(&sid, "replacement")
.is_empty());
}
#[test]
fn subscription_drop_connection_releases_map_guard_without_outer_row_race() {
let registry = Arc::new(SubscriptionRegistry::new());
let sid = SessionId::new();
let publisher = ConnectionId::new();
let subscriber = ConnectionId::new();
registry
.for_session(&sid)
.add(publisher, StreamId::new(), subscriber.clone());
let (done_tx, done_rx) = std::sync::mpsc::channel();
let registry_for_thread = Arc::clone(®istry);
std::thread::spawn(move || {
registry_for_thread.drop_connection(&subscriber);
let _ = done_tx.send(());
});
done_rx
.recv_timeout(Duration::from_secs(1))
.expect("drop_connection must not self-deadlock while removing an empty session");
let table = registry.for_session(&sid);
assert!(table
.rows()
.iter()
.all(|(_, _, subscribers)| subscribers.is_empty()));
registry.drop_session(&sid);
assert!(registry.for_session(&sid).is_empty());
}
}