Skip to main content

fraiseql_server/subscriptions/
broadcast.rs

1//! Ephemeral broadcast channels for realtime pub/sub.
2//!
3//! Provides in-memory named channels that clients can publish to via REST
4//! (`POST /realtime/v1/broadcast`) and subscribe to via `WebSocket`.
5//! No database persistence — messages are lost on server restart.
6
7use std::{
8    collections::HashMap,
9    sync::atomic::{AtomicU64, Ordering},
10};
11
12use tokio::{
13    sync::{RwLock, broadcast},
14    time::Instant,
15};
16use tracing::debug;
17
18/// Configuration for the broadcast subsystem.
19#[derive(Debug, Clone)]
20pub struct BroadcastConfig {
21    /// Per-channel buffer capacity (number of messages retained for slow subscribers).
22    pub channel_capacity: usize,
23
24    /// Maximum number of named channels that can exist simultaneously.
25    pub max_channels: usize,
26
27    /// Maximum message payload size in bytes.
28    pub max_message_bytes: usize,
29}
30
31impl BroadcastConfig {
32    /// Create config with production defaults.
33    #[must_use]
34    pub const fn new() -> Self {
35        Self {
36            channel_capacity:  128,
37            max_channels:      1_000,
38            max_message_bytes: 65_536,
39        }
40    }
41}
42
43impl Default for BroadcastConfig {
44    fn default() -> Self {
45        Self::new()
46    }
47}
48
49/// Statistics for the broadcast subsystem.
50#[derive(Debug, Clone)]
51pub struct BroadcastStats {
52    /// Total messages published across all channels.
53    pub messages_published: u64,
54
55    /// Number of active channels.
56    pub active_channels: usize,
57
58    /// Total active receivers across all channels.
59    pub active_receivers: usize,
60}
61
62/// A single named broadcast channel.
63#[derive(Debug)]
64struct BroadcastChannel {
65    sender:     broadcast::Sender<BroadcastMessage>,
66    created_at: Instant,
67}
68
69/// A message published to a broadcast channel.
70#[derive(Debug, Clone)]
71pub struct BroadcastMessage {
72    /// The channel this message was published to.
73    pub channel: String,
74
75    /// The event name (e.g., `cursor_move`, `typing`).
76    pub event: String,
77
78    /// Arbitrary JSON payload.
79    pub payload: serde_json::Value,
80}
81
82/// Manages named broadcast channels.
83///
84/// Thread-safe via interior mutability (`RwLock` for channel map, atomics for counters).
85#[derive(Debug)]
86pub struct BroadcastManager {
87    channels:           RwLock<HashMap<String, BroadcastChannel>>,
88    config:             BroadcastConfig,
89    messages_published: AtomicU64,
90}
91
92impl BroadcastManager {
93    /// Create a new broadcast manager.
94    #[must_use]
95    pub fn new(config: BroadcastConfig) -> Self {
96        Self {
97            channels: RwLock::new(HashMap::new()),
98            config,
99            messages_published: AtomicU64::new(0),
100        }
101    }
102
103    /// Publish a message to a named channel.
104    ///
105    /// Creates the channel if it doesn't exist. Returns the number of receivers
106    /// that were notified (0 if nobody is listening).
107    ///
108    /// # Errors
109    ///
110    /// Returns error if the channel limit is exceeded or the payload is too large.
111    pub async fn publish(
112        &self,
113        channel: &str,
114        event: String,
115        payload: serde_json::Value,
116    ) -> Result<usize, BroadcastError> {
117        // Validate payload size
118        let payload_str = serde_json::to_string(&payload)
119            .map_err(|e| BroadcastError::InvalidPayload(e.to_string()))?;
120        if payload_str.len() > self.config.max_message_bytes {
121            return Err(BroadcastError::PayloadTooLarge {
122                size: payload_str.len(),
123                max:  self.config.max_message_bytes,
124            });
125        }
126
127        let message = BroadcastMessage {
128            channel: channel.to_string(),
129            event,
130            payload,
131        };
132
133        // Try to send on existing channel first (read lock — fast path)
134        {
135            let channels = self.channels.read().await;
136            if let Some(ch) = channels.get(channel) {
137                let receivers = ch.sender.send(message).unwrap_or(0);
138                self.messages_published.fetch_add(1, Ordering::Relaxed);
139                debug!(channel, receivers, "broadcast message sent (existing channel)");
140                return Ok(receivers);
141            }
142        }
143
144        // Channel doesn't exist — create it (write lock)
145        let mut channels = self.channels.write().await;
146
147        // Double-check after acquiring write lock
148        if let Some(ch) = channels.get(channel) {
149            let receivers = ch.sender.send(message).unwrap_or(0);
150            self.messages_published.fetch_add(1, Ordering::Relaxed);
151            return Ok(receivers);
152        }
153
154        // Check channel limit
155        if channels.len() >= self.config.max_channels {
156            return Err(BroadcastError::TooManyChannels {
157                max: self.config.max_channels,
158            });
159        }
160
161        let (sender, _) = broadcast::channel(self.config.channel_capacity);
162        let receivers = sender.send(message).unwrap_or(0);
163        channels.insert(
164            channel.to_string(),
165            BroadcastChannel {
166                sender,
167                created_at: Instant::now(),
168            },
169        );
170        self.messages_published.fetch_add(1, Ordering::Relaxed);
171        debug!(channel, "broadcast channel created");
172
173        Ok(receivers)
174    }
175
176    /// Subscribe to a named channel. Creates the channel if it doesn't exist.
177    ///
178    /// # Errors
179    ///
180    /// Returns error if the channel limit is exceeded.
181    pub async fn subscribe(
182        &self,
183        channel: &str,
184    ) -> Result<broadcast::Receiver<BroadcastMessage>, BroadcastError> {
185        // Fast path: read lock
186        {
187            let channels = self.channels.read().await;
188            if let Some(ch) = channels.get(channel) {
189                return Ok(ch.sender.subscribe());
190            }
191        }
192
193        // Create channel
194        let mut channels = self.channels.write().await;
195
196        // Double-check
197        if let Some(ch) = channels.get(channel) {
198            return Ok(ch.sender.subscribe());
199        }
200
201        if channels.len() >= self.config.max_channels {
202            return Err(BroadcastError::TooManyChannels {
203                max: self.config.max_channels,
204            });
205        }
206
207        let (sender, receiver) = broadcast::channel(self.config.channel_capacity);
208        channels.insert(
209            channel.to_string(),
210            BroadcastChannel {
211                sender,
212                created_at: Instant::now(),
213            },
214        );
215        debug!(channel, "broadcast channel created for subscriber");
216
217        Ok(receiver)
218    }
219
220    /// Remove channels with no active subscribers to prevent memory leaks.
221    pub async fn gc_empty_channels(&self) -> usize {
222        let mut channels = self.channels.write().await;
223        let before = channels.len();
224        channels.retain(|name, ch| {
225            let has_receivers = ch.sender.receiver_count() > 0;
226            if !has_receivers {
227                debug!(channel = %name, age_secs = ch.created_at.elapsed().as_secs(), "gc: removing empty broadcast channel");
228            }
229            has_receivers
230        });
231        before - channels.len()
232    }
233
234    /// Get current broadcast statistics.
235    pub async fn stats(&self) -> BroadcastStats {
236        let channels = self.channels.read().await;
237        let active_receivers: usize = channels.values().map(|ch| ch.sender.receiver_count()).sum();
238        BroadcastStats {
239            messages_published: self.messages_published.load(Ordering::Relaxed),
240            active_channels: channels.len(),
241            active_receivers,
242        }
243    }
244
245    /// Get the number of active channels.
246    pub async fn channel_count(&self) -> usize {
247        self.channels.read().await.len()
248    }
249}
250
251/// Errors from broadcast operations.
252#[derive(Debug, thiserror::Error)]
253pub enum BroadcastError {
254    /// Payload exceeds maximum allowed size.
255    #[error("payload too large: {size} bytes exceeds max {max}")]
256    PayloadTooLarge {
257        /// Actual payload size.
258        size: usize,
259        /// Maximum allowed size.
260        max:  usize,
261    },
262
263    /// Too many named channels exist.
264    #[error("channel limit exceeded: max {max} channels")]
265    TooManyChannels {
266        /// Maximum allowed channels.
267        max: usize,
268    },
269
270    /// Invalid payload data.
271    #[error("invalid payload: {0}")]
272    InvalidPayload(String),
273}
274
275impl BroadcastError {
276    /// HTTP status code for this error.
277    #[must_use]
278    pub const fn status_code(&self) -> u16 {
279        match self {
280            Self::PayloadTooLarge { .. } => 413,
281            Self::TooManyChannels { .. } => 503,
282            Self::InvalidPayload(_) => 400,
283        }
284    }
285}