use super::handler::PgConnectionHandler;
use crate::EmbeddedDatabase;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, DuplexStream};
fn test_handler(db: Arc<EmbeddedDatabase>) -> (PgConnectionHandler<DuplexStream>, DuplexStream) {
let (server, client) = tokio::io::duplex(8 << 20);
(PgConnectionHandler::new_for_tests(db, server), client)
}
async fn drain(client: &mut DuplexStream) -> Vec<u8> {
let mut out = Vec::new();
let mut buf = [0u8; 65536];
loop {
match tokio::time::timeout(std::time::Duration::from_millis(50), client.read(&mut buf)).await {
Ok(Ok(0)) => break,
Ok(Ok(n)) => out.extend_from_slice(&buf[..n]),
_ => break,
}
}
out
}
fn parse_messages(bytes: &[u8]) -> Vec<(u8, Vec<u8>)> {
let mut out = Vec::new();
let mut pos = 0;
while pos + 5 <= bytes.len() {
let ty = bytes[pos];
let len = i32::from_be_bytes([bytes[pos + 1], bytes[pos + 2], bytes[pos + 3], bytes[pos + 4]]) as usize;
let end = pos + 1 + len;
assert!(end <= bytes.len(), "truncated message {ty:#x}");
out.push((ty, bytes[pos + 5..end].to_vec()));
pos = end;
}
out
}
fn decode_data_row(payload: &[u8]) -> Vec<Option<Vec<u8>>> {
let ncols = i16::from_be_bytes([payload[0], payload[1]]) as usize;
let mut pos = 2;
let mut cols = Vec::with_capacity(ncols);
for _ in 0..ncols {
let len = i32::from_be_bytes([payload[pos], payload[pos + 1], payload[pos + 2], payload[pos + 3]]);
pos += 4;
if len < 0 {
cols.push(None);
} else {
let len = len as usize;
cols.push(Some(payload[pos..pos + len].to_vec()));
pos += len;
}
}
cols
}
fn data_rows(bytes: &[u8]) -> Vec<Vec<Option<Vec<u8>>>> {
parse_messages(bytes)
.into_iter()
.filter(|(ty, _)| *ty == b'D')
.map(|(_, payload)| decode_data_row(&payload))
.collect()
}
fn wide_test_db(rows: usize) -> Arc<EmbeddedDatabase> {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute(
"CREATE TABLE wide (id INT PRIMARY KEY, a TEXT, b TEXT, c BIGINT, d DOUBLE PRECISION, \
e TEXT, f INT, g TEXT, h BIGINT, i TEXT, j DOUBLE PRECISION, k TEXT)",
)
.expect("create");
for n in 0..rows {
db.execute(&format!(
"INSERT INTO wide VALUES ({n}, 'alpha-{n}', 'beta-{n}', {}, {}.5, 'gamma-{n}', {}, \
'delta-{n}', {}, 'epsilon-{n}', {}.25, 'zeta-{n}')",
n * 1000,
n,
n % 97,
(n as i64) * 7,
n
))
.expect("insert");
}
db
}
#[tokio::test]
async fn extended_select_matches_simple_query_data_rows() {
let db = wide_test_db(50);
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_single_query("SELECT * FROM wide ORDER BY id")
.await
.expect("simple query");
let simple_rows = data_rows(&drain(&mut client).await);
assert_eq!(simple_rows.len(), 50);
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended("s1".into(), "SELECT * FROM wide ORDER BY id".into(), vec![])
.await
.expect("parse");
handler
.handle_bind_extended("p1".into(), "s1".into(), vec![], vec![], vec![])
.await
.expect("bind");
handler.handle_execute_extended("p1".into(), 0).await.expect("execute");
let extended_rows = data_rows(&drain(&mut client).await);
assert_eq!(extended_rows, simple_rows, "extended DataRows must be byte-identical");
}
#[tokio::test]
async fn extended_select_with_many_params() {
let db = wide_test_db(30);
let (mut handler, mut client) = test_handler(db);
let sql = "SELECT id, a, c FROM wide WHERE id = $1 OR id = $2 OR id = $3 OR id = $4 \
OR id = $5 OR id = $6 OR id = $7 OR id = $8 ORDER BY id";
handler
.handle_parse_extended("s2".into(), sql.into(), vec![23; 8])
.await
.expect("parse");
let params: Vec<Option<Vec<u8>>> = [1, 3, 5, 7, 11, 13, 17, 19]
.iter()
.map(|n: &i32| Some(n.to_string().into_bytes()))
.collect();
handler
.handle_bind_extended("p2".into(), "s2".into(), vec![0; 8], params, vec![])
.await
.expect("bind");
handler.handle_execute_extended("p2".into(), 0).await.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 8);
assert_eq!(rows[0][0].as_deref(), Some(b"1".as_ref()));
assert_eq!(rows[0][1].as_deref(), Some(b"alpha-1".as_ref()));
assert_eq!(rows[0][2].as_deref(), Some(b"1000".as_ref()));
assert_eq!(rows[7][0].as_deref(), Some(b"19".as_ref()));
assert_eq!(rows[7][2].as_deref(), Some(b"19000".as_ref()));
}
#[tokio::test]
async fn extended_select_null_handling() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE n (id INT PRIMARY KEY, v TEXT)")
.expect("create");
db.execute("INSERT INTO n VALUES (1, NULL), (2, 'x')").expect("insert");
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended("s3".into(), "SELECT v FROM n ORDER BY id".into(), vec![])
.await
.expect("parse");
handler
.handle_bind_extended("p3".into(), "s3".into(), vec![], vec![], vec![])
.await
.expect("bind");
handler.handle_execute_extended("p3".into(), 0).await.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 2);
assert_eq!(rows[0][0], None, "NULL must be the -1 sentinel");
assert_eq!(rows[1][0].as_deref(), Some(b"x".as_ref()));
}
#[tokio::test]
async fn extended_select_binary_format_fallback() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE bi (id INT PRIMARY KEY)").expect("create");
db.execute("INSERT INTO bi VALUES (305419896)").expect("insert");
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended("s4".into(), "SELECT id FROM bi".into(), vec![])
.await
.expect("parse");
handler
.handle_bind_extended("p4".into(), "s4".into(), vec![], vec![], vec![1])
.await
.expect("bind");
handler.handle_execute_extended("p4".into(), 0).await.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 1);
assert_eq!(
rows[0][0].as_deref(),
Some([0x12u8, 0x34, 0x56, 0x78].as_ref()),
"int4 must arrive as 4-byte big-endian binary"
);
}
#[tokio::test]
async fn repeated_execute_serves_identical_rows() {
let db = wide_test_db(20);
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended(
"rep".into(),
"SELECT id, a, c FROM wide WHERE id = $1 OR id = $2 ORDER BY id".into(),
vec![23, 23],
)
.await
.expect("parse");
let mut first_rows = None;
for i in 0..5 {
let portal = format!("rp{i}");
handler
.handle_bind_extended(
portal.clone(),
"rep".into(),
vec![0, 0],
vec![Some(b"3".to_vec()), Some(b"7".to_vec())],
vec![],
)
.await
.expect("bind");
handler.handle_execute_extended(portal, 0).await.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 2);
assert_eq!(rows[0][0].as_deref(), Some(b"3".as_ref()));
assert_eq!(rows[1][0].as_deref(), Some(b"7".as_ref()));
match &first_rows {
None => first_rows = Some(rows),
Some(expected) => assert_eq!(&rows, expected, "execute #{i} diverged"),
}
}
}
#[tokio::test]
async fn ddl_between_executes_invalidates_pinned_plan() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE evolve (id INT PRIMARY KEY)").expect("create");
db.execute("INSERT INTO evolve VALUES (1)").expect("insert");
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_parse_extended("ev".into(), "SELECT * FROM evolve".into(), vec![])
.await
.expect("parse");
handler
.handle_bind_extended("evp1".into(), "ev".into(), vec![], vec![], vec![])
.await
.expect("bind");
handler
.handle_execute_extended("evp1".into(), 0)
.await
.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].len(), 1, "one column before DDL");
db.execute("ALTER TABLE evolve ADD COLUMN extra TEXT").expect("alter");
handler
.handle_bind_extended("evp2".into(), "ev".into(), vec![], vec![], vec![])
.await
.expect("bind");
handler
.handle_execute_extended("evp2".into(), 0)
.await
.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 1);
assert_eq!(
rows[0].len(),
2,
"SELECT * must see the new column — stale pinned plan detected"
);
}
#[tokio::test]
async fn catalog_query_still_served_after_parse_decision() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended("cat".into(), "SELECT version()".into(), vec![])
.await
.expect("parse");
handler
.handle_bind_extended("catp".into(), "cat".into(), vec![], vec![], vec![])
.await
.expect("bind");
handler
.handle_execute_extended("catp".into(), 0)
.await
.expect("execute");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 1);
let version = String::from_utf8(rows[0][0].clone().expect("version text")).expect("utf8");
assert!(version.contains("PostgreSQL"), "catalog version() reply: {version}");
}
#[tokio::test]
#[ignore]
async fn probe_w2_repeated_prepared_execute() {
let db = wide_test_db(1_000);
let (server, mut client) = tokio::io::duplex(1 << 20);
let mut handler = PgConnectionHandler::new_for_tests(db, server);
let drain_task = tokio::spawn(async move {
let mut buf = vec![0u8; 1 << 20];
let mut total = 0u64;
while let Ok(n) = client.read(&mut buf).await {
if n == 0 {
break;
}
total += n as u64;
}
total
});
handler
.handle_parse_extended(
"probe2".into(),
"SELECT id, a, c FROM wide WHERE id = $1".into(),
vec![23],
)
.await
.expect("parse");
const ITERS: usize = 20_000;
let start = std::time::Instant::now();
for i in 0..ITERS {
let a = (i % 1000).to_string().into_bytes();
handler
.handle_bind_extended("".into(), "probe2".into(), vec![0], vec![Some(a)], vec![])
.await
.expect("bind");
handler.handle_execute_extended("".into(), 0).await.expect("execute");
}
let elapsed = start.elapsed();
drop(handler);
let bytes = drain_task.await.expect("drain");
println!(
"W2 probe: {ITERS} Bind+Execute of prepared point-SELECT in {:?} ({:.1} us/exec, {:.1} KB total)",
elapsed,
elapsed.as_secs_f64() * 1e6 / ITERS as f64,
bytes as f64 / 1024.0
);
}
#[tokio::test]
#[ignore]
async fn probe_w1_extended_select_10k_wide_rows() {
let db = wide_test_db(10_000);
let (server, mut client) = tokio::io::duplex(1 << 20);
let mut handler = PgConnectionHandler::new_for_tests(db, server);
let drain_task = tokio::spawn(async move {
let mut buf = vec![0u8; 1 << 20];
let mut total = 0u64;
while let Ok(n) = client.read(&mut buf).await {
if n == 0 {
break;
}
total += n as u64;
}
total
});
handler
.handle_parse_extended("probe".into(), "SELECT * FROM wide".into(), vec![])
.await
.expect("parse");
const ITERS: usize = 30;
let start = std::time::Instant::now();
for i in 0..ITERS {
let portal = format!("portal{i}");
handler
.handle_bind_extended(portal.clone(), "probe".into(), vec![], vec![], vec![])
.await
.expect("bind");
handler.handle_execute_extended(portal, 0).await.expect("execute");
}
let elapsed = start.elapsed();
drop(handler);
let bytes = drain_task.await.expect("drain");
println!(
"W1 probe: {ITERS} extended Executes of 10k×12 rows in {:?} ({:.2} ms/exec, {:.1} MB total)",
elapsed,
elapsed.as_secs_f64() * 1000.0 / ITERS as f64,
bytes as f64 / (1024.0 * 1024.0)
);
}
#[tokio::test]
async fn pipelined_executes_emit_exactly_one_ready_for_query() {
use super::messages::FrontendMessage;
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
db.execute("CREATE TABLE t (id INT PRIMARY KEY)").unwrap();
db.execute("INSERT INTO t VALUES (1),(2),(3)").unwrap();
let (mut handler, mut client) = test_handler(db);
let pipeline = vec![
FrontendMessage::Parse {
statement_name: "s1".into(),
query: "SELECT id FROM t ORDER BY id".into(),
param_types: vec![],
},
FrontendMessage::Bind {
portal_name: "p1".into(),
statement_name: "s1".into(),
param_formats: vec![],
params: vec![],
result_formats: vec![],
},
FrontendMessage::Execute {
portal_name: "p1".into(),
max_rows: 0,
},
FrontendMessage::Bind {
portal_name: "p2".into(),
statement_name: "s1".into(),
param_formats: vec![],
params: vec![],
result_formats: vec![],
},
FrontendMessage::Execute {
portal_name: "p2".into(),
max_rows: 0,
},
FrontendMessage::Sync,
];
for msg in pipeline {
handler.handle_message(msg).await.expect("dispatch");
}
let out = drain(&mut client).await;
let types: Vec<u8> = parse_messages(&out).iter().map(|(t, _)| *t).collect();
let render: String = types.iter().map(|&t| t as char).collect();
assert_eq!(
types.iter().filter(|&&t| t == b'Z').count(),
1,
"exactly one ReadyForQuery for the whole pipeline: {render}"
);
assert_eq!(*types.last().unwrap(), b'Z', "ReadyForQuery must be last: {render}");
assert_eq!(
types.iter().filter(|&&t| t == b'1').count(),
1,
"one ParseComplete: {render}"
);
assert_eq!(
types.iter().filter(|&&t| t == b'2').count(),
2,
"two BindComplete: {render}"
);
assert_eq!(
types.iter().filter(|&&t| t == b'C').count(),
2,
"two CommandComplete: {render}"
);
assert_eq!(
types.iter().filter(|&&t| t == b'D').count(),
6,
"six DataRows (3 per Execute): {render}"
);
}
fn command_tags(bytes: &[u8]) -> Vec<String> {
parse_messages(bytes)
.into_iter()
.filter(|(t, _)| *t == b'C')
.map(|(_, p)| String::from_utf8_lossy(p.split(|&b| b == 0).next().unwrap_or(&[])).to_string())
.collect()
}
fn param_status(bytes: &[u8]) -> Vec<(String, String)> {
parse_messages(bytes)
.into_iter()
.filter(|(t, _)| *t == b'S')
.map(|(_, p)| {
let mut it = p.split(|&b| b == 0);
let name = String::from_utf8_lossy(it.next().unwrap_or(&[])).to_string();
let val = String::from_utf8_lossy(it.next().unwrap_or(&[])).to_string();
(name, val)
})
.collect()
}
fn first_data_row_text(bytes: &[u8]) -> Option<String> {
data_rows(bytes)
.into_iter()
.next()
.and_then(|r| r.into_iter().next().flatten())
.map(|b| String::from_utf8_lossy(&b).to_string())
}
#[tokio::test]
async fn helios_fast_autocommit_set_show_roundtrip() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler
.handle_single_query("SET helios.fast_autocommit = on")
.await
.unwrap();
let out = drain(&mut client).await;
assert!(
param_status(&out)
.iter()
.any(|(n, v)| n == "helios.fast_autocommit" && v == "on"),
"expected GUC_REPORT helios.fast_autocommit=on, got {:?}",
param_status(&out)
);
assert!(command_tags(&out).iter().any(|t| t == "SET"), "expected SET tag");
handler
.handle_single_query("SHOW helios.fast_autocommit")
.await
.unwrap();
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("on"),
"SHOW must reflect the SET"
);
handler
.handle_single_query("SET helios.fast_autocommit = banana")
.await
.unwrap();
let out = drain(&mut client).await;
assert!(
parse_messages(&out).iter().any(|(t, _)| *t == b'E'),
"invalid value must produce an ErrorResponse"
);
}
#[tokio::test]
async fn discard_all_resets_session_guc() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler
.handle_single_query("SET helios.fast_autocommit = on")
.await
.unwrap();
let _ = drain(&mut client).await;
handler.handle_single_query("DISCARD ALL").await.unwrap();
let out = drain(&mut client).await;
assert!(
command_tags(&out).iter().any(|t| t == "DISCARD ALL"),
"expected DISCARD ALL tag, got {:?}",
command_tags(&out)
);
handler
.handle_single_query("SHOW helios.fast_autocommit")
.await
.unwrap();
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("off"),
"DISCARD ALL must reset helios.fast_autocommit to off"
);
}
#[tokio::test]
async fn bytea_text_output_is_hex_not_raw_bytes() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler.handle_single_query("CREATE TABLE wbt (b bytea)").await.unwrap();
let _ = drain(&mut client).await;
handler
.handle_single_query("INSERT INTO wbt VALUES ('\\x5a5b5c5d5e')")
.await
.unwrap();
let _ = drain(&mut client).await;
handler.handle_single_query("SELECT b FROM wbt").await.unwrap();
let out = drain(&mut client).await;
assert_eq!(
first_data_row_text(&out).as_deref(),
Some("\\x5a5b5c5d5e"),
"bytea text output must be `\\x`-hex encoded, not raw bytes"
);
}
fn row_description(bytes: &[u8]) -> Vec<(String, i32)> {
let mut out = Vec::new();
for (ty, payload) in parse_messages(bytes) {
if ty != b'T' {
continue;
}
let nfields = i16::from_be_bytes([payload[0], payload[1]]) as usize;
let mut pos = 2;
for _ in 0..nfields {
let name_end = pos + payload[pos..].iter().position(|&b| b == 0).expect("field name cstring");
let name = String::from_utf8_lossy(&payload[pos..name_end]).to_string();
pos = name_end + 1;
let oid_pos = pos + 4 + 2; let oid = i32::from_be_bytes([
payload[oid_pos],
payload[oid_pos + 1],
payload[oid_pos + 2],
payload[oid_pos + 3],
]);
out.push((name, oid));
pos += 18; }
}
out
}
#[tokio::test]
async fn parse_seeds_shared_plan_for_select() {
let db = wide_test_db(3);
let (mut handler, _client) = test_handler(db);
handler
.handle_parse_extended("sp".into(), "SELECT id, a FROM wide WHERE id = $1".into(), vec![23])
.await
.expect("parse");
let stmt = handler
.prepared_statements
.get_statement("sp")
.expect("get")
.expect("stmt present");
assert!(
stmt.cached_plan.is_some(),
"SELECT Parse must seed cached_plan from the shared parameterized plan cache"
);
}
#[tokio::test]
async fn dml_returning_keeps_private_schema_path() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE ins (id INT PRIMARY KEY, v TEXT)")
.expect("create");
let (mut handler, _client) = test_handler(db);
handler
.handle_parse_extended(
"dr".into(),
"INSERT INTO ins (id, v) VALUES ($1, $2) RETURNING id".into(),
vec![23, 25],
)
.await
.expect("parse");
let stmt = handler
.prepared_statements
.get_statement("dr")
.expect("get")
.expect("stmt present");
assert!(
stmt.cached_plan.is_none(),
"DML-RETURNING must not be seeded via the shared plan path (empty LogicalPlan::schema)"
);
let schema = stmt
.result_schema
.expect("RETURNING must still yield a result schema (RowDescription), not NoData");
assert_eq!(
schema.columns.iter().map(|c| c.name.as_str()).collect::<Vec<_>>(),
vec!["id"],
"RETURNING column must survive the private fallback path"
);
}
#[tokio::test]
async fn describe_reports_pg_type_oids() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE acct (id INT PRIMARY KEY, name TEXT, bal NUMERIC, big BIGINT, code VARCHAR(8))")
.expect("create");
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended(
"d".into(),
"SELECT id, name, bal, big, code FROM acct WHERE id = $1".into(),
vec![23],
)
.await
.expect("parse");
handler
.handle_describe_extended(super::messages::DescribeTarget::Statement, "d".into())
.await
.expect("describe");
let fields = row_description(&drain(&mut client).await);
assert_eq!(
fields,
vec![
("id".to_string(), 23),
("name".to_string(), 25),
("bal".to_string(), 1700),
("big".to_string(), 20),
("code".to_string(), 1043),
],
"Describe RowDescription names + pg_type OIDs (numeric MUST be 1700)"
);
}
#[tokio::test]
async fn describe_aggregate_alias_names_and_types() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE ev (id INT PRIMARY KEY, k TEXT)")
.expect("create");
db.execute("INSERT INTO ev VALUES (1,'a'),(2,'b'),(3,'a')")
.expect("insert");
let (mut handler, mut client) = test_handler(db);
handler
.handle_parse_extended("ag".into(), "SELECT count(*) AS n FROM ev".into(), vec![])
.await
.expect("parse");
handler
.handle_describe_extended(super::messages::DescribeTarget::Statement, "ag".into())
.await
.expect("describe");
let fields = row_description(&drain(&mut client).await);
assert_eq!(
fields,
vec![("n".to_string(), 20)],
"count(*) AS n → int8 (OID 20) named n"
);
}
#[tokio::test]
async fn view_redefine_invalidates_describe_schema() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE vt (id INT PRIMARY KEY, name TEXT)")
.expect("create table");
db.execute("CREATE VIEW vv AS SELECT id FROM vt").expect("create view");
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_parse_extended("v1".into(), "SELECT * FROM vv".into(), vec![])
.await
.expect("parse v1");
handler
.handle_describe_extended(super::messages::DescribeTarget::Statement, "v1".into())
.await
.expect("describe v1");
let before = row_description(&drain(&mut client).await);
assert_eq!(
before,
vec![("id".to_string(), 23)],
"view Describe before redefine: single int4 column id"
);
db.execute("CREATE OR REPLACE VIEW vv AS SELECT id, name FROM vt")
.expect("replace view");
handler
.handle_parse_extended("v2".into(), "SELECT * FROM vv".into(), vec![])
.await
.expect("parse v2");
handler
.handle_describe_extended(super::messages::DescribeTarget::Statement, "v2".into())
.await
.expect("describe v2");
let after = row_description(&drain(&mut client).await);
assert_eq!(
after,
vec![("id".to_string(), 23), ("name".to_string(), 25)],
"view Describe after CREATE OR REPLACE must reflect the added TEXT column \
(a stale shared-plan-cache entry would still report only id)"
);
}
#[tokio::test]
async fn view_redefine_via_extended_protocol_invalidates_describe_schema() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE vt (id INT PRIMARY KEY, name TEXT)")
.expect("create table");
db.execute("CREATE VIEW vv AS SELECT id FROM vt").expect("create view");
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_parse_extended("v1".into(), "SELECT * FROM vv".into(), vec![])
.await
.expect("parse v1");
handler
.handle_describe_extended(super::messages::DescribeTarget::Statement, "v1".into())
.await
.expect("describe v1");
let before = row_description(&drain(&mut client).await);
assert_eq!(
before,
vec![("id".to_string(), 23)],
"view Describe before redefine: single int4 column id"
);
handler
.handle_parse_extended(
"ddl".into(),
"CREATE OR REPLACE VIEW vv AS SELECT id, name FROM vt".into(),
vec![],
)
.await
.expect("parse ddl");
handler
.handle_bind_extended("ddlp".into(), "ddl".into(), vec![], vec![], vec![])
.await
.expect("bind ddl");
handler
.handle_execute_extended("ddlp".into(), 0)
.await
.expect("execute ddl");
let _ = drain(&mut client).await;
handler
.handle_parse_extended("v2".into(), "SELECT * FROM vv".into(), vec![])
.await
.expect("parse v2");
handler
.handle_describe_extended(super::messages::DescribeTarget::Statement, "v2".into())
.await
.expect("describe v2");
let after = row_description(&drain(&mut client).await);
assert_eq!(
after,
vec![("id".to_string(), 23), ("name".to_string(), 25)],
"view Describe after CREATE OR REPLACE via the extended protocol must reflect \
the added TEXT column (a stale shared-plan-cache entry would still report only id)"
);
}
#[tokio::test]
async fn extended_dml_returning_failure_keeps_connection_usable() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
db.execute("CREATE TABLE ov (id BIGINT PRIMARY KEY)").expect("create");
db.execute("INSERT INTO ov VALUES (9223372036854775807)")
.expect("insert");
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_parse_extended("bad".into(), "UPDATE ov SET id = id + 1 RETURNING id".into(), vec![])
.await
.expect("parse");
handler
.handle_bind_extended("bp".into(), "bad".into(), vec![], vec![], vec![])
.await
.expect("bind");
let err = handler
.handle_execute_extended("bp".into(), 0)
.await
.expect_err("overflowing UPDATE ... RETURNING must surface an error, not panic/drop");
assert_eq!(
super::handler::sqlstate_for_error(&err),
"XX000",
"a failure on the guarded DML-RETURNING path must map to a recoverable SQLSTATE; got {err}"
);
let _ = drain(&mut client).await;
handler
.handle_parse_extended("ok".into(), "SELECT id FROM ov ORDER BY id".into(), vec![])
.await
.expect("parse after error");
handler
.handle_bind_extended("op".into(), "ok".into(), vec![], vec![], vec![])
.await
.expect("bind after error");
handler
.handle_execute_extended("op".into(), 0)
.await
.expect("a fresh statement after the error must execute normally");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 1, "the row must be unchanged and the connection healthy");
assert_eq!(
rows[0][0].as_deref(),
Some(b"9223372036854775807".as_ref()),
"the overflowing UPDATE must not have mutated the row"
);
}
#[tokio::test]
#[allow(clippy::expect_used)] async fn partition_of_and_attach_detach_over_the_wire() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().expect("db"));
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_single_query("CREATE TABLE w_parent (id INT, label TEXT) PARTITION BY RANGE (id)")
.await
.expect("parent create");
handler
.handle_single_query("CREATE TABLE w_child PARTITION OF w_parent FOR VALUES FROM (0) TO (100)")
.await
.expect("child create");
let tags = command_tags(&drain(&mut client).await);
assert!(
tags.iter().any(|t| t == "CREATE TABLE"),
"PARTITION OF child must complete as CREATE TABLE, got {tags:?}"
);
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_single_query("INSERT INTO w_child (id, label) VALUES (5, 'hi')")
.await
.expect("insert");
let _ = drain(&mut client).await;
handler
.handle_single_query("SELECT id, label FROM w_child")
.await
.expect("select");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 1, "child SELECT must return the inserted row");
let (mut handler, mut client) = test_handler(Arc::clone(&db));
handler
.handle_single_query("ALTER TABLE w_parent ATTACH PARTITION w_child FOR VALUES FROM (0) TO (100)")
.await
.expect("attach");
handler
.handle_single_query("ALTER TABLE w_parent DETACH PARTITION w_child")
.await
.expect("detach");
let tags = command_tags(&drain(&mut client).await);
assert_eq!(
tags.iter().filter(|t| *t == "ALTER TABLE").count(),
2,
"ATTACH and DETACH PARTITION must each complete as ALTER TABLE, got {tags:?}"
);
}
#[tokio::test]
async fn search_path_scopes_bare_names_over_wire() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
for stmt in [
"CREATE SCHEMA wa",
"SET search_path TO wa",
"CREATE TABLE wt (v INT)",
"INSERT INTO wt (v) VALUES (42)",
] {
handler.handle_single_query(stmt).await.expect("setup stmt");
let _ = drain(&mut client).await;
}
handler.handle_single_query("SHOW search_path").await.expect("show");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("wa, public"),
"SHOW search_path must reflect the SET"
);
handler
.handle_single_query("SELECT v FROM wt")
.await
.expect("bare select");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("42"),
"bare name resolves under search_path"
);
handler
.handle_single_query("SELECT v FROM wa.wt")
.await
.expect("qualified select");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("42"),
"qualified name resolves"
);
handler
.handle_single_query("SET search_path TO public")
.await
.expect("reset to public");
let _ = drain(&mut client).await;
handler
.handle_single_query("SELECT v FROM wt")
.await
.expect("bare select public");
let out = drain(&mut client).await;
assert!(
parse_messages(&out).iter().any(|(t, _)| *t == b'E'),
"bare `wt` must error under public search_path"
);
}
#[test]
fn search_path_is_isolated_per_connection() {
std::thread::Builder::new()
.stack_size(16 * 1024 * 1024)
.spawn(|| {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("runtime")
.block_on(async move {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut a, mut ca) = test_handler(db.clone());
let (mut b, mut cb) = test_handler(db);
for stmt in ["CREATE SCHEMA sa", "CREATE SCHEMA sb"] {
a.handle_single_query(stmt).await.expect("schema");
let _ = drain(&mut ca).await;
}
a.handle_single_query("SET search_path TO sa").await.expect("A set");
let _ = drain(&mut ca).await;
b.handle_single_query("SET search_path TO sb").await.expect("B set");
let _ = drain(&mut cb).await;
a.handle_single_query("CREATE TABLE orders (v INT)")
.await
.expect("A create");
let _ = drain(&mut ca).await;
b.handle_single_query("CREATE TABLE orders (v INT)")
.await
.expect("B create");
let _ = drain(&mut cb).await;
a.handle_single_query("INSERT INTO orders VALUES (1)")
.await
.expect("A insert");
let _ = drain(&mut ca).await;
b.handle_single_query("INSERT INTO orders VALUES (2)")
.await
.expect("B insert");
let _ = drain(&mut cb).await;
a.handle_single_query("SELECT v FROM sa.orders").await.expect("read sa");
assert_eq!(
first_data_row_text(&drain(&mut ca).await).as_deref(),
Some("1"),
"A's row is in sa.orders"
);
b.handle_single_query("SELECT v FROM sb.orders").await.expect("read sb");
assert_eq!(
first_data_row_text(&drain(&mut cb).await).as_deref(),
Some("2"),
"B's row is in sb.orders"
);
a.handle_single_query("SELECT count(*) FROM sa.orders")
.await
.expect("count sa");
assert_eq!(
first_data_row_text(&drain(&mut ca).await).as_deref(),
Some("1"),
"sa.orders holds exactly one row"
);
b.handle_single_query("SELECT count(*) FROM sb.orders")
.await
.expect("count sb");
assert_eq!(
first_data_row_text(&drain(&mut cb).await).as_deref(),
Some("1"),
"sb.orders holds exactly one row"
);
a.handle_single_query("SELECT v FROM orders").await.expect("A bare");
assert_eq!(
first_data_row_text(&drain(&mut ca).await).as_deref(),
Some("1"),
"A bare orders -> sa"
);
b.handle_single_query("SELECT v FROM orders").await.expect("B bare");
assert_eq!(
first_data_row_text(&drain(&mut cb).await).as_deref(),
Some("2"),
"B bare orders -> sb"
);
a.handle_single_query("SHOW search_path").await.expect("A show");
assert_eq!(
first_data_row_text(&drain(&mut ca).await).as_deref(),
Some("sa, public")
);
b.handle_single_query("SHOW search_path").await.expect("B show");
assert_eq!(
first_data_row_text(&drain(&mut cb).await).as_deref(),
Some("sb, public")
);
});
})
.expect("spawn")
.join()
.expect("join");
}
#[tokio::test]
async fn set_constraints_defers_fk_over_wire() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(Arc::clone(&db));
for stmt in [
"CREATE SCHEMA fkpart9w",
"SET search_path TO fkpart9w",
"CREATE TABLE pk (a int PRIMARY KEY) PARTITION BY LIST (a)",
"CREATE TABLE pk1 PARTITION OF pk FOR VALUES IN (1, 2) PARTITION BY LIST (a)",
"CREATE TABLE pk11 PARTITION OF pk1 FOR VALUES IN (1)",
"CREATE TABLE pk3 PARTITION OF pk FOR VALUES IN (3)",
"CREATE TABLE fk (a int REFERENCES pk DEFERRABLE INITIALLY IMMEDIATE)",
] {
handler.handle_single_query(stmt).await.expect("setup stmt");
let _ = drain(&mut client).await;
}
handler
.handle_single_query("INSERT INTO fk VALUES (1)")
.await
.expect_err("immediate FK must reject a reference to the empty parent");
let _ = drain(&mut client).await;
for stmt in [
"BEGIN",
"SET CONSTRAINTS ALL DEFERRED",
"INSERT INTO fk VALUES (1)", "INSERT INTO pk VALUES (1)", "COMMIT", ] {
handler
.handle_single_query(stmt)
.await
.unwrap_or_else(|e| panic!("deferred flow `{stmt}` must succeed over the wire: {e}"));
let _ = drain(&mut client).await;
}
handler
.handle_single_query("SELECT a FROM fk")
.await
.expect("select fk");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("1"),
"the deferred child row must be committed"
);
handler
.handle_single_query("SELECT a FROM pk")
.await
.expect("select pk");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("1"),
"the parent row must be committed"
);
handler
.handle_single_query("INSERT INTO fk VALUES (2)")
.await
.expect_err("deferral must not leak past the transaction it was set in");
}
#[tokio::test]
async fn test_wire_bare_version_still_works() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler.handle_single_query("SELECT version()").await.unwrap();
let out = drain(&mut client).await;
let rows = data_rows(&out);
assert_eq!(rows.len(), 1, "version() must return exactly one row");
assert_eq!(rows[0].len(), 1, "version() must return exactly one column");
let text = first_data_row_text(&out).expect("version text");
assert!(
text.starts_with("PostgreSQL 16.0"),
"version() must return the PostgreSQL version banner, got {text:?}"
);
}
#[tokio::test]
async fn test_wire_bare_current_database_still_works() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler.handle_single_query("SELECT current_database()").await.unwrap();
let out = drain(&mut client).await;
let rows = data_rows(&out);
assert_eq!(rows.len(), 1, "current_database() must return exactly one row");
assert_eq!(rows[0].len(), 1, "current_database() must return exactly one column");
assert_eq!(
first_data_row_text(&out).as_deref(),
Some("heliosdb"),
"current_database() must return the database name"
);
}
#[tokio::test]
async fn test_wire_bare_current_user_still_works() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler.handle_single_query("SELECT current_user").await.unwrap();
let out = drain(&mut client).await;
let rows = data_rows(&out);
assert_eq!(rows.len(), 1, "current_user must return exactly one row");
assert_eq!(rows[0].len(), 1, "current_user must return exactly one column");
assert_eq!(
first_data_row_text(&out).as_deref(),
Some("heliosdb"),
"bare current_user keyword must return the user name via the real evaluator"
);
}
#[tokio::test]
async fn test_wire_current_database_with_operator_not_hijacked() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler
.handle_single_query("SELECT current_database() ~ 'hel'")
.await
.unwrap();
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("t"),
"current_database() ~ 'hel' must evaluate to boolean true, not return the db name"
);
}
#[tokio::test]
async fn test_wire_current_user_with_operator_not_hijacked() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler
.handle_single_query("SELECT current_user ~ 'nomatchxyz'")
.await
.unwrap();
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("f"),
"current_user ~ 'nomatchxyz' must evaluate to boolean false, not return the user name"
);
}
#[tokio::test]
async fn test_wire_version_wrapped_in_function_not_hijacked() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler
.handle_single_query("SELECT length(version()) > 0")
.await
.unwrap();
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("t"),
"length(version()) > 0 must evaluate to boolean true, not return the raw version string"
);
}
#[tokio::test]
async fn test_wire_where_clause_with_current_user_scans_real_table() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
for stmt in [
"CREATE TABLE wcu (id INT PRIMARY KEY, name TEXT)",
"INSERT INTO wcu VALUES (1, 'alice'), (2, 'bob')",
] {
handler.handle_single_query(stmt).await.expect("setup");
let _ = drain(&mut client).await;
}
handler
.handle_single_query("SELECT * FROM wcu WHERE current_user ~ 'nomatchxyz'")
.await
.expect("filtered select");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(
rows.len(),
0,
"non-matching current_user predicate must scan the table and return zero rows, not a fake row"
);
handler
.handle_single_query("SELECT id, name FROM wcu WHERE current_user ~ 'helios' ORDER BY id")
.await
.expect("matching select");
let rows = data_rows(&drain(&mut client).await);
assert_eq!(rows.len(), 2, "matching predicate must return the real table rows");
assert_eq!(rows[0][0].as_deref(), Some(b"1".as_ref()));
assert_eq!(rows[0][1].as_deref(), Some(b"alice".as_ref()));
assert_eq!(rows[1][0].as_deref(), Some(b"2".as_ref()));
assert_eq!(rows[1][1].as_deref(), Some(b"bob".as_ref()));
}
#[tokio::test]
async fn test_wire_update_with_current_database_in_where_actually_executes() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
for stmt in [
"CREATE TABLE wud (id INT PRIMARY KEY, name TEXT)",
"INSERT INTO wud VALUES (1, 'before')",
] {
handler.handle_single_query(stmt).await.expect("setup");
let _ = drain(&mut client).await;
}
handler
.handle_single_query("UPDATE wud SET name = 'after' WHERE current_database() = 'nonexistent_db_name'")
.await
.expect("update with always-false predicate");
let tags = command_tags(&drain(&mut client).await);
assert!(
tags.iter().any(|t| t == "UPDATE 0"),
"an UPDATE whose WHERE mentions current_database() must execute and report `UPDATE 0`, got {tags:?}"
);
handler
.handle_single_query("SELECT name FROM wud WHERE id = 1")
.await
.expect("verify unchanged");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("before"),
"the always-false UPDATE must not have mutated the row"
);
handler
.handle_single_query("UPDATE wud SET name = 'after' WHERE current_database() = 'heliosdb'")
.await
.expect("update with always-true predicate");
let tags = command_tags(&drain(&mut client).await);
assert!(
tags.iter().any(|t| t == "UPDATE 1"),
"the always-true UPDATE must execute and report `UPDATE 1`, got {tags:?}"
);
handler
.handle_single_query("SELECT name FROM wud WHERE id = 1")
.await
.expect("verify changed");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("after"),
"the always-true UPDATE must have genuinely applied the new value"
);
}
#[tokio::test]
async fn test_wire_multi_element_create_schema() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
handler
.handle_single_query(
"CREATE SCHEMA wms \
CREATE TABLE tbl1(f1 int PRIMARY KEY) \
CREATE TABLE tbl2(f1 int REFERENCES tbl1)",
)
.await
.expect("multi-element create schema over wire");
let tags = command_tags(&drain(&mut client).await);
assert!(
tags.iter().any(|t| t == "CREATE SCHEMA"),
"multi-element CREATE SCHEMA must complete as `CREATE SCHEMA`, got {tags:?}"
);
handler
.handle_single_query("INSERT INTO wms.tbl1 (f1) VALUES (11)")
.await
.expect("insert parent");
let _ = drain(&mut client).await;
handler
.handle_single_query("INSERT INTO wms.tbl2 (f1) VALUES (11)")
.await
.expect("insert child referencing parent");
let _ = drain(&mut client).await;
handler
.handle_single_query("SELECT f1 FROM wms.tbl2")
.await
.expect("select from newly created schema table");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("11"),
"the wire-created multi-element schema tables must be queryable and hold the inserted row"
);
let err = handler
.handle_single_query("INSERT INTO wms.tbl2 (f1) VALUES (999)")
.await
.expect_err("a dangling FK reference must be rejected");
assert!(
err.to_string().contains("wms.tbl1"),
"the FK violation must reference wms.tbl1, proving tbl2's FK resolved into the new schema, got: {err}"
);
}
#[tokio::test]
async fn test_wire_update_with_pg_tables_in_literal_actually_executes() {
let db = Arc::new(EmbeddedDatabase::new_in_memory().unwrap());
let (mut handler, mut client) = test_handler(db);
for stmt in [
"CREATE TABLE inventory (id INT PRIMARY KEY, note TEXT)",
"INSERT INTO inventory VALUES (1, 'before')",
] {
handler.handle_single_query(stmt).await.expect("setup");
let _ = drain(&mut client).await;
}
handler
.handle_single_query("UPDATE inventory SET note='see pg_tables' WHERE id=1")
.await
.expect("update with pg_tables in literal must execute");
let tags = command_tags(&drain(&mut client).await);
assert!(
tags.iter().any(|t| t == "UPDATE 1"),
"an UPDATE whose literal mentions pg_tables must execute and report `UPDATE 1`, got {tags:?}"
);
handler
.handle_single_query("SELECT note FROM inventory WHERE id = 1")
.await
.expect("read back");
assert_eq!(
first_data_row_text(&drain(&mut client).await).as_deref(),
Some("see pg_tables"),
"the UPDATE must have genuinely written the new note value"
);
}