Skip to main content

postrust_graphql/subscription/
broker.rs

1//! PostgreSQL NOTIFY message broker for GraphQL subscriptions.
2//!
3//! This module provides a broker that listens to PostgreSQL NOTIFY events
4//! and broadcasts them to GraphQL subscription clients.
5
6use futures::stream::{Stream, StreamExt};
7use sqlx::postgres::PgListener;
8use sqlx::PgPool;
9use std::collections::HashMap;
10use std::pin::Pin;
11use std::sync::Arc;
12use tokio::sync::broadcast;
13use tokio::sync::RwLock;
14use tracing::{debug, error, info, warn};
15
16/// Default channel capacity for broadcast channels
17const DEFAULT_CHANNEL_CAPACITY: usize = 256;
18
19/// A notification from PostgreSQL
20#[derive(Debug, Clone)]
21pub struct PgNotification {
22    /// The channel name (table name or custom channel)
23    pub channel: String,
24    /// The payload (usually JSON)
25    pub payload: String,
26    /// Process ID that sent the notification
27    pub process_id: u32,
28}
29
30/// Message broker that distributes PostgreSQL NOTIFY events to subscribers.
31pub struct NotifyBroker {
32    /// Database connection pool
33    pool: PgPool,
34    /// Channel senders keyed by channel name
35    channels: Arc<RwLock<HashMap<String, broadcast::Sender<PgNotification>>>>,
36    /// Capacity for new broadcast channels
37    channel_capacity: usize,
38    /// Whether the broker is running
39    running: Arc<RwLock<bool>>,
40}
41
42impl NotifyBroker {
43    /// Create a new notification broker.
44    pub fn new(pool: PgPool) -> Self {
45        Self {
46            pool,
47            channels: Arc::new(RwLock::new(HashMap::new())),
48            channel_capacity: DEFAULT_CHANNEL_CAPACITY,
49            running: Arc::new(RwLock::new(false)),
50        }
51    }
52
53    /// Create a new notification broker with custom channel capacity.
54    pub fn with_capacity(pool: PgPool, capacity: usize) -> Self {
55        Self {
56            pool,
57            channels: Arc::new(RwLock::new(HashMap::new())),
58            channel_capacity: capacity,
59            running: Arc::new(RwLock::new(false)),
60        }
61    }
62
63    /// Start listening for notifications on the given channels.
64    ///
65    /// This spawns a background task that listens for PostgreSQL NOTIFY events
66    /// and broadcasts them to all subscribers.
67    pub async fn start(&self, listen_channels: Vec<String>) -> Result<(), BrokerError> {
68        // Check if already running
69        {
70            let running = self.running.read().await;
71            if *running {
72                return Err(BrokerError::AlreadyRunning);
73            }
74        }
75
76        // Mark as running
77        {
78            let mut running = self.running.write().await;
79            *running = true;
80        }
81
82        // Create channels for each listen channel
83        {
84            let mut channels = self.channels.write().await;
85            for channel_name in &listen_channels {
86                if !channels.contains_key(channel_name) {
87                    let (tx, _) = broadcast::channel(self.channel_capacity);
88                    channels.insert(channel_name.clone(), tx);
89                }
90            }
91        }
92
93        // Create listener
94        let mut listener = PgListener::connect_with(&self.pool)
95            .await
96            .map_err(BrokerError::Database)?;
97
98        // Subscribe to all channels
99        for channel in &listen_channels {
100            listener
101                .listen(channel)
102                .await
103                .map_err(BrokerError::Database)?;
104            info!("Listening on PostgreSQL channel: {}", channel);
105        }
106
107        // Clone for the spawned task
108        let channels = Arc::clone(&self.channels);
109        let running = Arc::clone(&self.running);
110
111        // Spawn listener task
112        tokio::spawn(async move {
113            loop {
114                // Check if we should stop
115                {
116                    let is_running = running.read().await;
117                    if !*is_running {
118                        info!("Broker stopped, exiting listener loop");
119                        break;
120                    }
121                }
122
123                match listener.try_recv().await {
124                    Ok(Some(notification)) => {
125                        let pg_notification = PgNotification {
126                            channel: notification.channel().to_string(),
127                            payload: notification.payload().to_string(),
128                            process_id: notification.process_id(),
129                        };
130
131                        debug!(
132                            "Received notification on channel '{}': {}",
133                            pg_notification.channel,
134                            &pg_notification.payload[..pg_notification.payload.len().min(100)]
135                        );
136
137                        // Broadcast to subscribers
138                        let channels_read = channels.read().await;
139                        if let Some(sender) = channels_read.get(&pg_notification.channel) {
140                            // Ignore send errors - means no active receivers
141                            let _ = sender.send(pg_notification);
142                        }
143                    }
144                    Ok(None) => {
145                        // No notification available, continue
146                        tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
147                    }
148                    Err(e) => {
149                        error!("Error receiving notification: {:?}", e);
150                        // Try to reconnect after a delay
151                        tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
152                    }
153                }
154            }
155        });
156
157        Ok(())
158    }
159
160    /// Stop the broker.
161    pub async fn stop(&self) {
162        let mut running = self.running.write().await;
163        *running = false;
164        info!("Broker stop requested");
165    }
166
167    /// Subscribe to notifications for a specific channel.
168    ///
169    /// Returns a stream of notifications for the given channel.
170    pub async fn subscribe(
171        &self,
172        channel: &str,
173    ) -> Result<Pin<Box<dyn Stream<Item = PgNotification> + Send>>, BrokerError> {
174        let channels = self.channels.read().await;
175
176        let sender = channels
177            .get(channel)
178            .ok_or_else(|| BrokerError::ChannelNotFound(channel.to_string()))?;
179
180        let receiver = sender.subscribe();
181
182        // Convert broadcast receiver to stream
183        let stream = tokio_stream::wrappers::BroadcastStream::new(receiver)
184            .filter_map(|result| futures::future::ready(result.ok()));
185
186        Ok(Box::pin(stream))
187    }
188
189    /// Subscribe to a channel, creating it if it doesn't exist.
190    ///
191    /// Note: This only creates a broadcast channel. You must also call
192    /// `listen_channel` to start receiving PostgreSQL notifications.
193    pub async fn subscribe_or_create(
194        &self,
195        channel: &str,
196    ) -> Pin<Box<dyn Stream<Item = PgNotification> + Send>> {
197        // First try to get existing channel
198        {
199            let channels = self.channels.read().await;
200            if let Some(sender) = channels.get(channel) {
201                let receiver = sender.subscribe();
202                let stream = tokio_stream::wrappers::BroadcastStream::new(receiver)
203                    .filter_map(|result| futures::future::ready(result.ok()));
204                return Box::pin(stream);
205            }
206        }
207
208        // Create new channel
209        {
210            let mut channels = self.channels.write().await;
211            // Double-check after acquiring write lock
212            if !channels.contains_key(channel) {
213                let (tx, _) = broadcast::channel(self.channel_capacity);
214                channels.insert(channel.to_string(), tx);
215            }
216        }
217
218        // Now subscribe
219        let channels = self.channels.read().await;
220        let sender = channels.get(channel).expect("just created");
221        let receiver = sender.subscribe();
222        let stream = tokio_stream::wrappers::BroadcastStream::new(receiver)
223            .filter_map(|result| futures::future::ready(result.ok()));
224        Box::pin(stream)
225    }
226
227    /// Add a new channel to listen on dynamically.
228    pub async fn listen_channel(&self, channel: &str) -> Result<(), BrokerError> {
229        // Create a new listener for this channel
230        let mut listener = PgListener::connect_with(&self.pool)
231            .await
232            .map_err(BrokerError::Database)?;
233
234        listener
235            .listen(channel)
236            .await
237            .map_err(BrokerError::Database)?;
238
239        // Ensure broadcast channel exists
240        {
241            let mut channels = self.channels.write().await;
242            if !channels.contains_key(channel) {
243                let (tx, _) = broadcast::channel(self.channel_capacity);
244                channels.insert(channel.to_string(), tx);
245            }
246        }
247
248        let channels = Arc::clone(&self.channels);
249        let running = Arc::clone(&self.running);
250        let channel_name = channel.to_string();
251
252        // Spawn a listener for this channel
253        tokio::spawn(async move {
254            info!("Started dynamic listener for channel: {}", channel_name);
255
256            loop {
257                {
258                    let is_running = running.read().await;
259                    if !*is_running {
260                        break;
261                    }
262                }
263
264                match listener.try_recv().await {
265                    Ok(Some(notification)) => {
266                        let pg_notification = PgNotification {
267                            channel: notification.channel().to_string(),
268                            payload: notification.payload().to_string(),
269                            process_id: notification.process_id(),
270                        };
271
272                        let channels_read = channels.read().await;
273                        if let Some(sender) = channels_read.get(&pg_notification.channel) {
274                            let _ = sender.send(pg_notification);
275                        }
276                    }
277                    Ok(None) => {
278                        tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
279                    }
280                    Err(e) => {
281                        warn!("Error on channel {}: {:?}", channel_name, e);
282                        tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
283                    }
284                }
285            }
286
287            info!("Stopped dynamic listener for channel: {}", channel_name);
288        });
289
290        Ok(())
291    }
292
293    /// Check if the broker is currently running.
294    pub async fn is_running(&self) -> bool {
295        *self.running.read().await
296    }
297
298    /// Get the number of active channels.
299    pub async fn channel_count(&self) -> usize {
300        self.channels.read().await.len()
301    }
302}
303
304/// Errors that can occur in the broker.
305#[derive(Debug, thiserror::Error)]
306pub enum BrokerError {
307    #[error("Database error: {0}")]
308    Database(#[from] sqlx::Error),
309
310    #[error("Channel not found: {0}")]
311    ChannelNotFound(String),
312
313    #[error("Broker is already running")]
314    AlreadyRunning,
315}
316
317/// Generate a channel name for table change notifications.
318pub fn table_channel_name(schema: &str, table: &str) -> String {
319    format!("postrust_{}_{}", schema, table)
320}
321
322/// Generate SQL to create a notification trigger for a table.
323pub fn create_notify_trigger_sql(schema: &str, table: &str) -> String {
324    let channel = table_channel_name(schema, table);
325    let trigger_name = format!("postrust_notify_{}_{}", schema, table);
326    let function_name = format!("postrust_notify_{}_{}_fn", schema, table);
327
328    format!(
329        r#"
330-- Create notification function
331CREATE OR REPLACE FUNCTION {schema}.{function_name}()
332RETURNS TRIGGER AS $$
333DECLARE
334    payload jsonb;
335BEGIN
336    IF TG_OP = 'DELETE' THEN
337        payload := jsonb_build_object(
338            'operation', 'DELETE',
339            'table', TG_TABLE_NAME,
340            'schema', TG_TABLE_SCHEMA,
341            'old', row_to_json(OLD)
342        );
343    ELSIF TG_OP = 'UPDATE' THEN
344        payload := jsonb_build_object(
345            'operation', 'UPDATE',
346            'table', TG_TABLE_NAME,
347            'schema', TG_TABLE_SCHEMA,
348            'old', row_to_json(OLD),
349            'new', row_to_json(NEW)
350        );
351    ELSIF TG_OP = 'INSERT' THEN
352        payload := jsonb_build_object(
353            'operation', 'INSERT',
354            'table', TG_TABLE_NAME,
355            'schema', TG_TABLE_SCHEMA,
356            'new', row_to_json(NEW)
357        );
358    END IF;
359
360    PERFORM pg_notify('{channel}', payload::text);
361
362    RETURN COALESCE(NEW, OLD);
363END;
364$$ LANGUAGE plpgsql;
365
366-- Create trigger
367DROP TRIGGER IF EXISTS {trigger_name} ON {schema}.{table};
368CREATE TRIGGER {trigger_name}
369    AFTER INSERT OR UPDATE OR DELETE ON {schema}.{table}
370    FOR EACH ROW
371    EXECUTE FUNCTION {schema}.{function_name}();
372"#,
373        schema = schema,
374        table = table,
375        channel = channel,
376        function_name = function_name,
377        trigger_name = trigger_name
378    )
379}
380
381/// Generate SQL to drop a notification trigger for a table.
382pub fn drop_notify_trigger_sql(schema: &str, table: &str) -> String {
383    let trigger_name = format!("postrust_notify_{}_{}", schema, table);
384    let function_name = format!("postrust_notify_{}_{}_fn", schema, table);
385
386    format!(
387        r#"
388DROP TRIGGER IF EXISTS {trigger_name} ON {schema}.{table};
389DROP FUNCTION IF EXISTS {schema}.{function_name}();
390"#,
391        schema = schema,
392        table = table,
393        trigger_name = trigger_name,
394        function_name = function_name
395    )
396}
397
398#[cfg(test)]
399mod tests {
400    use super::*;
401
402    #[test]
403    fn test_table_channel_name() {
404        assert_eq!(
405            table_channel_name("public", "users"),
406            "postrust_public_users"
407        );
408        assert_eq!(table_channel_name("api", "orders"), "postrust_api_orders");
409    }
410
411    #[test]
412    fn test_create_notify_trigger_sql() {
413        let sql = create_notify_trigger_sql("public", "users");
414        assert!(sql.contains("CREATE OR REPLACE FUNCTION"));
415        assert!(sql.contains("postrust_notify_public_users_fn"));
416        assert!(sql.contains("CREATE TRIGGER"));
417        assert!(sql.contains("pg_notify"));
418        assert!(sql.contains("postrust_public_users"));
419    }
420
421    #[test]
422    fn test_drop_notify_trigger_sql() {
423        let sql = drop_notify_trigger_sql("public", "users");
424        assert!(sql.contains("DROP TRIGGER IF EXISTS"));
425        assert!(sql.contains("DROP FUNCTION IF EXISTS"));
426    }
427}