mod common;
use common::native_harness::{do_handshake, send_sql};
use common::pgwire_harness::TestServer;
use nodedb_types::protocol::HelloFrame;
use nodedb_types::protocol::opcodes::ResponseStatus;
use nodedb_types::value::Value;
use tokio::net::TcpStream;
async fn native_session(srv: &TestServer) -> TcpStream {
let addr = format!("127.0.0.1:{}", srv.native_port)
.parse()
.expect("native addr");
let (stream, _ack) = do_handshake(addr, &HelloFrame::current())
.await
.expect("native handshake");
stream
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn native_rollback_to_savepoint_reverts_later_staged_write() {
let server = TestServer::start().await;
server
.exec(
"CREATE COLLECTION native_sp_kv (key TEXT PRIMARY KEY, n INT) \
WITH (engine='kv')",
)
.await
.unwrap();
let mut stream = native_session(&server).await;
let mut seq = 1u64;
let begin = send_sql(&mut stream, seq, "BEGIN").await;
assert_eq!(begin.status, ResponseStatus::Ok, "BEGIN must succeed");
seq += 1;
let ins_a = send_sql(
&mut stream,
seq,
"INSERT INTO native_sp_kv (key, n) VALUES ('a', 1)",
)
.await;
assert_eq!(ins_a.status, ResponseStatus::Ok);
seq += 1;
let sp = send_sql(&mut stream, seq, "SAVEPOINT s1").await;
assert_eq!(sp.status, ResponseStatus::Ok, "SAVEPOINT must succeed");
seq += 1;
let ins_b = send_sql(
&mut stream,
seq,
"INSERT INTO native_sp_kv (key, n) VALUES ('b', 2)",
)
.await;
assert_eq!(ins_b.status, ResponseStatus::Ok);
seq += 1;
let both = send_sql(
&mut stream,
seq,
"SELECT n FROM native_sp_kv WHERE key = 'b'",
)
.await;
assert_eq!(both.status, ResponseStatus::Ok);
assert_eq!(
both.rows
.expect("b visible pre-rollback")
.first()
.and_then(|r| r.first())
.cloned(),
Some(Value::Integer(2)),
"row staged after SAVEPOINT must be visible before ROLLBACK TO"
);
seq += 1;
let rb = send_sql(&mut stream, seq, "ROLLBACK TO SAVEPOINT s1").await;
assert_eq!(
rb.status,
ResponseStatus::Ok,
"ROLLBACK TO SAVEPOINT must succeed"
);
seq += 1;
let after_b = send_sql(
&mut stream,
seq,
"SELECT n FROM native_sp_kv WHERE key = 'b'",
)
.await;
assert_eq!(after_b.status, ResponseStatus::Ok);
assert!(
after_b.rows.map(|r| r.is_empty()).unwrap_or(true),
"row staged after SAVEPOINT must be gone after ROLLBACK TO"
);
seq += 1;
let after_a = send_sql(
&mut stream,
seq,
"SELECT n FROM native_sp_kv WHERE key = 'a'",
)
.await;
assert_eq!(after_a.status, ResponseStatus::Ok);
assert_eq!(
after_a
.rows
.expect("a visible post-rollback")
.first()
.and_then(|r| r.first())
.cloned(),
Some(Value::Integer(1)),
"row staged before SAVEPOINT must survive ROLLBACK TO"
);
seq += 1;
let commit = send_sql(&mut stream, seq, "COMMIT").await;
assert_eq!(commit.status, ResponseStatus::Ok, "COMMIT must succeed");
let committed_a = server
.query_text("SELECT n FROM native_sp_kv WHERE key = 'a'")
.await
.unwrap();
assert_eq!(committed_a, vec!["1".to_string()], "'a' must persist");
let committed_b = server
.query_text("SELECT n FROM native_sp_kv WHERE key = 'b'")
.await
.unwrap();
assert!(
committed_b.is_empty(),
"'b' must NOT persist after ROLLBACK TO, found: {committed_b:?}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn native_release_unknown_savepoint_errors_3b001() {
let server = TestServer::start().await;
let mut stream = native_session(&server).await;
let mut seq = 1u64;
let begin = send_sql(&mut stream, seq, "BEGIN").await;
assert_eq!(begin.status, ResponseStatus::Ok);
seq += 1;
let release = send_sql(&mut stream, seq, "RELEASE SAVEPOINT nope").await;
assert_eq!(
release.status,
ResponseStatus::Error,
"RELEASE of an unknown savepoint must fail"
);
let err = release.error.expect("error payload expected");
assert_eq!(
err.code, "3B001",
"unknown savepoint must surface SQLSTATE 3B001, got {}",
err.code
);
seq += 1;
let rollback = send_sql(&mut stream, seq, "ROLLBACK").await;
assert_eq!(rollback.status, ResponseStatus::Ok);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn native_savepoint_outside_transaction_errors_25p01() {
let server = TestServer::start().await;
let mut stream = native_session(&server).await;
let sp = send_sql(&mut stream, 1, "SAVEPOINT s1").await;
assert_eq!(
sp.status,
ResponseStatus::Error,
"SAVEPOINT outside a transaction must fail"
);
let err = sp.error.expect("error payload expected");
assert_eq!(
err.code, "25P01",
"SAVEPOINT outside a transaction must surface SQLSTATE 25P01, got {}",
err.code
);
}