#![cfg(any(feature = "libpq", feature = "rustls-tls"))]
use pg_walstream::{
CancellationToken, ChangeEvent, EventType, LogicalReplicationStream, PgReplicationConnection,
ReplicationSlotOptions, ReplicationStreamConfig, StreamingMode, WalRouter,
};
use serde::Deserialize;
use std::sync::{Arc, Mutex};
use std::time::Duration;
fn replication_conn_string() -> String {
std::env::var("DATABASE_URL").unwrap_or_else(|_| {
"postgresql://postgres:postgres@localhost:5432/postgres?replication=database".to_string()
})
}
fn regular_conn_string() -> String {
if let Ok(s) = std::env::var("DATABASE_URL_REGULAR") {
return s;
}
let url = replication_conn_string();
match url.split_once('?') {
None => url,
Some((base, query)) => {
let kept: Vec<&str> = query
.split('&')
.filter(|p| *p != "replication=database")
.collect();
if kept.is_empty() {
base.to_string()
} else {
format!("{base}?{}", kept.join("&"))
}
}
}
}
fn drop_slot(slot_name: &str) {
if let Ok(mut conn) = PgReplicationConnection::connect(&replication_conn_string()) {
let _ = conn.exec(&format!(
"SELECT pg_drop_replication_slot('{slot_name}') \
WHERE EXISTS (SELECT 1 FROM pg_replication_slots WHERE slot_name = '{slot_name}')"
));
}
}
fn ergonomics_config(slot_name: &str, pub_name: &str) -> ReplicationStreamConfig {
ReplicationStreamConfig::builder(slot_name, pub_name)
.with_protocol_version(2)
.with_streaming_mode(StreamingMode::On)
.with_slot_options(ReplicationSlotOptions {
temporary: true,
..Default::default()
})
}
const TABLE: &str = "ergonomics_router";
const DDL: &str = "CREATE TABLE IF NOT EXISTS ergonomics_router (\
id BIGINT PRIMARY KEY, \
label TEXT NOT NULL\
)";
fn setup_table(regular: &mut PgReplicationConnection, pub_name: &str) {
let _ = regular.exec(DDL);
let _ = regular.exec("TRUNCATE ergonomics_router");
let _ = regular.exec("ALTER TABLE ergonomics_router REPLICA IDENTITY FULL");
let _ = regular.exec(&format!("DROP PUBLICATION IF EXISTS {pub_name}"));
let _ = regular.exec(&format!(
"CREATE PUBLICATION {pub_name} FOR TABLE ergonomics_router"
));
}
fn teardown(regular: &mut PgReplicationConnection, pub_name: &str, slot: &str) {
let _ = regular.exec(&format!("DROP PUBLICATION IF EXISTS {pub_name}"));
let _ = regular.exec("DROP TABLE IF EXISTS ergonomics_router");
drop_slot(slot);
}
#[derive(Debug, Deserialize, PartialEq)]
struct Row {
id: i64,
}
fn cancel_after(token: CancellationToken, secs: u64) {
tokio::spawn(async move {
tokio::time::sleep(Duration::from_secs(secs)).await;
token.cancel();
});
}
#[tokio::test]
#[ignore = "requires live PostgreSQL with wal_level=logical"]
async fn wal_router_dispatches_live_events() {
let slot = "it_ergo_router";
let pub_name = "ergo_router_pub";
drop_slot(slot);
let mut regular =
PgReplicationConnection::connect(®ular_conn_string()).expect("regular connection");
setup_table(&mut regular, pub_name);
let mut stream = LogicalReplicationStream::new(
&replication_conn_string(),
ergonomics_config(slot, pub_name),
)
.await
.expect("replication stream");
stream.start(None).await.expect("start");
let cancel_token = CancellationToken::new();
let mut event_stream = stream.into_stream(cancel_token.clone());
let inserted: Arc<Mutex<Vec<i64>>> = Arc::new(Mutex::new(Vec::new()));
let updated: Arc<Mutex<Vec<i64>>> = Arc::new(Mutex::new(Vec::new()));
let deleted: Arc<Mutex<Vec<i64>>> = Arc::new(Mutex::new(Vec::new()));
let mut router = WalRouter::new();
{
let ins = inserted.clone();
router.on_insert::<Row, _, _>(TABLE, move |r| {
let ins = ins.clone();
async move {
ins.lock().unwrap().push(r.id);
Ok(())
}
});
}
{
let upd = updated.clone();
router.on_update::<Row, _, _>(TABLE, move |_old, new| {
let upd = upd.clone();
async move {
upd.lock().unwrap().push(new.id);
Ok(())
}
});
}
{
let del = deleted.clone();
router.on_delete::<Row, _, _>(TABLE, move |r| {
let del = del.clone();
async move {
del.lock().unwrap().push(r.id);
Ok(())
}
});
}
regular
.exec("INSERT INTO ergonomics_router (id, label) VALUES (1, 'a'), (2, 'b')")
.expect("INSERT");
regular
.exec("UPDATE ergonomics_router SET label = 'b2' WHERE id = 2")
.expect("UPDATE");
regular
.exec("DELETE FROM ergonomics_router WHERE id = 1")
.expect("DELETE");
let counted = inserted.clone();
let cancel_clone = cancel_token.clone();
tokio::spawn(async move {
for _ in 0..200 {
if counted.lock().unwrap().len() >= 2 {
tokio::time::sleep(Duration::from_millis(500)).await;
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
cancel_clone.cancel();
});
cancel_after(cancel_token.clone(), 20);
router.run(&mut event_stream).await.expect("router.run");
assert_eq!(
*inserted.lock().unwrap(),
vec![1, 2],
"router should record both inserted ids in order"
);
assert_eq!(
*updated.lock().unwrap(),
vec![2],
"router should record update"
);
assert_eq!(
*deleted.lock().unwrap(),
vec![1],
"router should record delete"
);
let (_flushed, applied) = event_stream.get_feedback_lsn();
assert!(
applied > 0,
"applied LSN should have advanced, got {applied}"
);
teardown(&mut regular, pub_name, slot);
}
#[cfg(feature = "derive")]
mod derive_layer {
use super::*;
use pg_walstream::{WalRouter, WalTable};
use std::sync::{Arc, Mutex};
#[derive(serde::Deserialize, WalTable)]
#[wal(table = "ergonomics_router")]
struct DerivedRow {
id: i64,
}
#[tokio::test]
#[ignore = "requires live PostgreSQL with wal_level=logical"]
async fn on_insert_of_dispatches_live_events() {
let slot = "it_ergo_derive";
let pub_name = "ergo_derive_pub";
drop_slot(slot);
let mut regular =
PgReplicationConnection::connect(®ular_conn_string()).expect("regular connection");
setup_table(&mut regular, pub_name);
let mut stream = LogicalReplicationStream::new(
&replication_conn_string(),
ergonomics_config(slot, pub_name),
)
.await
.expect("replication stream");
stream.start(None).await.expect("start");
let cancel_token = CancellationToken::new();
let mut event_stream = stream.into_stream(cancel_token.clone());
let seen = Arc::new(Mutex::new(Vec::<i64>::new()));
let mut router = WalRouter::new();
{
let s = seen.clone();
router.on_insert_of::<DerivedRow, _>(move |row| {
let s = s.clone();
async move {
s.lock().unwrap().push(row.id);
Ok(())
}
});
}
assert_eq!(<DerivedRow as WalTable>::TABLE, "ergonomics_router");
regular
.exec("INSERT INTO ergonomics_router (id, label) VALUES (1, 'a'), (2, 'b')")
.expect("INSERT");
let counted = seen.clone();
let cancel_clone = cancel_token.clone();
tokio::spawn(async move {
for _ in 0..200 {
if counted.lock().unwrap().len() >= 2 {
break;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
cancel_clone.cancel();
});
cancel_after(cancel_token.clone(), 20);
router.run(&mut event_stream).await.expect("router.run");
assert_eq!(
*seen.lock().unwrap(),
vec![1, 2],
"on_insert_of-driven handler should record both inserted ids in order"
);
let (_flushed, applied) = event_stream.get_feedback_lsn();
assert!(
applied > 0,
"applied LSN should have advanced, got {applied}"
);
teardown(&mut regular, pub_name, slot);
}
}
#[tokio::test]
#[ignore = "requires live PostgreSQL with wal_level=logical"]
async fn for_each_event_auto_acks_live() {
let slot = "it_ergo_for_each";
let pub_name = "ergo_for_each_pub";
drop_slot(slot);
let mut regular =
PgReplicationConnection::connect(®ular_conn_string()).expect("regular connection");
setup_table(&mut regular, pub_name);
let mut stream = LogicalReplicationStream::new(
&replication_conn_string(),
ergonomics_config(slot, pub_name),
)
.await
.expect("replication stream");
stream.start(None).await.expect("start");
let cancel_token = CancellationToken::new();
let mut event_stream = stream.into_stream(cancel_token.clone());
regular
.exec("INSERT INTO ergonomics_router (id, label) VALUES (10, 'x'), (11, 'y'), (12, 'z')")
.expect("INSERT");
let seen = Arc::new(Mutex::new(0usize));
let seen_ids = Arc::new(Mutex::new(Vec::<i64>::new()));
let cancel_clone = cancel_token.clone();
cancel_after(cancel_token.clone(), 20);
let m = 3usize;
{
let seen = seen.clone();
let seen_ids = seen_ids.clone();
event_stream
.for_each_event(move |ev: ChangeEvent| {
let seen = seen.clone();
let seen_ids = seen_ids.clone();
let cancel = cancel_clone.clone();
async move {
if let EventType::Insert { .. } = ev.event_type {
if let Ok(row) = ev.deserialize_insert::<Row>() {
seen_ids.lock().unwrap().push(row.id);
}
let mut n = seen.lock().unwrap();
*n += 1;
if *n >= m {
cancel.cancel();
}
}
Ok(())
}
})
.await
.expect("for_each_event");
}
assert!(
*seen.lock().unwrap() >= m,
"should have seen at least {m} insert events, saw {}",
*seen.lock().unwrap()
);
assert_eq!(*seen_ids.lock().unwrap(), vec![10, 11, 12]);
let (_flushed, applied) = event_stream.get_feedback_lsn();
assert!(
applied > 0,
"applied LSN should have advanced, got {applied}"
);
teardown(&mut regular, pub_name, slot);
}