mod common;
use common::pgwire_harness::TestServer;
async fn setup(server: &TestServer) {
server
.exec("CREATE COLLECTION acct (key TEXT PRIMARY KEY, balance INT) WITH (engine='kv')")
.await
.unwrap();
server
.exec("CREATE COLLECTION items (key TEXT PRIMARY KEY, name TEXT) WITH (engine='kv')")
.await
.unwrap();
}
fn json_of(rows: &[String]) -> serde_json::Value {
serde_json::from_str(&rows[0]).expect("TRANSFER* result must be JSON")
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn transfer_in_tx_chains_and_commits() {
let server = TestServer::start().await;
setup(&server).await;
server
.exec("INSERT INTO acct (key, balance) VALUES ('a', 100)")
.await
.unwrap();
server
.exec("INSERT INTO acct (key, balance) VALUES ('b', 10)")
.await
.unwrap();
server.exec("BEGIN").await.unwrap();
let rows = server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 30)")
.await
.unwrap();
let v = json_of(&rows);
assert_eq!(v["source_balance"], 70.0);
assert_eq!(v["dest_balance"], 40.0);
let rows2 = server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 10)")
.await
.unwrap();
let v2 = json_of(&rows2);
assert_eq!(
v2["source_balance"], 60.0,
"second in-tx transfer must chain off the first staged source balance"
);
assert_eq!(
v2["dest_balance"], 50.0,
"second in-tx transfer must chain off the first staged dest balance"
);
server.exec("COMMIT").await.unwrap();
let after = json_of(
&server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 1)")
.await
.unwrap(),
);
assert_eq!(
after["source_balance"], 59.0,
"post-COMMIT source balance must be the chained value 60 (−1 = 59): {after}"
);
assert_eq!(
after["dest_balance"], 51.0,
"post-COMMIT dest balance must be the chained value 50 (+1 = 51): {after}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn transfer_in_tx_rollback_reverts_both_balances() {
let server = TestServer::start().await;
setup(&server).await;
server
.exec("INSERT INTO acct (key, balance) VALUES ('a', 100)")
.await
.unwrap();
server
.exec("INSERT INTO acct (key, balance) VALUES ('b', 10)")
.await
.unwrap();
server.exec("BEGIN").await.unwrap();
let rows = server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 50)")
.await
.unwrap();
assert_eq!(json_of(&rows)["source_balance"], 50.0);
server.exec("ROLLBACK").await.unwrap();
let after = server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 90)")
.await
.unwrap();
assert_eq!(
json_of(&after)["source_balance"],
10.0,
"ROLLBACK must revert 'a' to its original balance (100 - 90 = 10)"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn transfer_in_tx_insufficient_balance_fails_at_statement_time() {
let server = TestServer::start().await;
setup(&server).await;
server
.exec("INSERT INTO acct (key, balance) VALUES ('a', 100)")
.await
.unwrap();
server
.exec("INSERT INTO acct (key, balance) VALUES ('b', 0)")
.await
.unwrap();
server.exec("BEGIN").await.unwrap();
server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 80)")
.await
.unwrap();
let err = server
.query_text("SELECT TRANSFER('acct', 'a', 'b', 'balance', 21)")
.await
.unwrap_err();
assert!(
err.to_string().contains("20"),
"insufficient-balance error must report the STAGED balance (20), got: {err}"
);
server.exec("ROLLBACK").await.unwrap();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn transfer_item_in_tx_visible_to_a_second_call_and_reverts_on_rollback() {
let server = TestServer::start().await;
setup(&server).await;
server
.exec("INSERT INTO items (key, name) VALUES ('ownerA:sword', 'Sword')")
.await
.unwrap();
server.exec("BEGIN").await.unwrap();
let rows = server
.query_text("SELECT TRANSFER_ITEM('items', 'items', 'sword', 'ownerA', 'ownerB')")
.await
.unwrap();
let v = json_of(&rows);
assert_eq!(v["item_key"], "ownerA:sword");
assert_eq!(v["dest_key"], "ownerB:sword");
let err = server
.query_text("SELECT TRANSFER_ITEM('items', 'items', 'sword', 'ownerA', 'ownerC')")
.await
.unwrap_err();
assert!(
err.to_string().to_lowercase().contains("not found")
|| err.to_string().contains("22023")
|| err.to_string().contains("NOT_FOUND"),
"second move from the now-tombstoned source must fail NotFound, got: {err}"
);
server.exec("ROLLBACK").await.unwrap();
let after = server
.query_text("SELECT TRANSFER_ITEM('items', 'items', 'sword', 'ownerA', 'ownerB')")
.await
.unwrap();
let v_after = json_of(&after);
assert_eq!(v_after["item_key"], "ownerA:sword");
}