use std::{
collections::HashMap,
sync::atomic::{AtomicU64, Ordering},
};
use tokio::{
sync::{RwLock, broadcast},
time::Instant,
};
use tracing::debug;
#[derive(Debug, Clone)]
pub struct BroadcastConfig {
pub channel_capacity: usize,
pub max_channels: usize,
pub max_message_bytes: usize,
}
impl BroadcastConfig {
#[must_use]
pub const fn new() -> Self {
Self {
channel_capacity: 128,
max_channels: 1_000,
max_message_bytes: 65_536,
}
}
}
impl Default for BroadcastConfig {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct BroadcastStats {
pub messages_published: u64,
pub active_channels: usize,
pub active_receivers: usize,
}
#[derive(Debug)]
struct BroadcastChannel {
sender: broadcast::Sender<BroadcastMessage>,
created_at: Instant,
}
#[derive(Debug, Clone)]
pub struct BroadcastMessage {
pub channel: String,
pub event: String,
pub payload: serde_json::Value,
}
#[derive(Debug)]
pub struct BroadcastManager {
channels: RwLock<HashMap<String, BroadcastChannel>>,
config: BroadcastConfig,
messages_published: AtomicU64,
}
impl BroadcastManager {
#[must_use]
pub fn new(config: BroadcastConfig) -> Self {
Self {
channels: RwLock::new(HashMap::new()),
config,
messages_published: AtomicU64::new(0),
}
}
pub async fn publish(
&self,
channel: &str,
event: String,
payload: serde_json::Value,
) -> Result<usize, BroadcastError> {
let payload_str = serde_json::to_string(&payload)
.map_err(|e| BroadcastError::InvalidPayload(e.to_string()))?;
if payload_str.len() > self.config.max_message_bytes {
return Err(BroadcastError::PayloadTooLarge {
size: payload_str.len(),
max: self.config.max_message_bytes,
});
}
let message = BroadcastMessage {
channel: channel.to_string(),
event,
payload,
};
{
let channels = self.channels.read().await;
if let Some(ch) = channels.get(channel) {
let receivers = ch.sender.send(message).unwrap_or(0);
self.messages_published.fetch_add(1, Ordering::Relaxed);
debug!(channel, receivers, "broadcast message sent (existing channel)");
return Ok(receivers);
}
}
let mut channels = self.channels.write().await;
if let Some(ch) = channels.get(channel) {
let receivers = ch.sender.send(message).unwrap_or(0);
self.messages_published.fetch_add(1, Ordering::Relaxed);
return Ok(receivers);
}
if channels.len() >= self.config.max_channels {
return Err(BroadcastError::TooManyChannels {
max: self.config.max_channels,
});
}
let (sender, _) = broadcast::channel(self.config.channel_capacity);
let receivers = sender.send(message).unwrap_or(0);
channels.insert(
channel.to_string(),
BroadcastChannel {
sender,
created_at: Instant::now(),
},
);
self.messages_published.fetch_add(1, Ordering::Relaxed);
debug!(channel, "broadcast channel created");
Ok(receivers)
}
pub async fn subscribe(
&self,
channel: &str,
) -> Result<broadcast::Receiver<BroadcastMessage>, BroadcastError> {
{
let channels = self.channels.read().await;
if let Some(ch) = channels.get(channel) {
return Ok(ch.sender.subscribe());
}
}
let mut channels = self.channels.write().await;
if let Some(ch) = channels.get(channel) {
return Ok(ch.sender.subscribe());
}
if channels.len() >= self.config.max_channels {
return Err(BroadcastError::TooManyChannels {
max: self.config.max_channels,
});
}
let (sender, receiver) = broadcast::channel(self.config.channel_capacity);
channels.insert(
channel.to_string(),
BroadcastChannel {
sender,
created_at: Instant::now(),
},
);
debug!(channel, "broadcast channel created for subscriber");
Ok(receiver)
}
pub async fn gc_empty_channels(&self) -> usize {
let mut channels = self.channels.write().await;
let before = channels.len();
channels.retain(|name, ch| {
let has_receivers = ch.sender.receiver_count() > 0;
if !has_receivers {
debug!(channel = %name, age_secs = ch.created_at.elapsed().as_secs(), "gc: removing empty broadcast channel");
}
has_receivers
});
before - channels.len()
}
pub async fn stats(&self) -> BroadcastStats {
let channels = self.channels.read().await;
let active_receivers: usize = channels.values().map(|ch| ch.sender.receiver_count()).sum();
BroadcastStats {
messages_published: self.messages_published.load(Ordering::Relaxed),
active_channels: channels.len(),
active_receivers,
}
}
pub async fn channel_count(&self) -> usize {
self.channels.read().await.len()
}
}
#[derive(Debug, thiserror::Error)]
pub enum BroadcastError {
#[error("payload too large: {size} bytes exceeds max {max}")]
PayloadTooLarge {
size: usize,
max: usize,
},
#[error("channel limit exceeded: max {max} channels")]
TooManyChannels {
max: usize,
},
#[error("invalid payload: {0}")]
InvalidPayload(String),
}
impl BroadcastError {
#[must_use]
pub const fn status_code(&self) -> u16 {
match self {
Self::PayloadTooLarge { .. } => 413,
Self::TooManyChannels { .. } => 503,
Self::InvalidPayload(_) => 400,
}
}
}