use crate::error::Result;
use crate::wire::ExtensionOids;
use crate::{execute_typed_on_raw, query_typed_named_on_raw, query_typed_on_raw};
use pylon_value::DecodedValue;
use tokio_postgres::AsyncMessage;
pub use tokio_postgres::Notification;
pub struct PgListener {
client: tokio_postgres::Client,
types: ExtensionOids,
}
impl PgListener {
pub async fn connect<F>(dsn: &str, on_notification: F) -> Result<Self>
where
F: Fn(Notification) + Send + 'static,
{
let pg_config = crate::session_config(dsn)?;
let (client, mut connection) = pg_config.connect(tokio_postgres::NoTls).await?;
tokio::spawn(async move {
loop {
match std::future::poll_fn(|cx| connection.poll_message(cx)).await {
Some(Ok(AsyncMessage::Notification(n))) => on_notification(n),
Some(Ok(_)) => continue,
Some(Err(_)) | None => break,
}
}
});
let types = crate::discover_types(&client).await?;
Ok(Self { client, types })
}
pub fn types(&self) -> &ExtensionOids {
&self.types
}
pub async fn listen(&self, channel: &str) -> Result<()> {
self.client
.batch_execute(&format!("LISTEN {}", quote_ident(channel)))
.await?;
Ok(())
}
pub async fn unlisten(&self, channel: &str) -> Result<()> {
self.client
.batch_execute(&format!("UNLISTEN {}", quote_ident(channel)))
.await?;
Ok(())
}
pub async fn query_typed(
&self,
sql: &str,
params: &[DecodedValue],
ext: &ExtensionOids,
) -> Result<Vec<DecodedValue>> {
query_typed_on_raw(&self.client, sql, params, ext).await
}
pub async fn query_typed_named(
&self,
sql: &str,
params: &[DecodedValue],
ext: &ExtensionOids,
) -> Result<Vec<DecodedValue>> {
query_typed_named_on_raw(&self.client, sql, params, ext).await
}
pub async fn batch_execute(&self, sql: &str) -> Result<()> {
self.client.batch_execute(sql).await?;
Ok(())
}
pub async fn execute_typed(&self, sql: &str, params: &[DecodedValue]) -> Result<u64> {
execute_typed_on_raw(&self.client, sql, params).await
}
}
pub(crate) fn quote_ident(ident: &str) -> String {
format!("\"{}\"", ident.replace('"', "\"\""))
}
#[cfg(test)]
mod tests {
use super::*;
fn test_dsn() -> String {
std::env::var("PYLON_PGCON_TEST_DSN").expect("PYLON_PGCON_TEST_DSN must be set to run live-Postgres tests")
}
#[tokio::test]
#[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
async fn receives_a_notification_on_a_listened_channel() {
use std::sync::{Arc, Mutex};
let received: Arc<Mutex<Vec<(String, String)>>> = Arc::new(Mutex::new(Vec::new()));
let received_clone = received.clone();
let listener = PgListener::connect(&test_dsn(), move |n| {
received_clone
.lock()
.unwrap()
.push((n.channel().to_string(), n.payload().to_string()));
})
.await
.unwrap();
listener.listen("pgcon_listener_test").await.unwrap();
let notifier = crate::PgPool::connect(&test_dsn(), 1).await.unwrap();
notifier.query_raw("NOTIFY pgcon_listener_test, 'hello'").await.unwrap();
for _ in 0..50 {
if !received.lock().unwrap().is_empty() {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
}
let got = received.lock().unwrap().clone();
assert_eq!(got, vec![("pgcon_listener_test".to_string(), "hello".to_string())]);
}
#[tokio::test]
#[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
async fn does_not_receive_notifications_on_channels_never_listened_to() {
use std::sync::{Arc, Mutex};
let received: Arc<Mutex<Vec<String>>> = Arc::new(Mutex::new(Vec::new()));
let received_clone = received.clone();
let _listener = PgListener::connect(&test_dsn(), move |n| {
received_clone.lock().unwrap().push(n.payload().to_string());
})
.await
.unwrap();
let notifier = crate::PgPool::connect(&test_dsn(), 1).await.unwrap();
notifier
.query_raw("NOTIFY pgcon_listener_unheard_test, 'should not arrive'")
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert!(received.lock().unwrap().is_empty());
}
#[tokio::test]
#[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
async fn unlisten_stops_further_notifications() {
use std::sync::{Arc, Mutex};
let count: Arc<Mutex<usize>> = Arc::new(Mutex::new(0));
let count_clone = count.clone();
let listener = PgListener::connect(&test_dsn(), move |_n| {
*count_clone.lock().unwrap() += 1;
})
.await
.unwrap();
listener.listen("pgcon_listener_unlisten_test").await.unwrap();
let notifier = crate::PgPool::connect(&test_dsn(), 1).await.unwrap();
notifier
.query_raw("NOTIFY pgcon_listener_unlisten_test, 'first'")
.await
.unwrap();
for _ in 0..50 {
if *count.lock().unwrap() >= 1 {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
}
assert_eq!(*count.lock().unwrap(), 1);
listener.unlisten("pgcon_listener_unlisten_test").await.unwrap();
notifier
.query_raw("NOTIFY pgcon_listener_unlisten_test, 'second'")
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert_eq!(
*count.lock().unwrap(),
1,
"unlisten should have stopped further notifications"
);
}
#[tokio::test]
#[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
async fn queries_run_on_the_same_connection_as_listen() {
let listener = PgListener::connect(&test_dsn(), |_n| {}).await.unwrap();
listener
.execute_typed("CREATE TEMP TABLE pgcon_listener_query_test (id int8)", &[])
.await
.unwrap();
listener
.execute_typed(
"INSERT INTO pgcon_listener_query_test (id) VALUES ($1::int8)",
&[DecodedValue::I64(7)],
)
.await
.unwrap();
let rows = listener
.query_typed(
"SELECT (id) AS result FROM pgcon_listener_query_test",
&[],
&ExtensionOids::default(),
)
.await
.unwrap();
assert_eq!(rows, vec![DecodedValue::I64(7)]);
}
#[tokio::test]
#[ignore = "requires a live Postgres via PYLON_PGCON_TEST_DSN"]
async fn query_typed_named_decodes_every_column_by_name() {
let listener = PgListener::connect(&test_dsn(), |_n| {}).await.unwrap();
let rows = listener
.query_typed_named(
"SELECT $1::int8 AS id, $2::text AS type_name, $3::text AS index_name",
&[
DecodedValue::I64(42),
DecodedValue::Str("default::Product".to_string()),
DecodedValue::Null,
],
&ExtensionOids::default(),
)
.await
.unwrap();
assert_eq!(
rows,
vec![DecodedValue::Object(vec![
("id".to_string(), DecodedValue::I64(42)),
(
"type_name".to_string(),
DecodedValue::Str("default::Product".to_string())
),
("index_name".to_string(), DecodedValue::Null),
])]
);
}
}