use crate::common;
use std::io::{Read, Write};
use std::net::TcpStream;
use std::path::PathBuf;
use std::time::Duration;
const READ_TIMEOUT: Duration = Duration::from_secs(5);
fn unique_tmpdir(label: &str) -> PathBuf {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let pid = std::process::id();
let p = std::env::temp_dir().join(format!("spg-e2e-mysqlwire-query-{label}-{pid}-{nanos}"));
std::fs::create_dir_all(&p).unwrap();
p
}
fn read_packet(stream: &mut TcpStream) -> (u8, Vec<u8>) {
let mut hdr = [0u8; 4];
stream.read_exact(&mut hdr).expect("read header");
let len = u32::from(hdr[0]) | (u32::from(hdr[1]) << 8) | (u32::from(hdr[2]) << 16);
let seqno = hdr[3];
let mut payload = vec![0u8; len as usize];
stream.read_exact(&mut payload).expect("read payload");
(seqno, payload)
}
fn write_query(stream: &mut TcpStream, sql: &str) -> std::io::Result<()> {
let mut payload = Vec::with_capacity(1 + sql.len());
payload.push(0x03); payload.extend_from_slice(sql.as_bytes());
let len = u32::try_from(payload.len()).expect("query fits a packet");
let hdr = [len as u8, (len >> 8) as u8, (len >> 16) as u8, 0u8];
stream.write_all(&hdr)?;
stream.write_all(&payload)?;
Ok(())
}
fn write_packet(stream: &mut TcpStream, seqno: u8, payload: &[u8]) {
let len = payload.len() as u32;
let hdr = [len as u8, (len >> 8) as u8, (len >> 16) as u8, seqno];
stream.write_all(&hdr).expect("write hdr");
stream.write_all(payload).expect("write payload");
}
fn build_handshake_response(username: &str) -> Vec<u8> {
let caps: u32 = 0x0000_0200 | 0x0000_8000 | 0x0008_0000;
let mut payload = Vec::new();
payload.extend_from_slice(&caps.to_le_bytes());
payload.extend_from_slice(&16_777_215u32.to_le_bytes());
payload.push(0xff);
payload.extend_from_slice(&[0u8; 23]);
payload.extend_from_slice(username.as_bytes());
payload.push(0);
payload.push(0); payload.extend_from_slice(b"mysql_native_password\0");
payload
}
fn auth_open_mode(addr: &str) -> TcpStream {
let mut s = common::connect_to(addr);
s.set_read_timeout(Some(READ_TIMEOUT)).unwrap();
let (_seqno, _greeting) = read_packet(&mut s);
write_packet(&mut s, 1, &build_handshake_response("anyone"));
let (_seqno, ok) = read_packet(&mut s);
assert_eq!(ok[0], 0x00, "expected OK after auth, got {:#x}", ok[0]);
s
}
fn send_query(s: &mut TcpStream, sql: &str) {
let mut payload = Vec::with_capacity(1 + sql.len());
payload.push(0x03); payload.extend_from_slice(sql.as_bytes());
write_packet(s, 0, &payload);
}
fn spawn() -> (common::ChildGuard, String) {
let dir = unique_tmpdir("svc");
let db = dir.join("spg.db");
let (child, addrs) = common::ServerBuilder::new()
.arg_path(&db)
.with_mysqlwire()
.spawn();
let addr = addrs.mysqlwire.expect("mysql-wire addr");
(common::ChildGuard(child), addr)
}
fn read_lenenc(buf: &[u8], pos: usize) -> (u64, usize) {
let first = buf[pos];
match first {
0xfb => (0, 1), 0xfc => {
let v = u16::from_le_bytes(buf[pos + 1..pos + 3].try_into().unwrap());
(u64::from(v), 3)
}
0xfd => {
let mut bytes = [0u8; 4];
bytes[..3].copy_from_slice(&buf[pos + 1..pos + 4]);
let v = u32::from_le_bytes(bytes);
(u64::from(v), 4)
}
0xfe => {
let v = u64::from_le_bytes(buf[pos + 1..pos + 9].try_into().unwrap());
(v, 9)
}
n => (u64::from(n), 1),
}
}
fn read_lenenc_string(buf: &[u8], pos: usize) -> (Vec<u8>, usize) {
let (n, consumed) = read_lenenc(buf, pos);
let s = buf[pos + consumed..pos + consumed + n as usize].to_vec();
(s, consumed + n as usize)
}
fn read_columns_eof(s: &mut TcpStream) {
let (_seq, pkt) = read_packet(s);
assert_eq!(pkt[0], 0xfe, "EOF closes the column definitions");
assert_eq!(pkt.len(), 5, "protocol-41 EOF: header + warnings + status");
}
fn read_result_eof(s: &mut TcpStream) {
read_columns_eof(s);
}
fn is_result_eof(pkt: &[u8]) -> bool {
pkt[0] == 0xfe && pkt.len() < 9
}
#[test]
fn select_literal_int_returns_one_column_one_row() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "SELECT 42 AS answer");
let (_seq, cc_pkt) = read_packet(&mut s);
let (col_count, _) = read_lenenc(&cc_pkt, 0);
assert_eq!(col_count, 1, "1 projection column");
let (_seq, col_def) = read_packet(&mut s);
let mut pos = 0;
let (_catalog, c) = read_lenenc_string(&col_def, pos);
pos += c;
let (_schema, c) = read_lenenc_string(&col_def, pos);
pos += c;
let (_table, c) = read_lenenc_string(&col_def, pos);
pos += c;
let (_org_table, c) = read_lenenc_string(&col_def, pos);
pos += c;
let (name, c) = read_lenenc_string(&col_def, pos);
pos += c;
assert_eq!(name, b"answer");
let (_org_name, c) = read_lenenc_string(&col_def, pos);
pos += c;
assert_eq!(col_def[pos], 0x0c, "fixed-length marker");
let type_byte = col_def[pos + 1 + 2 + 4];
assert_eq!(type_byte, 0x03, "MYSQL_TYPE_LONG");
read_columns_eof(&mut s);
let (_seq, row) = read_packet(&mut s);
let (value, _) = read_lenenc_string(&row, 0);
assert_eq!(value, b"42");
let (_seq, eof) = read_packet(&mut s);
assert_eq!(eof[0], 0xfe, "trailing EOF");
}
#[test]
fn select_text_literal_round_trips_through_lenenc_string() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "SELECT 'hello, mysql' AS greeting");
let (_seq, _cc) = read_packet(&mut s);
let (_seq, _col) = read_packet(&mut s);
read_columns_eof(&mut s);
let (_seq, row) = read_packet(&mut s);
let (value, _) = read_lenenc_string(&row, 0);
assert_eq!(value, b"hello, mysql");
read_result_eof(&mut s);
}
#[test]
fn ddl_returns_ok_packet_with_affected_zero() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "CREATE TABLE t (id INT NOT NULL)");
let (_seq, ok) = read_packet(&mut s);
assert_eq!(ok[0], 0x00, "OK header");
assert_eq!(ok[1], 0, "affected = 0");
}
#[test]
fn dml_returns_ok_packet_with_affected_count() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "CREATE TABLE t (id INT NOT NULL)");
let (_seq, _ok) = read_packet(&mut s);
send_query(&mut s, "INSERT INTO t VALUES (1), (2), (3)");
let (_seq, ok) = read_packet(&mut s);
assert_eq!(ok[0], 0x00);
let (affected, _) = read_lenenc(&ok, 1);
assert_eq!(affected, 3);
}
#[test]
fn select_from_table_returns_correct_rows() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "CREATE TABLE products (id INT NOT NULL, name TEXT)");
let (_seq, _ok) = read_packet(&mut s);
send_query(
&mut s,
"INSERT INTO products VALUES (10, 'widget'), (20, 'gadget'), (30, 'doohickey')",
);
let (_seq, _ok) = read_packet(&mut s);
send_query(&mut s, "SELECT id, name FROM products ORDER BY id");
let (_seq, cc) = read_packet(&mut s);
let (col_count, _) = read_lenenc(&cc, 0);
assert_eq!(col_count, 2);
let (_seq, _col1) = read_packet(&mut s);
let (_seq, _col2) = read_packet(&mut s);
read_columns_eof(&mut s);
let mut got_rows: Vec<(String, String)> = Vec::new();
loop {
let (_seq, pkt) = read_packet(&mut s);
if is_result_eof(&pkt) {
break;
}
let (id_bytes, c) = read_lenenc_string(&pkt, 0);
let (name_bytes, _) = read_lenenc_string(&pkt, c);
got_rows.push((
String::from_utf8(id_bytes).unwrap(),
String::from_utf8(name_bytes).unwrap(),
));
}
assert_eq!(
got_rows,
vec![
("10".to_string(), "widget".to_string()),
("20".to_string(), "gadget".to_string()),
("30".to_string(), "doohickey".to_string()),
]
);
}
#[test]
fn null_values_decode_as_0xfb_byte() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "CREATE TABLE t (a INT, b TEXT)");
let (_seq, _ok) = read_packet(&mut s);
send_query(&mut s, "INSERT INTO t VALUES (NULL, 'x'), (1, NULL)");
let (_seq, _ok) = read_packet(&mut s);
send_query(&mut s, "SELECT a, b FROM t ORDER BY b");
let (_seq, _cc) = read_packet(&mut s);
let (_seq, _c1) = read_packet(&mut s);
let (_seq, _c2) = read_packet(&mut s);
read_columns_eof(&mut s);
let (_seq, row1) = read_packet(&mut s);
let (a1, c) = read_lenenc_string(&row1, 0);
assert_eq!(a1, b"1");
assert_eq!(row1[c], 0xfb, "NULL text column");
let (_seq, row2) = read_packet(&mut s);
assert_eq!(row2[0], 0xfb, "NULL int column");
let (b2, _) = read_lenenc_string(&row2, 1);
assert_eq!(b2, b"x");
read_result_eof(&mut s);
}
#[test]
fn parse_error_returns_err_packet() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
send_query(&mut s, "SELEKT 1");
let (_seq, err) = read_packet(&mut s);
assert_eq!(err[0], 0xff);
let errno = u16::from_le_bytes(err[1..3].try_into().unwrap());
assert_eq!(errno, 1064, "ER_PARSE_ERROR");
assert_eq!(&err[4..9], b"42000");
}
#[test]
fn com_quit_closes_connection_cleanly() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
write_packet(&mut s, 0, &[0x01]);
let mut buf = [0u8; 4];
let n = s.read(&mut buf).unwrap_or(0);
assert_eq!(n, 0, "server closed connection after COM_QUIT");
}
fn query_scalar(s: &mut TcpStream, sql: &str) -> String {
send_query(s, sql);
let (_seq, cc) = read_packet(s);
let (col_count, _) = read_lenenc(&cc, 0);
for _ in 0..col_count {
let _ = read_packet(s);
}
read_columns_eof(s);
let (_seq, row) = read_packet(s);
let (val, _) = read_lenenc_string(&row, 0);
read_result_eof(s);
String::from_utf8(val).unwrap()
}
fn exec_ok(s: &mut TcpStream, sql: &str) {
send_query(s, sql);
let (_seq, ok) = read_packet(s);
assert_eq!(ok[0], 0x00, "expected OK for `{sql}`, got {:#x}", ok[0]);
}
#[test]
fn mysql_connection_defaults_to_backslash_escape_dialect() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\n')"), "1");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\t')"), "1");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\\')"), "1");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\'')"), "1");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('a\nb')"), "3");
}
#[test]
fn mysql_no_backslash_escapes_sql_mode_disables_escapes() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\n')"), "1");
exec_ok(&mut s, "SET sql_mode='NO_BACKSLASH_ESCAPES'");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\n')"), "2");
exec_ok(&mut s, "SET sql_mode='STRICT_TRANS_TABLES'");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\n')"), "1");
exec_ok(&mut s, "SET sql_mode='ANSI_QUOTES,NO_BACKSLASH_ESCAPES'");
assert_eq!(query_scalar(&mut s, r"SELECT LENGTH('\n')"), "2");
}
#[test]
fn mysql_dialect_is_isolated_per_connection() {
let (_guard, addr) = spawn();
let mut a = auth_open_mode(&addr);
exec_ok(&mut a, "SET sql_mode='NO_BACKSLASH_ESCAPES'");
assert_eq!(query_scalar(&mut a, r"SELECT LENGTH('\n')"), "2");
let mut b = auth_open_mode(&addr);
assert_eq!(query_scalar(&mut b, r"SELECT LENGTH('\n')"), "1");
assert_eq!(query_scalar(&mut a, r"SELECT LENGTH('\n')"), "2");
}
#[test]
fn two_mysql_connections_can_each_hold_a_transaction() {
let (_guard, addr) = spawn();
let mut a = auth_open_mode(&addr);
let mut b = auth_open_mode(&addr);
exec_ok(&mut a, "CREATE TABLE t (id INT PRIMARY KEY, v INT)");
exec_ok(&mut a, "BEGIN");
exec_ok(&mut b, "BEGIN");
exec_ok(&mut a, "INSERT INTO t VALUES (1, 100)");
exec_ok(&mut b, "INSERT INTO t VALUES (2, 200)");
assert_eq!(query_scalar(&mut a, "SELECT COUNT(*) FROM t"), "1");
assert_eq!(query_scalar(&mut b, "SELECT COUNT(*) FROM t"), "1");
exec_ok(&mut a, "COMMIT");
exec_ok(&mut b, "COMMIT");
assert_eq!(query_scalar(&mut a, "SELECT COUNT(*) FROM t"), "2");
}
#[test]
fn mysql_transaction_rollback_discards_only_its_own_writes() {
let (_guard, addr) = spawn();
let mut a = auth_open_mode(&addr);
let mut b = auth_open_mode(&addr);
exec_ok(&mut a, "CREATE TABLE t (id INT PRIMARY KEY)");
exec_ok(&mut a, "INSERT INTO t VALUES (1)");
exec_ok(&mut b, "BEGIN");
exec_ok(&mut b, "INSERT INTO t VALUES (2)");
exec_ok(&mut b, "ROLLBACK");
assert_eq!(query_scalar(&mut a, "SELECT COUNT(*) FROM t"), "1");
}
#[test]
fn mysql_open_transaction_rolls_back_on_disconnect() {
let (_guard, addr) = spawn();
let mut a = auth_open_mode(&addr);
exec_ok(&mut a, "CREATE TABLE t (id INT PRIMARY KEY)");
exec_ok(&mut a, "BEGIN");
exec_ok(&mut a, "INSERT INTO t VALUES (1)");
write_packet(&mut a, 0, &[0x01]); drop(a);
let mut b = auth_open_mode(&addr);
assert_eq!(query_scalar(&mut b, "SELECT COUNT(*) FROM t"), "0");
}
#[test]
fn unknown_command_returns_err_packet() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
write_packet(&mut s, 0, &[0x99, b'X']);
let (_seq, err) = read_packet(&mut s);
assert_eq!(err[0], 0xff);
let errno = u16::from_le_bytes(err[1..3].try_into().unwrap());
assert_eq!(errno, 1047);
}
fn ok_status(pkt: &[u8]) -> u16 {
assert_eq!(pkt[0], 0x00, "expected an OK packet, got {:#x}", pkt[0]);
let mut pos = 1;
for _ in 0..2 {
let (_, used) = read_lenenc(pkt, pos);
pos += used;
}
u16::from_le_bytes([pkt[pos], pkt[pos + 1]])
}
fn status_of(s: &mut TcpStream, sql: &str) -> u16 {
send_query(s, sql);
let (_seq, pkt) = read_packet(s);
ok_status(&pkt)
}
fn eof_status(pkt: &[u8]) -> u16 {
assert_eq!(pkt[0], 0xfe, "expected an EOF packet, got {:#x}", pkt[0]);
u16::from_le_bytes([pkt[3], pkt[4]])
}
fn status_of_select(s: &mut TcpStream, sql: &str) -> u16 {
send_query(s, sql);
let (_seq, cc) = read_packet(s);
let (col_count, _) = read_lenenc(&cc, 0);
for _ in 0..col_count {
let _ = read_packet(s);
}
read_columns_eof(s);
loop {
let (_seq, pkt) = read_packet(s);
if is_result_eof(&pkt) {
return eof_status(&pkt);
}
}
}
const IN_TRANS: u16 = 0x0001;
const AUTOCOMMIT: u16 = 0x0002;
#[test]
fn the_ok_packet_reports_the_transaction_state() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
exec_ok(&mut s, "CREATE TABLE t316 (id INT)");
assert_eq!(status_of_select(&mut s, "SELECT 1"), AUTOCOMMIT, "idle");
assert_eq!(
status_of(&mut s, "BEGIN"),
AUTOCOMMIT | IN_TRANS,
"BEGIN's own reply"
);
assert_eq!(
status_of(&mut s, "INSERT INTO t316 VALUES (1)"),
AUTOCOMMIT | IN_TRANS,
"inside the block"
);
assert_eq!(status_of(&mut s, "COMMIT"), AUTOCOMMIT, "COMMIT clears it");
assert_eq!(status_of(&mut s, "BEGIN"), AUTOCOMMIT | IN_TRANS);
assert_eq!(
status_of(&mut s, "ROLLBACK"),
AUTOCOMMIT,
"ROLLBACK clears it"
);
assert_eq!(
status_of_select(&mut s, "SELECT 1"),
AUTOCOMMIT,
"idle again"
);
}
#[test]
fn a_result_sets_terminator_reports_it_too() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
assert_eq!(status_of_select(&mut s, "SELECT 1"), AUTOCOMMIT);
exec_ok(&mut s, "BEGIN");
assert_eq!(
status_of_select(&mut s, "SELECT 1"),
AUTOCOMMIT | IN_TRANS,
"SELECT inside a block"
);
exec_ok(&mut s, "COMMIT");
assert_eq!(status_of_select(&mut s, "SELECT 1"), AUTOCOMMIT);
}
#[test]
fn com_ping_reports_the_transaction_state() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
write_packet(&mut s, 0, &[0x0e]);
let (_seq, pkt) = read_packet(&mut s);
assert_eq!(ok_status(&pkt), AUTOCOMMIT, "ping when idle");
exec_ok(&mut s, "BEGIN");
write_packet(&mut s, 0, &[0x0e]);
let (_seq, pkt) = read_packet(&mut s);
assert_eq!(
ok_status(&pkt),
AUTOCOMMIT | IN_TRANS,
"ping inside a block"
);
exec_ok(&mut s, "ROLLBACK");
}
#[test]
fn one_connections_block_does_not_colour_anothers_status() {
let (_guard, addr) = spawn();
let mut a = auth_open_mode(&addr);
let mut b = auth_open_mode(&addr);
exec_ok(&mut a, "BEGIN");
assert_eq!(status_of_select(&mut a, "SELECT 1"), AUTOCOMMIT | IN_TRANS);
assert_eq!(
status_of_select(&mut b, "SELECT 1"),
AUTOCOMMIT,
"B is idle and must say so while A holds a block"
);
exec_ok(&mut a, "ROLLBACK");
}
fn auth_open_mode_with_id(addr: &str) -> (TcpStream, u32) {
let mut s = common::connect_to(addr);
s.set_read_timeout(Some(READ_TIMEOUT)).unwrap();
let (_seqno, greeting) = read_packet(&mut s);
let nul = 1 + greeting[1..]
.iter()
.position(|&b| b == 0)
.expect("version NUL");
let idpos = nul + 1;
let conn_id = u32::from_le_bytes(greeting[idpos..idpos + 4].try_into().unwrap());
write_packet(&mut s, 1, &build_handshake_response("anyone"));
let (_seqno, ok) = read_packet(&mut s);
assert_eq!(ok[0], 0x00, "expected OK after auth, got {:#x}", ok[0]);
(s, conn_id)
}
fn query_rows(s: &mut TcpStream, sql: &str) -> Vec<Vec<Option<String>>> {
send_query(s, sql);
let (_seq, cc) = read_packet(s);
let (col_count, _) = read_lenenc(&cc, 0);
for _ in 0..col_count {
let _ = read_packet(s);
}
read_columns_eof(s);
let mut out = Vec::new();
loop {
let (_seq, pkt) = read_packet(s);
if is_result_eof(&pkt) {
return out;
}
let mut pos = 0;
let mut row = Vec::with_capacity(col_count as usize);
for _ in 0..col_count {
if pkt[pos] == 0xfb {
row.push(None);
pos += 1;
} else {
let (v, used) = read_lenenc_string(&pkt, pos);
pos += used;
row.push(Some(String::from_utf8(v).unwrap()));
}
}
out.push(row);
}
}
#[test]
fn connection_id_is_per_connection_and_matches_the_greeting() {
let (_guard, addr) = spawn();
let (mut a, a_id) = auth_open_mode_with_id(&addr);
let (mut b, b_id) = auth_open_mode_with_id(&addr);
assert_ne!(a_id, b_id, "two live connections must not share an id");
let a_seen = query_scalar(&mut a, "SELECT connection_id()");
let b_seen = query_scalar(&mut b, "SELECT connection_id()");
assert_eq!(
a_seen,
a_id.to_string(),
"CONNECTION_ID() must be the id the greeting announced"
);
assert_eq!(b_seen, b_id.to_string());
assert_eq!(
query_scalar(&mut a, "SELECT connection_id()"),
a_seen,
"stable within the connection"
);
}
#[test]
fn show_processlist_lists_the_live_connections() {
let (_guard, addr) = spawn();
let (mut a, a_id) = auth_open_mode_with_id(&addr);
let (mut b, b_id) = auth_open_mode_with_id(&addr);
assert_eq!(query_scalar(&mut b, "SELECT 1"), "1");
let rows = query_rows(&mut a, "SHOW PROCESSLIST");
let find = |id: u32| {
rows.iter()
.find(|r| r[0].as_deref() == Some(id.to_string().as_str()))
.unwrap_or_else(|| panic!("no row for connection {id} in {rows:?}"))
.clone()
};
let a_row = find(a_id);
let b_row = find(b_id);
assert_eq!(
a_row[4].as_deref(),
Some("Query"),
"the asker is running one"
);
assert_eq!(
a_row[7].as_deref(),
Some("SHOW PROCESSLIST"),
"its own Info is the statement it is running"
);
assert_eq!(
b_row[4].as_deref(),
Some("Sleep"),
"B is between statements"
);
assert_eq!(b_row[7], None, "an idle connection has no Info");
}
#[test]
fn a_closed_connection_drops_out_of_the_processlist() {
let (_guard, addr) = spawn();
let (mut a, _a_id) = auth_open_mode_with_id(&addr);
let (mut b, b_id) = auth_open_mode_with_id(&addr);
assert_eq!(query_scalar(&mut b, "SELECT 1"), "1");
assert!(
query_rows(&mut a, "SHOW PROCESSLIST")
.iter()
.any(|r| r[0].as_deref() == Some(b_id.to_string().as_str())),
"B is live"
);
write_packet(&mut b, 0, &[0x01]);
drop(b);
let mut gone = false;
for _ in 0..100 {
gone = !query_rows(&mut a, "SHOW PROCESSLIST")
.iter()
.any(|r| r[0].as_deref() == Some(b_id.to_string().as_str()));
if gone {
break;
}
std::thread::sleep(Duration::from_millis(20));
}
assert!(gone, "a disconnected connection must leave the processlist");
}
fn err_parts(pkt: &[u8]) -> (u16, String, String) {
assert_eq!(pkt[0], 0xff, "expected an ERR packet, got {:#x}", pkt[0]);
let errno = u16::from_le_bytes([pkt[1], pkt[2]]);
assert_eq!(pkt[3], b'#', "SQLSTATE marker");
let sqlstate = String::from_utf8(pkt[4..9].to_vec()).unwrap();
let msg = String::from_utf8(pkt[9..].to_vec()).unwrap();
(errno, sqlstate, msg)
}
fn err_of(s: &mut TcpStream, sql: &str) -> (u16, String, String) {
send_query(s, sql);
let (_seq, pkt) = read_packet(s);
err_parts(&pkt)
}
#[test]
fn kill_of_an_unknown_thread_id_is_error_1094() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
for sql in ["KILL 999999", "KILL QUERY 999999", "KILL CONNECTION 999999"] {
let (errno, sqlstate, msg) = err_of(&mut s, sql);
assert_eq!((errno, sqlstate.as_str()), (1094, "HY000"), "for `{sql}`");
assert_eq!(msg, "Unknown thread id: 999999", "for `{sql}`");
}
assert_eq!(query_scalar(&mut s, "SELECT 1"), "1");
}
#[test]
fn kill_connection_drops_the_named_connection() {
let (_guard, addr) = spawn();
let (mut killer, _) = auth_open_mode_with_id(&addr);
let (mut victim, victim_id) = auth_open_mode_with_id(&addr);
assert_eq!(query_scalar(&mut victim, "SELECT 1"), "1");
exec_ok(&mut killer, &format!("KILL CONNECTION {victim_id}"));
let mut gone = false;
for _ in 0..100 {
gone = !query_rows(&mut killer, "SHOW PROCESSLIST")
.iter()
.any(|r| r[0].as_deref() == Some(victim_id.to_string().as_str()));
if gone {
break;
}
std::thread::sleep(Duration::from_millis(20));
}
assert!(gone, "a killed connection must leave the processlist");
let (errno, _, _) = err_of(&mut killer, &format!("KILL CONNECTION {victim_id}"));
assert_eq!(errno, 1094, "the id is no longer live");
let mut hdr = [0u8; 4];
let exchange_failed =
write_query(&mut victim, "SELECT 1").is_err() || victim.read_exact(&mut hdr).is_err();
assert!(
exchange_failed,
"the killed connection must be closed, on write or on read"
);
assert_eq!(query_scalar(&mut killer, "SELECT 1"), "1");
}
#[test]
fn kill_of_your_own_connection_reports_1927_and_closes() {
let (_guard, addr) = spawn();
let (mut s, my_id) = auth_open_mode_with_id(&addr);
let (errno, sqlstate, msg) = err_of(&mut s, &format!("KILL CONNECTION {my_id}"));
assert_eq!((errno, sqlstate.as_str()), (1927, "70100"));
assert_eq!(msg, "Connection was killed");
{
use std::io::Write;
let payload = [0x03, b'S', b'E', b'L', b'E', b'C', b'T', b' ', b'1'];
let hdr = [payload.len() as u8, 0, 0, 0];
let wrote = s.write_all(&hdr).and_then(|()| s.write_all(&payload));
if wrote.is_ok() {
let mut hdr = [0u8; 4];
assert!(
s.read_exact(&mut hdr).is_err(),
"the connection must be closed after killing itself"
);
}
}
}
#[test]
fn kill_query_leaves_the_connection_alive() {
let (_guard, addr) = spawn();
let (mut s, my_id) = auth_open_mode_with_id(&addr);
exec_ok(&mut s, &format!("KILL QUERY {my_id}"));
assert_eq!(
query_scalar(&mut s, "SELECT 1"),
"1",
"KILL QUERY must not end the connection"
);
}
#[test]
fn processlist_host_and_db_describe_the_connection() {
let (_guard, addr) = spawn();
let (mut s, my_id) = auth_open_mode_with_id(&addr);
let local = s.local_addr().unwrap();
let own_row = |s: &mut TcpStream| {
query_rows(s, "SHOW PROCESSLIST")
.into_iter()
.find(|r| r[0].as_deref() == Some(my_id.to_string().as_str()))
.expect("our own row")
};
let row = own_row(&mut s);
assert_eq!(
row[2].as_deref(),
Some(format!("{}:{}", local.ip(), local.port()).as_str()),
"Host is the peer address"
);
assert_eq!(row[3], None, "no database selected yet");
let mut pkt = vec![0x02];
pkt.extend_from_slice(b"shop");
write_packet(&mut s, 0, &pkt);
let (_seq, ok) = read_packet(&mut s);
assert_eq!(ok[0], 0x00, "COM_INIT_DB accepted");
assert_eq!(
own_row(&mut s)[3].as_deref(),
Some("shop"),
"db follows the selected database"
);
}
#[test]
fn a_parse_error_carries_no_internal_prefix() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
let (errno, sqlstate, msg) = err_of(&mut s, "SELECT * FROM");
assert_eq!((errno, sqlstate.as_str()), (1064, "42000"));
assert!(
!msg.contains("parse error at token"),
"SPG's token index leaked: {msg}"
);
for prefix in ["parse: ", "eval: ", "unsupported: ", "storage: ", "lex: "] {
assert!(
!msg.starts_with(prefix),
"SPG's internal class vocabulary leaked: {msg}"
);
}
assert_eq!(query_scalar(&mut s, "SELECT 1"), "1");
}
#[test]
fn a_runtime_error_carries_no_internal_prefix() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
let (_errno, _sqlstate, msg) = err_of(&mut s, "SELECT * FROM no_such_table");
for prefix in ["parse: ", "eval: ", "unsupported: ", "storage: ", "lex: "] {
assert!(
!msg.starts_with(prefix),
"SPG's internal class vocabulary leaked: {msg}"
);
}
}
#[test]
fn set_autocommit_off_defers_the_commit() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
exec_ok(&mut s, "CREATE TABLE ac (id INT NOT NULL)");
exec_ok(&mut s, "SET autocommit=0");
exec_ok(&mut s, "INSERT INTO ac VALUES (1)");
assert_eq!(
query_scalar(&mut s, "SELECT COUNT(*) FROM ac"),
"1",
"the session sees its own uncommitted write"
);
exec_ok(&mut s, "ROLLBACK");
assert_eq!(
query_scalar(&mut s, "SELECT COUNT(*) FROM ac"),
"0",
"ROLLBACK discards it — with autocommit on it would already be permanent"
);
exec_ok(&mut s, "INSERT INTO ac VALUES (2)");
exec_ok(&mut s, "COMMIT");
assert_eq!(query_scalar(&mut s, "SELECT COUNT(*) FROM ac"), "1");
}
#[test]
fn at_at_autocommit_reports_the_setting() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
assert_eq!(query_scalar(&mut s, "SELECT @@autocommit"), "1");
exec_ok(&mut s, "SET autocommit=0");
assert_eq!(query_scalar(&mut s, "SELECT @@autocommit"), "0");
exec_ok(&mut s, "SET autocommit=1");
assert_eq!(query_scalar(&mut s, "SELECT @@autocommit"), "1");
}
#[test]
fn the_status_flags_drop_autocommit_when_it_is_off() {
let (_guard, addr) = spawn();
let mut s = auth_open_mode(&addr);
assert_eq!(status_of_select(&mut s, "SELECT 1"), AUTOCOMMIT);
assert_eq!(
status_of(&mut s, "SET autocommit=0"),
0,
"the bit is cleared"
);
exec_ok(&mut s, "CREATE TABLE ac2 (id INT NOT NULL)");
assert_eq!(status_of(&mut s, "INSERT INTO ac2 VALUES (1)"), IN_TRANS);
exec_ok(&mut s, "ROLLBACK");
assert_eq!(status_of(&mut s, "SET autocommit=1"), AUTOCOMMIT);
}
#[test]
fn a_disconnect_under_autocommit_off_rolls_back() {
let (_guard, addr) = spawn();
let mut a = auth_open_mode(&addr);
exec_ok(&mut a, "CREATE TABLE ac3 (id INT NOT NULL)");
{
let mut b = auth_open_mode(&addr);
exec_ok(&mut b, "SET autocommit=0");
exec_ok(&mut b, "INSERT INTO ac3 VALUES (1)");
write_packet(&mut b, 0, &[0x01]); }
for _ in 0..100 {
if query_scalar(&mut a, "SELECT COUNT(*) FROM ac3") == "0" {
return;
}
std::thread::sleep(Duration::from_millis(20));
}
panic!("an uncommitted write survived the disconnect");
}