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, src: &str, tgt: &str, engine: &str) {
for coll in [src, tgt] {
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 {src} (id, n) VALUES ('{id}', {n})"))
.await
.unwrap();
}
}
async fn insert_select_stages_and_is_visible(engine: &str, src: &str, tgt: &str) {
let server = TestServer::start().await;
setup(&server, src, tgt, engine).await;
server.exec("BEGIN").await.unwrap();
let msgs = server
.client
.simple_query(&format!(
"INSERT INTO {tgt} SELECT * FROM {src} WHERE n = 1"
))
.await
.expect("in-tx insert-select should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(2),
"{engine}: INSERT ... SELECT must report the real copied-row count, not OK"
);
let seen = scan_ints(&server, &format!("SELECT n FROM {tgt}")).await;
assert_eq!(
seen,
vec![1, 1],
"{engine}: in-tx scan of target must observe the staged copy"
);
server.client.simple_query("COMMIT").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {tgt}")).await;
assert_eq!(
after,
vec![1, 1],
"{engine}: committed insert-select must persist"
);
}
async fn insert_select_rollback_discards_staged_copy(engine: &str, src: &str, tgt: &str) {
let server = TestServer::start().await;
setup(&server, src, tgt, engine).await;
server.exec("BEGIN").await.unwrap();
server
.exec(&format!(
"INSERT INTO {tgt} SELECT * FROM {src} WHERE n = 1"
))
.await
.unwrap();
let staged = scan_ints(&server, &format!("SELECT n FROM {tgt}")).await;
assert_eq!(
staged,
vec![1, 1],
"{engine}: staged copy must be visible pre-rollback"
);
server.client.simple_query("ROLLBACK").await.unwrap();
let after = scan_ints(&server, &format!("SELECT n FROM {tgt}")).await;
assert!(
after.is_empty(),
"{engine}: ROLLBACK must discard the staged copy, got {after:?}"
);
}
async fn insert_select_respects_limit(engine: &str, src: &str, tgt: &str) {
let server = TestServer::start().await;
setup(&server, src, tgt, engine).await;
server.exec("BEGIN").await.unwrap();
let msgs = server
.client
.simple_query(&format!(
"INSERT INTO {tgt} SELECT * FROM {src} WHERE n = 1 LIMIT 1"
))
.await
.expect("in-tx insert-select with LIMIT should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(1),
"{engine}: LIMIT must cap the staged copied-row count"
);
let seen = scan_ints(&server, &format!("SELECT n FROM {tgt}")).await;
assert_eq!(
seen,
vec![1],
"{engine}: in-tx scan of target must show only the LIMIT-ed copy"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
async fn insert_select_sees_source_rows_staged_earlier_in_txn(engine: &str, src: &str, tgt: &str) {
let server = TestServer::start().await;
setup(&server, src, tgt, engine).await;
server.exec("BEGIN").await.unwrap();
server
.exec(&format!("INSERT INTO {src} (id, n) VALUES ('fresh', 1)"))
.await
.unwrap();
let msgs = server
.client
.simple_query(&format!(
"INSERT INTO {tgt} SELECT * FROM {src} WHERE n = 1"
))
.await
.expect("in-tx insert-select should succeed at the statement");
assert_eq!(
command_count(&msgs),
Some(3),
"{engine}: copy must include the source row staged earlier in this txn"
);
let seen = scan_ints(&server, &format!("SELECT n FROM {tgt}")).await;
assert_eq!(
seen,
vec![1, 1, 1],
"{engine}: target must contain the copy of the in-txn-staged source row"
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_insert_select_stages_and_is_visible() {
insert_select_stages_and_is_visible("document_schemaless", "is_sc_src", "is_sc_tgt").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_insert_select_rollback_discards_staged_copy() {
insert_select_rollback_discards_staged_copy(
"document_schemaless",
"is_sc_rb_src",
"is_sc_rb_tgt",
)
.await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_insert_select_respects_limit() {
insert_select_respects_limit("document_schemaless", "is_sc_lim_src", "is_sc_lim_tgt").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn schemaless_insert_select_sees_source_rows_staged_earlier_in_txn() {
insert_select_sees_source_rows_staged_earlier_in_txn(
"document_schemaless",
"is_sc_ov_src",
"is_sc_ov_tgt",
)
.await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_insert_select_stages_and_is_visible() {
insert_select_stages_and_is_visible("document_strict", "is_st_src", "is_st_tgt").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_insert_select_rollback_discards_staged_copy() {
insert_select_rollback_discards_staged_copy("document_strict", "is_st_rb_src", "is_st_rb_tgt")
.await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_insert_select_respects_limit() {
insert_select_respects_limit("document_strict", "is_st_lim_src", "is_st_lim_tgt").await;
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn strict_insert_select_sees_source_rows_staged_earlier_in_txn() {
insert_select_sees_source_rows_staged_earlier_in_txn(
"document_strict",
"is_st_ov_src",
"is_st_ov_tgt",
)
.await;
}