mod common;
use common::pgwire_harness::TestServer;
use tokio_postgres::SimpleQueryMessage;
fn command_count(msgs: &[SimpleQueryMessage]) -> Option<u64> {
msgs.iter().find_map(|m| match m {
SimpleQueryMessage::CommandComplete(n) => Some(*n),
_ => None,
})
}
async fn scan_ints(server: &TestServer, sql: &str) -> Vec<i64> {
let mut v: Vec<i64> = server
.query_text(sql)
.await
.unwrap()
.into_iter()
.map(|s| s.parse().unwrap())
.collect();
v.sort_unstable();
v
}
async fn setup(server: &TestServer, coll: &str, engine: &str) {
server
.exec(&format!(
"CREATE COLLECTION {coll} \
(id STRING NOT NULL PRIMARY KEY, n INT) WITH (engine='{engine}')"
))
.await
.unwrap();
for (id, n) in [("a", 1), ("b", 1), ("c", 2), ("unrelated", 100)] {
server
.exec(&format!("INSERT INTO {coll} (id, n) VALUES ('{id}', {n})"))
.await
.unwrap();
}
}
async fn bulk_update_stages_and_is_visible(engine: &str, coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, engine).await;
server.exec("BEGIN").await.unwrap();
let msgs = server
.client
.simple_query(&format!("UPDATE {coll} SET n = 7 WHERE n = 1"))
.await
.expect("in-tx bulk update should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(2),
"{engine}: predicate UPDATE must report the real matched-row count, not OK"
);
let seen = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 7")).await;
assert_eq!(
seen,
vec![7, 7],
"{engine}: in-tx scan must observe the staged bulk update"
);
let old = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 1")).await;
assert!(
old.is_empty(),
"{engine}: rows updated away from n=1 must no longer match n=1, got {old:?}"
);
server.client.simple_query("COMMIT").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 7")).await;
assert_eq!(
after,
vec![7, 7],
"{engine}: committed bulk update must persist"
);
}
async fn bulk_delete_stages_and_is_visible(engine: &str, coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, engine).await;
server.exec("BEGIN").await.unwrap();
let msgs = server
.client
.simple_query(&format!("DELETE FROM {coll} WHERE n = 1"))
.await
.expect("in-tx bulk delete should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(2),
"{engine}: predicate DELETE must report the real matched-row count, not OK"
);
let remaining = scan_ints(&server, &format!("SELECT n FROM {coll}")).await;
assert_eq!(
remaining,
vec![2, 100],
"{engine}: in-tx scan must hide the staged bulk delete"
);
server.client.simple_query("COMMIT").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {coll}")).await;
assert_eq!(
after,
vec![2, 100],
"{engine}: committed bulk delete must persist"
);
}
async fn bulk_dml_rollback_discards_staged_changes(engine: &str, coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, engine).await;
server.exec("BEGIN").await.unwrap();
server
.exec(&format!("UPDATE {coll} SET n = 7 WHERE n = 1"))
.await
.unwrap();
server
.exec(&format!("DELETE FROM {coll} WHERE n = 2"))
.await
.unwrap();
let staged = scan_ints(&server, &format!("SELECT n FROM {coll}")).await;
assert_eq!(
staged,
vec![7, 7, 100],
"{engine}: staged changes must be visible pre-rollback"
);
server.client.simple_query("ROLLBACK").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {coll}")).await;
assert_eq!(
after,
vec![1, 1, 2, 100],
"{engine}: ROLLBACK must restore the original rows, got {after:?}"
);
}
async fn bulk_update_visible_via_index_lookup(engine: &str, coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, engine).await;
server
.exec(&format!("CREATE INDEX ON {coll}(n)"))
.await
.unwrap();
server.exec("BEGIN").await.unwrap();
server
.exec(&format!("UPDATE {coll} SET n = 42 WHERE n = 1"))
.await
.unwrap();
let via_scan = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 42")).await;
assert_eq!(
via_scan,
vec![42, 42],
"{engine}: table scan must see the staged bulk update"
);
let via_index = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 42")).await;
assert_eq!(
via_index,
vec![42, 42],
"{engine}: indexed equality lookup must see the staged bulk update"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
async fn bulk_update_visible_via_point_get_by_pk(engine: &str, coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, engine).await;
server.exec("BEGIN").await.unwrap();
server
.exec(&format!("UPDATE {coll} SET n = 7 WHERE n = 1"))
.await
.unwrap();
let a = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE id = 'a'")).await;
assert_eq!(
a,
vec![7],
"{engine}: PK point-get must see the staged bulk update"
);
server.client.simple_query("ROLLBACK").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE id = 'a'")).await;
assert_eq!(after, vec![1], "{engine}: ROLLBACK restores the base row");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_bulk_update_visible_via_point_get_by_pk() {
bulk_update_visible_via_point_get_by_pk("document_schemaless", "bu_sc_pk").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_update_visible_via_point_get_by_pk() {
bulk_update_visible_via_point_get_by_pk("document_strict", "bu_st_pk").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_bulk_update_stages_and_is_visible() {
bulk_update_stages_and_is_visible("document_schemaless", "bu_sc_upd").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_bulk_delete_stages_and_is_visible() {
bulk_delete_stages_and_is_visible("document_schemaless", "bu_sc_del").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_bulk_dml_rollback_discards_staged_changes() {
bulk_dml_rollback_discards_staged_changes("document_schemaless", "bu_sc_rb").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_bulk_update_visible_via_index_lookup() {
bulk_update_visible_via_index_lookup("document_schemaless", "bu_sc_idx").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_update_stages_and_is_visible() {
bulk_update_stages_and_is_visible("document_strict", "bu_st_upd").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_delete_stages_and_is_visible() {
bulk_delete_stages_and_is_visible("document_strict", "bu_st_del").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_dml_rollback_discards_staged_changes() {
bulk_dml_rollback_discards_staged_changes("document_strict", "bu_st_rb").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_update_visible_via_index_lookup() {
bulk_update_visible_via_index_lookup("document_strict", "bu_st_idx").await;
}
async fn strict_bulk_update_sees_row_staged_earlier_in_txn(coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, "document_strict").await;
server.exec("BEGIN").await.unwrap();
server
.exec(&format!("INSERT INTO {coll} (id, n) VALUES ('fresh', 1)"))
.await
.unwrap();
let msgs = server
.client
.simple_query(&format!("UPDATE {coll} SET n = 9 WHERE n = 1"))
.await
.expect("in-tx bulk update should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(3),
"strict: predicate UPDATE must include the row staged earlier \
in this txn by a point INSERT (base 'a','b' + staged 'fresh')"
);
let seen = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 9")).await;
assert_eq!(
seen,
vec![9, 9, 9],
"strict: in-tx scan must show the staged-earlier row updated too"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
async fn strict_bulk_delete_sees_row_staged_earlier_in_txn(coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, "document_strict").await;
server.exec("BEGIN").await.unwrap();
server
.exec(&format!("INSERT INTO {coll} (id, n) VALUES ('fresh', 1)"))
.await
.unwrap();
let msgs = server
.client
.simple_query(&format!("DELETE FROM {coll} WHERE n = 1"))
.await
.expect("in-tx bulk delete should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(3),
"strict: predicate DELETE must include the row staged earlier \
in this txn by a point INSERT (base 'a','b' + staged 'fresh')"
);
let remaining = scan_ints(&server, &format!("SELECT n FROM {coll}")).await;
assert_eq!(
remaining,
vec![2, 100],
"strict: in-tx scan must hide both base and staged-earlier deleted rows"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_update_sees_row_staged_earlier_in_txn_case() {
strict_bulk_update_sees_row_staged_earlier_in_txn("bu_st_ov_upd").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_bulk_delete_sees_row_staged_earlier_in_txn_case() {
strict_bulk_delete_sees_row_staged_earlier_in_txn("bu_st_ov_del").await;
}
async fn bulk_update_returning_in_txn_reports_count_not_ok(engine: &str, coll: &str) {
let server = TestServer::start().await;
setup(&server, coll, engine).await;
server.exec("BEGIN").await.unwrap();
let msgs = server
.client
.simple_query(&format!(
"UPDATE {coll} SET n = 7 WHERE n = 1 RETURNING id, n"
))
.await
.expect("in-tx RETURNING bulk update should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(2),
"{engine}: in-tx UPDATE ... RETURNING must report the real matched-row count, not OK"
);
let row_count = msgs
.iter()
.filter(|m| matches!(m, SimpleQueryMessage::Row(_)))
.count();
assert_eq!(
row_count, 0,
"{engine}: in-tx RETURNING rows are not projected back to the client"
);
let seen = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 7")).await;
assert_eq!(
seen,
vec![7, 7],
"{engine}: in-tx scan must observe the staged RETURNING update"
);
server.client.simple_query("COMMIT").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {coll} WHERE n = 7")).await;
assert_eq!(
after,
vec![7, 7],
"{engine}: committed RETURNING update must persist"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn bulk_update_returning_in_txn_reports_count_not_ok_case() {
bulk_update_returning_in_txn_reports_count_not_ok("document_schemaless", "bu_ret_ov_upd").await;
}