1use crate::error::Result;
33use crate::wire::ExtensionOids;
34use crate::{execute_typed_on_raw, query_typed_named_on_raw, query_typed_on_raw};
35use pylon_value::DecodedValue;
36use tokio_postgres::AsyncMessage;
37
38pub use tokio_postgres::Notification;
39
40pub struct PgListener {
46 client: tokio_postgres::Client,
47 types: std::sync::RwLock<std::sync::Arc<ExtensionOids>>,
50}
51
52impl PgListener {
53 pub async fn connect<F>(dsn: &str, on_notification: F) -> Result<Self>
61 where
62 F: Fn(Notification) + Send + 'static,
63 {
64 let pg_config = crate::session_config(dsn)?;
65 let (client, mut connection) = pg_config.connect(tokio_postgres::NoTls).await?;
66 tokio::spawn(async move {
67 loop {
68 match std::future::poll_fn(|cx| connection.poll_message(cx)).await {
69 Some(Ok(AsyncMessage::Notification(n))) => on_notification(n),
70 Some(Ok(_)) => continue,
71 Some(Err(_)) | None => break,
72 }
73 }
74 });
75 let types = crate::discover_types(&client).await?;
76 Ok(Self {
77 client,
78 types: std::sync::RwLock::new(std::sync::Arc::new(types)),
79 })
80 }
81
82 pub fn types(&self) -> std::sync::Arc<ExtensionOids> {
85 self.types.read().unwrap().clone()
86 }
87
88 pub async fn refresh_types(&self) -> Result<()> {
92 let fresh = crate::discover_types(&self.client).await?;
93 *self.types.write().unwrap() = std::sync::Arc::new(fresh);
94 Ok(())
95 }
96
97 async fn heal_types(&self) -> Result<std::sync::Arc<ExtensionOids>> {
99 let fresh = std::sync::Arc::new(crate::discover_types(&self.client).await?);
100 *self.types.write().unwrap() = fresh.clone();
101 Ok(fresh)
102 }
103
104 pub async fn listen(&self, channel: &str) -> Result<()> {
105 self.client
106 .batch_execute(&format!("LISTEN {}", quote_ident(channel)))
107 .await?;
108 Ok(())
109 }
110
111 pub async fn unlisten(&self, channel: &str) -> Result<()> {
112 self.client
113 .batch_execute(&format!("UNLISTEN {}", quote_ident(channel)))
114 .await?;
115 Ok(())
116 }
117
118 pub async fn query_typed(
119 &self,
120 sql: &str,
121 params: &[DecodedValue],
122 ext: &ExtensionOids,
123 ) -> Result<Vec<DecodedValue>> {
124 match query_typed_on_raw(&self.client, sql, params, ext).await {
125 Err(e) if crate::is_unknown_oid(&e) => {
126 let healed = self.heal_types().await?;
127 query_typed_on_raw(&self.client, sql, params, &healed).await
128 }
129 other => other,
130 }
131 }
132
133 pub async fn query_typed_named(
138 &self,
139 sql: &str,
140 params: &[DecodedValue],
141 ext: &ExtensionOids,
142 ) -> Result<Vec<DecodedValue>> {
143 match query_typed_named_on_raw(&self.client, sql, params, ext).await {
144 Err(e) if crate::is_unknown_oid(&e) => {
145 let healed = self.heal_types().await?;
146 query_typed_named_on_raw(&self.client, sql, params, &healed).await
147 }
148 other => other,
149 }
150 }
151
152 pub async fn batch_execute(&self, sql: &str) -> Result<()> {
156 self.client.batch_execute(sql).await?;
157 Ok(())
158 }
159
160 pub async fn execute_typed(&self, sql: &str, params: &[DecodedValue]) -> Result<u64> {
161 execute_typed_on_raw(&self.client, sql, params).await
162 }
163}
164
165pub(crate) fn quote_ident(ident: &str) -> String {
171 format!("\"{}\"", ident.replace('"', "\"\""))
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177
178 fn test_dsn() -> String {
179 std::env::var("PYLON_PGCON_TEST_DSN").expect("PYLON_PGCON_TEST_DSN must be set to run live-Postgres tests")
180 }
181
182 #[tokio::test]
183 #[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
184 async fn receives_a_notification_on_a_listened_channel() {
185 use std::sync::{Arc, Mutex};
186 let received: Arc<Mutex<Vec<(String, String)>>> = Arc::new(Mutex::new(Vec::new()));
187 let received_clone = received.clone();
188
189 let listener = PgListener::connect(&test_dsn(), move |n| {
190 received_clone
191 .lock()
192 .unwrap()
193 .push((n.channel().to_string(), n.payload().to_string()));
194 })
195 .await
196 .unwrap();
197 listener.listen("pgcon_listener_test").await.unwrap();
198
199 let notifier = crate::PgPool::connect(&test_dsn(), 1).await.unwrap();
202 notifier.query_raw("NOTIFY pgcon_listener_test, 'hello'").await.unwrap();
203
204 for _ in 0..50 {
208 if !received.lock().unwrap().is_empty() {
209 break;
210 }
211 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
212 }
213
214 let got = received.lock().unwrap().clone();
215 assert_eq!(got, vec![("pgcon_listener_test".to_string(), "hello".to_string())]);
216 }
217
218 #[tokio::test]
219 #[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
220 async fn does_not_receive_notifications_on_channels_never_listened_to() {
221 use std::sync::{Arc, Mutex};
222 let received: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
223 let received_clone = received.clone();
224
225 let _listener = PgListener::connect(&test_dsn(), move |n| {
226 received_clone.lock().unwrap().push(n.payload().to_string());
227 })
228 .await
229 .unwrap();
230 let notifier = crate::PgPool::connect(&test_dsn(), 1).await.unwrap();
233 notifier
234 .query_raw("NOTIFY pgcon_listener_unheard_test, 'should not arrive'")
235 .await
236 .unwrap();
237
238 tokio::time::sleep(std::time::Duration::from_millis(200)).await;
239 assert!(received.lock().unwrap().is_empty());
240 }
241
242 #[tokio::test]
243 #[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
244 async fn unlisten_stops_further_notifications() {
245 use std::sync::{Arc, Mutex};
246 let count: Arc<Mutex<usize>> = Arc::new(Mutex::new(0));
247 let count_clone = count.clone();
248
249 let listener = PgListener::connect(&test_dsn(), move |_n| {
250 *count_clone.lock().unwrap() += 1;
251 })
252 .await
253 .unwrap();
254 listener.listen("pgcon_listener_unlisten_test").await.unwrap();
255
256 let notifier = crate::PgPool::connect(&test_dsn(), 1).await.unwrap();
257 notifier
258 .query_raw("NOTIFY pgcon_listener_unlisten_test, 'first'")
259 .await
260 .unwrap();
261 for _ in 0..50 {
262 if *count.lock().unwrap() >= 1 {
263 break;
264 }
265 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
266 }
267 assert_eq!(*count.lock().unwrap(), 1);
268
269 listener.unlisten("pgcon_listener_unlisten_test").await.unwrap();
270 notifier
271 .query_raw("NOTIFY pgcon_listener_unlisten_test, 'second'")
272 .await
273 .unwrap();
274 tokio::time::sleep(std::time::Duration::from_millis(200)).await;
275 assert_eq!(
276 *count.lock().unwrap(),
277 1,
278 "unlisten should have stopped further notifications"
279 );
280 }
281
282 #[tokio::test]
283 #[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
284 async fn queries_run_on_the_same_connection_as_listen() {
285 let listener = PgListener::connect(&test_dsn(), |_n| {}).await.unwrap();
286 listener
287 .execute_typed("CREATE TEMP TABLE pgcon_listener_query_test (id int8)", &[])
288 .await
289 .unwrap();
290 listener
291 .execute_typed(
292 "INSERT INTO pgcon_listener_query_test (id) VALUES ($1::int8)",
293 &[DecodedValue::I64(7)],
294 )
295 .await
296 .unwrap();
297 let rows = listener
298 .query_typed(
299 "SELECT (id) AS result FROM pgcon_listener_query_test",
300 &[],
301 &ExtensionOids::default(),
302 )
303 .await
304 .unwrap();
305 assert_eq!(rows, vec![DecodedValue::I64(7)]);
306 }
307
308 #[tokio::test]
309 #[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
310 async fn query_typed_named_decodes_every_column_by_name() {
311 let listener = PgListener::connect(&test_dsn(), |_n| {}).await.unwrap();
315 let rows = listener
316 .query_typed_named(
317 "SELECT $1::int8 AS id, $2::text AS type_name, $3::text AS index_name",
318 &[
319 DecodedValue::I64(42),
320 DecodedValue::Str("default::Product".to_string()),
321 DecodedValue::Null,
322 ],
323 &ExtensionOids::default(),
324 )
325 .await
326 .unwrap();
327 assert_eq!(
328 rows,
329 vec![DecodedValue::Object(vec![
330 ("id".to_string(), DecodedValue::I64(42)),
331 (
332 "type_name".to_string(),
333 DecodedValue::Str("default::Product".to_string())
334 ),
335 ("index_name".to_string(), DecodedValue::Null),
336 ])]
337 );
338 }
339}