postrust_graphql/subscription/
broker.rs1use 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
16const DEFAULT_CHANNEL_CAPACITY: usize = 256;
18
19#[derive(Debug, Clone)]
21pub struct PgNotification {
22 pub channel: String,
24 pub payload: String,
26 pub process_id: u32,
28}
29
30pub struct NotifyBroker {
32 pool: PgPool,
34 channels: Arc<RwLock<HashMap<String, broadcast::Sender<PgNotification>>>>,
36 channel_capacity: usize,
38 running: Arc<RwLock<bool>>,
40}
41
42impl NotifyBroker {
43 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 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 pub async fn start(&self, listen_channels: Vec<String>) -> Result<(), BrokerError> {
68 {
70 let running = self.running.read().await;
71 if *running {
72 return Err(BrokerError::AlreadyRunning);
73 }
74 }
75
76 {
78 let mut running = self.running.write().await;
79 *running = true;
80 }
81
82 {
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 let mut listener = PgListener::connect_with(&self.pool)
95 .await
96 .map_err(BrokerError::Database)?;
97
98 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 let channels = Arc::clone(&self.channels);
109 let running = Arc::clone(&self.running);
110
111 tokio::spawn(async move {
113 loop {
114 {
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 let channels_read = channels.read().await;
139 if let Some(sender) = channels_read.get(&pg_notification.channel) {
140 let _ = sender.send(pg_notification);
142 }
143 }
144 Ok(None) => {
145 tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
147 }
148 Err(e) => {
149 error!("Error receiving notification: {:?}", e);
150 tokio::time::sleep(tokio::time::Duration::from_secs(1)).await;
152 }
153 }
154 }
155 });
156
157 Ok(())
158 }
159
160 pub async fn stop(&self) {
162 let mut running = self.running.write().await;
163 *running = false;
164 info!("Broker stop requested");
165 }
166
167 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 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 pub async fn subscribe_or_create(
194 &self,
195 channel: &str,
196 ) -> Pin<Box<dyn Stream<Item = PgNotification> + Send>> {
197 {
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 {
210 let mut channels = self.channels.write().await;
211 if !channels.contains_key(channel) {
213 let (tx, _) = broadcast::channel(self.channel_capacity);
214 channels.insert(channel.to_string(), tx);
215 }
216 }
217
218 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 pub async fn listen_channel(&self, channel: &str) -> Result<(), BrokerError> {
229 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 {
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 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 pub async fn is_running(&self) -> bool {
295 *self.running.read().await
296 }
297
298 pub async fn channel_count(&self) -> usize {
300 self.channels.read().await.len()
301 }
302}
303
304#[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
317pub fn table_channel_name(schema: &str, table: &str) -> String {
319 format!("postrust_{}_{}", schema, table)
320}
321
322pub 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
381pub 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}