Skip to main content

systemprompt_runtime/reporting/
snapshot_wakeup.rs

1//! One `PostgreSQL` `LISTEN` relay per process wakes the bounded feedback
2//! snapshot streams; the context owns it and joins it on shutdown.
3//!
4//! Copyright (c) systemprompt.io — Business Source License 1.1.
5//! See <https://systemprompt.io> for licensing details.
6
7use std::time::Duration;
8
9use sqlx::PgPool;
10use sqlx::postgres::PgListener;
11use systemprompt_database::DbPool;
12use tokio::sync::{Mutex, watch};
13use tokio::task::JoinHandle;
14
15const CHANNEL: &str = "feedback_snapshots";
16const TICK: Duration = Duration::from_secs(5);
17const CONNECT_TIMEOUT: Duration = Duration::from_secs(2);
18
19/// Lazily started relay from the `feedback_snapshots` channel to every
20/// subscribed stream.
21///
22/// The relay is spawned on the first subscription, idles without a listener
23/// while nobody is subscribed, and runs until [`SnapshotWakeup::shutdown`]
24/// joins it.
25#[derive(Debug)]
26pub struct SnapshotWakeup {
27    hints: watch::Sender<u64>,
28    relay: Mutex<Option<JoinHandle<()>>>,
29}
30
31impl Default for SnapshotWakeup {
32    fn default() -> Self {
33        Self {
34            hints: watch::channel(0).0,
35            relay: Mutex::const_new(None),
36        }
37    }
38}
39
40impl SnapshotWakeup {
41    pub async fn subscribe(&self, db: &DbPool) -> watch::Receiver<u64> {
42        let receiver = self.hints.subscribe();
43        let mut relay = self.relay.lock().await;
44        if relay.as_ref().is_none_or(JoinHandle::is_finished) {
45            let pool = db.write_pool().as_ref().clone();
46            let hints = self.hints.clone();
47            *relay = Some(tokio::spawn(run(pool, hints)));
48        }
49        receiver
50    }
51
52    pub async fn shutdown(&self) {
53        let Some(handle) = self.relay.lock().await.take() else {
54            return;
55        };
56        handle.abort();
57        if let Err(error) = handle.await
58            && !error.is_cancelled()
59        {
60            tracing::warn!(error = %error, "Snapshot wakeup relay ended abnormally");
61        }
62    }
63}
64
65async fn run(pool: PgPool, hints: watch::Sender<u64>) {
66    let mut interval = tokio::time::interval(TICK);
67    interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
68    let mut listener = None;
69    let mut failing = false;
70    loop {
71        tokio::select! {
72            _tick = interval.tick() => {},
73            () = notification(&mut listener) => {
74                hints.send_modify(|generation| *generation = generation.wrapping_add(1));
75            },
76        }
77        if hints.receiver_count() == 0 {
78            listener = None;
79            continue;
80        }
81        if listener.is_some() {
82            continue;
83        }
84        match tokio::time::timeout(CONNECT_TIMEOUT, connect(&pool)).await {
85            Ok(Ok(connected)) => {
86                if failing {
87                    tracing::info!("Snapshot wakeup listener recovered");
88                    failing = false;
89                }
90                listener = Some(connected);
91            },
92            Ok(Err(error)) => {
93                if !failing {
94                    tracing::warn!(
95                        error = %error,
96                        "Snapshot wakeup listener unavailable; streams fall back to polling"
97                    );
98                    failing = true;
99                }
100            },
101            Err(_elapsed) => {
102                if !failing {
103                    tracing::warn!(
104                        "Snapshot wakeup listener connect timed out; streams fall back to polling"
105                    );
106                    failing = true;
107                }
108            },
109        }
110    }
111}
112
113async fn notification(listener: &mut Option<PgListener>) {
114    match listener {
115        Some(connection) => {
116            if let Err(error) = connection.recv().await {
117                tracing::warn!(error = %error, "Snapshot wakeup listener dropped; reconnecting");
118                *listener = None;
119            }
120        },
121        None => std::future::pending::<()>().await,
122    }
123}
124
125async fn connect(pool: &PgPool) -> Result<PgListener, sqlx::Error> {
126    let mut listener = PgListener::connect_with(pool).await?;
127    listener.listen(CHANNEL).await?;
128    Ok(listener)
129}