systemprompt_runtime/reporting/
snapshot_wakeup.rs1use 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#[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}