mod common;
use common::pgwire_harness::TestServer;
use tokio_postgres::SimpleQueryMessage;
async fn read_n(server: &TestServer, sql: &str) -> Option<String> {
let msgs = server
.client
.simple_query(sql)
.await
.expect("in-tx read should succeed");
msgs.iter().find_map(|m| match m {
SimpleQueryMessage::Row(r) => r.get("n").map(str::to_string),
_ => None,
})
}
async fn setup_doc(server: &TestServer) {
server
.exec(
"CREATE COLLECTION t \
(id STRING NOT NULL PRIMARY KEY, n INT) WITH (engine='document_strict')",
)
.await
.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('a', 1)")
.await
.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('b', 2)")
.await
.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn rollback_to_reverts_post_savepoint_insert_keeps_pre_savepoint() {
let server = TestServer::start().await;
setup_doc(&server).await;
server.exec("BEGIN").await.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('p', 10)")
.await
.unwrap();
server.exec("SAVEPOINT s1").await.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('q', 20)")
.await
.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'p'").await,
Some("10".to_string())
);
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'q'").await,
Some("20".to_string())
);
server.exec("ROLLBACK TO SAVEPOINT s1").await.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'p'").await,
Some("10".to_string()),
"pre-savepoint insert must survive ROLLBACK TO"
);
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'q'").await,
None,
"post-savepoint insert must be invisible after ROLLBACK TO (the bug)"
);
server
.exec("INSERT INTO t (id, n) VALUES ('r', 30)")
.await
.unwrap();
server.exec("COMMIT").await.unwrap();
assert_eq!(
server
.query_text("SELECT n FROM t WHERE id = 'p'")
.await
.unwrap(),
vec!["10"]
);
assert_eq!(
server
.query_text("SELECT n FROM t WHERE id = 'r'")
.await
.unwrap(),
vec!["30"]
);
assert!(
server
.query_text("SELECT n FROM t WHERE id = 'q'")
.await
.unwrap()
.is_empty(),
"rolled-back insert must not persist"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn nested_rollback_to_outer_discards_inner_savepoint() {
let server = TestServer::start().await;
setup_doc(&server).await;
server.exec("BEGIN").await.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('p', 10)")
.await
.unwrap();
server.exec("SAVEPOINT s1").await.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('q', 20)")
.await
.unwrap();
server.exec("SAVEPOINT s2").await.unwrap();
server
.exec("INSERT INTO t (id, n) VALUES ('r', 30)")
.await
.unwrap();
server.exec("ROLLBACK TO SAVEPOINT s1").await.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'p'").await,
Some("10".to_string())
);
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'q'").await,
None,
"inner savepoint write must be discarded"
);
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'r'").await,
None,
"innermost write must be discarded"
);
server.exec("COMMIT").await.unwrap();
assert_eq!(
server
.query_text("SELECT n FROM t WHERE id = 'p'")
.await
.unwrap(),
vec!["10"]
);
assert!(
server
.query_text("SELECT n FROM t WHERE id = 'q'")
.await
.unwrap()
.is_empty()
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn rollback_to_reverts_update_to_prior_value() {
let server = TestServer::start().await;
setup_doc(&server).await;
server.exec("BEGIN").await.unwrap();
server
.exec("UPDATE t SET n = 10 WHERE id = 'a'")
.await
.unwrap();
server.exec("SAVEPOINT s").await.unwrap();
server
.exec("UPDATE t SET n = 20 WHERE id = 'a'")
.await
.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'a'").await,
Some("20".to_string())
);
server.exec("ROLLBACK TO SAVEPOINT s").await.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'a'").await,
Some("10".to_string()),
"ROLLBACK TO must restore the value staged before the savepoint, not the base"
);
server.exec("COMMIT").await.unwrap();
assert_eq!(
server
.query_text("SELECT n FROM t WHERE id = 'a'")
.await
.unwrap(),
vec!["10"]
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn rollback_to_reverts_delete() {
let server = TestServer::start().await;
setup_doc(&server).await;
server.exec("BEGIN").await.unwrap();
server.exec("SAVEPOINT s").await.unwrap();
server.exec("DELETE FROM t WHERE id = 'b'").await.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'b'").await,
None,
"delete must be visible in-transaction before rollback"
);
server.exec("ROLLBACK TO SAVEPOINT s").await.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM t WHERE id = 'b'").await,
Some("2".to_string()),
"ROLLBACK TO must restore the deleted row"
);
server.exec("COMMIT").await.unwrap();
assert_eq!(
server
.query_text("SELECT n FROM t WHERE id = 'b'")
.await
.unwrap(),
vec!["2"]
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn release_then_rollback_to_errors_3b001() {
let server = TestServer::start().await;
setup_doc(&server).await;
server.exec("BEGIN").await.unwrap();
server.exec("SAVEPOINT s1").await.unwrap();
server.exec("RELEASE SAVEPOINT s1").await.unwrap();
let err = server
.client
.simple_query("ROLLBACK TO SAVEPOINT s1")
.await
.expect_err("ROLLBACK TO a released savepoint must error");
let db_err = err.as_db_error().expect("expected a DbError");
assert_eq!(
db_err.code().code(),
"3B001",
"expected 3B001 for an unknown savepoint, got {}",
db_err.code().code()
);
server.client.simple_query("ROLLBACK").await.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn release_unknown_savepoint_errors_3b001() {
let server = TestServer::start().await;
setup_doc(&server).await;
server.exec("BEGIN").await.unwrap();
let err = server
.client
.simple_query("RELEASE SAVEPOINT nope")
.await
.expect_err("RELEASE of an unknown savepoint must error");
let db_err = err.as_db_error().expect("expected a DbError");
assert_eq!(db_err.code().code(), "3B001");
server.client.simple_query("ROLLBACK").await.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn savepoint_outside_transaction_errors_25p01() {
let server = TestServer::start().await;
setup_doc(&server).await;
for sql in [
"SAVEPOINT s1",
"ROLLBACK TO SAVEPOINT s1",
"RELEASE SAVEPOINT s1",
] {
let err = match server.client.simple_query(sql).await {
Ok(_) => panic!("`{sql}` outside a transaction should have errored"),
Err(e) => e,
};
let db_err = err.as_db_error().expect("expected a DbError");
assert_eq!(
db_err.code().code(),
"25P01",
"`{sql}` outside a transaction must be 25P01, got {}",
db_err.code().code()
);
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn kv_value_change_reverted_by_rollback_to() {
let server = TestServer::start().await;
server
.exec("CREATE COLLECTION c (key TEXT PRIMARY KEY, n INT) WITH (engine='kv')")
.await
.unwrap();
server
.exec("INSERT INTO c (key, n) VALUES ('a', 1)")
.await
.unwrap();
server.exec("BEGIN").await.unwrap();
server.exec("SAVEPOINT s").await.unwrap();
server
.exec("UPDATE c SET n = 99 WHERE key = 'a'")
.await
.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM c WHERE key = 'a'").await,
Some("99".to_string())
);
server.exec("ROLLBACK TO SAVEPOINT s").await.unwrap();
assert_eq!(
read_n(&server, "SELECT n FROM c WHERE key = 'a'").await,
Some("1".to_string()),
"ROLLBACK TO must revert the KV value overlay to its pre-savepoint state"
);
server.exec("COMMIT").await.unwrap();
assert_eq!(
server
.query_text("SELECT n FROM c WHERE key = 'a'")
.await
.unwrap(),
vec!["1"]
);
}