use std::env;
use std::path::PathBuf;
use std::sync::Arc;
use sqlx::sqlite::{SqliteConnectOptions, SqlitePoolOptions};
use sqlx::{AssertSqlSafe, Executor};
use uuid::Uuid;
use super::query::Cell;
use super::schema::ReferentialAction;
use super::{
CatalogKind, Connection, ConnectionConfig, Credentials, DatabaseObject, Engine, ObjectKind,
QueryDigest, RowKey, SafetyMode, SshAuth, SshConfig, SslConfig, SslMode, TagColor,
keyword_literal, quote_identifier, typed_placeholder,
};
use super::{
Decision, OnFailure, PinnedConnection, QueryOutcome, QuerySource, ScriptFailure, ScriptMode,
ScriptOutcome, ScriptRun, Step, TxnState,
};
pub(crate) struct TempDatabase {
path: PathBuf,
}
impl TempDatabase {
pub(crate) async fn new() -> Self {
let path = env::temp_dir().join(format!("zippa-db-test-{}.sqlite", Uuid::new_v4()));
let pool = SqlitePoolOptions::new()
.connect_with(
SqliteConnectOptions::new()
.filename(&path)
.create_if_missing(true),
)
.await
.expect("could not create the test database");
pool.execute(AssertSqlSafe(
"CREATE TABLE items (
id INTEGER PRIMARY KEY,
name TEXT,
score REAL,
payload BLOB
);
INSERT INTO items VALUES (1, 'alpha', 1.5, x'001122');
INSERT INTO items VALUES (2, NULL, NULL, NULL);
CREATE VIEW named_items AS SELECT id, name FROM items WHERE name IS NOT NULL;"
.to_string(),
))
.await
.expect("could not seed the test database");
pool.close().await;
Self { path }
}
pub(crate) fn config(&self) -> ConnectionConfig {
ConnectionConfig {
database: self.path.to_string_lossy().to_string(),
..ConnectionConfig::new(Engine::Sqlite)
}
}
}
impl Drop for TempDatabase {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.path);
}
}
#[tokio::test]
async fn reads_columns_values_and_nulls() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let result = connection
.run_query("SELECT id, name, score, payload FROM items ORDER BY id")
.await
.expect("query failed");
assert_eq!(result.columns, ["id", "name", "score", "payload"]);
assert_eq!(
result.rows[0],
[
Some("1".to_string()),
Some("alpha".to_string()),
Some("1.5".to_string()),
Some("<3 bytes>".to_string()),
]
);
assert_eq!(result.rows[1][1], None);
assert_eq!(result.rows[1][2], None);
assert_eq!(result.row_count(), 2);
connection.close().await;
}
#[tokio::test]
async fn fetch_binary_reads_the_real_bytes_a_query_only_describes() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let described = connection
.run_query("SELECT payload FROM items WHERE id = 1")
.await
.expect("query failed");
assert_eq!(described.rows[0][0], Some("<3 bytes>".to_string()));
let bytes = connection
.fetch_binary(
"SELECT payload FROM items WHERE id = ?",
vec![Some("1".to_string())],
)
.await
.expect("fetch failed");
assert_eq!(bytes, Some(vec![0x00, 0x11, 0x22]));
let null = connection
.fetch_binary(
"SELECT payload FROM items WHERE id = ?",
vec![Some("2".to_string())],
)
.await
.expect("fetch failed");
assert_eq!(null, None);
let missing = connection
.fetch_binary(
"SELECT payload FROM items WHERE id = ?",
vec![Some("99".to_string())],
)
.await
.expect("fetch failed");
assert_eq!(missing, None);
connection.close().await;
}
#[tokio::test]
async fn sqlite_has_no_process_list() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let error = connection
.processes()
.await
.expect_err("SQLite has no server to list processes for");
assert!(error.to_string().contains("no server processes"));
let error = connection
.kill_process("1")
.await
.expect_err("SQLite has no server process to end");
assert!(error.to_string().contains("no server processes"));
connection.close().await;
}
#[tokio::test]
async fn sqlite_has_no_server_variables() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let error = connection
.server_variables()
.await
.expect_err("SQLite has no server-side configuration to list");
assert!(error.to_string().contains("no server variables"));
connection.close().await;
}
#[tokio::test]
async fn sqlite_has_no_query_digest() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let error = connection
.query_digest()
.await
.expect_err("SQLite has no query instrumentation to read");
assert!(error.to_string().contains("no query digest"));
connection.close().await;
}
#[tokio::test]
async fn sqlite_maintenance_runs_each_task() {
use super::Maintenance;
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let check = connection
.run_maintenance(Maintenance::IntegrityCheck)
.await
.expect("integrity check should run");
assert_eq!(check.rows, vec![vec![Some("ok".to_string())]]);
for task in Maintenance::ALL {
connection
.run_maintenance(task)
.await
.unwrap_or_else(|error| panic!("{task:?} failed: {error:#}"));
}
connection.close().await;
}
#[tokio::test]
async fn read_only_sqlite_refuses_the_maintenance_that_writes() {
use super::Maintenance;
let database = TempDatabase::new().await;
let mut config = database.config();
config.safety = SafetyMode::ReadOnly;
let connection = Connection::open(config, None)
.await
.expect("could not open the test database");
for task in Maintenance::ALL {
let outcome = connection.run_maintenance(task).await;
if task.writes() {
let error = outcome.expect_err("a write should be refused");
assert!(error.to_string().contains("read-only"), "{task:?}");
} else {
outcome.unwrap_or_else(|error| panic!("{task:?} is a read: {error:#}"));
}
}
connection.close().await;
}
#[tokio::test]
async fn reports_errors_from_the_server() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let error = connection
.run_query("SELECT * FROM missing_table")
.await
.expect_err("a missing table should fail");
assert!(error.to_string().contains("missing_table"));
connection.close().await;
}
#[tokio::test]
async fn missing_file_is_an_error() {
let mut config = ConnectionConfig::new(Engine::Sqlite);
config.database = env::temp_dir()
.join(format!("zippa-db-missing-{}.sqlite", Uuid::new_v4()))
.to_string_lossy()
.to_string();
let error = Connection::open(config, None)
.await
.expect_err("opening a missing file should fail");
assert!(error.to_string().contains("could not open"));
}
#[tokio::test]
async fn lists_databases_and_objects() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
assert_eq!(
connection
.databases()
.await
.expect("could not list databases"),
["main"]
);
let objects = connection.objects().await.expect("could not list objects");
let described: Vec<_> = objects
.iter()
.map(|object| (object.label(), object.kind))
.collect();
assert_eq!(
described,
[
("items".to_string(), ObjectKind::Table),
("named_items".to_string(), ObjectKind::View),
]
);
assert!(objects.iter().all(|object| object.schema.is_none()));
connection.close().await;
}
#[tokio::test]
async fn catalog_reads_columns_indexes_and_triggers() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
connection
.run_query("CREATE INDEX items_name_idx ON items(name)")
.await
.expect("could not create the index");
connection
.run_query("CREATE TRIGGER items_guard AFTER INSERT ON items BEGIN SELECT 1; END")
.await
.expect("could not create the trigger");
let catalog = connection
.catalog()
.await
.expect("could not read the catalog");
let find = |kind, name: &str| {
catalog
.entries
.iter()
.find(|entry| entry.kind == kind && entry.name == name)
};
let id = find(CatalogKind::Column, "id").expect("the id column is missing");
assert_eq!(id.detail, "INTEGER");
assert_eq!(id.owner().map(|object| object.name.as_str()), Some("items"));
let index = find(CatalogKind::Index, "items_name_idx").expect("the index is missing");
assert_eq!(index.detail, "name");
assert_eq!(
index.owner().map(|object| object.name.as_str()),
Some("items")
);
let trigger = find(CatalogKind::Trigger, "items_guard").expect("the trigger is missing");
assert_eq!(
trigger.owner().map(|object| object.name.as_str()),
Some("items")
);
assert!(
catalog
.entries
.iter()
.any(|entry| entry.kind == CatalogKind::View && entry.name == "named_items")
);
assert!(
!catalog
.entries
.iter()
.any(|entry| entry.name.starts_with("sqlite_"))
);
assert_eq!(
catalog
.entries
.iter()
.filter(|entry| entry.kind == CatalogKind::Table)
.count(),
1
);
assert_eq!(catalog.truncated(), None);
connection.close().await;
}
#[tokio::test]
async fn switching_database_is_a_new_connection() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let switched = connection
.with_database(connection.database())
.await
.expect("could not reopen the connection");
assert_eq!(switched.database(), connection.database());
switched.close().await;
connection.close().await;
}
#[test]
fn identifiers_are_quoted_only_when_they_have_to_be() {
assert_eq!(quote_identifier("users", Engine::Postgres), "users");
assert_eq!(quote_identifier("user_roles", Engine::MySql), "user_roles");
assert_eq!(quote_identifier("Users", Engine::Postgres), "\"Users\"");
assert_eq!(
quote_identifier("order items", Engine::MySql),
"`order items`"
);
assert_eq!(quote_identifier("2fa", Engine::Sqlite), "\"2fa\"");
assert_eq!(
quote_identifier("we\"ird", Engine::Postgres),
"\"we\"\"ird\""
);
assert_eq!(quote_identifier("we`ird", Engine::MySql), "`we``ird`");
}
fn table(name: &str) -> DatabaseObject {
DatabaseObject {
schema: None,
name: name.to_string(),
kind: ObjectKind::Table,
}
}
#[tokio::test]
async fn column_types_come_back_with_the_columns() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let result = connection
.run_query("SELECT * FROM items ORDER BY id")
.await
.expect("the query failed");
assert_eq!(result.column_types, ["INTEGER", "TEXT", "REAL", "BLOB"]);
connection.close().await;
}
#[tokio::test]
async fn primary_key_columns_are_found() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
assert_eq!(
connection.row_key(&table("items")).await.expect("no key"),
RowKey::Columns(vec!["id".to_string()])
);
connection.close().await;
}
#[tokio::test]
async fn a_view_cannot_be_edited() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let view = DatabaseObject {
kind: ObjectKind::View,
..table("named_items")
};
assert!(matches!(
connection.row_key(&view).await.expect("no key"),
RowKey::Unavailable(_)
));
connection.close().await;
}
#[tokio::test]
async fn a_table_without_a_primary_key_falls_back_to_rowid() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
connection
.execute("CREATE TABLE notes (body TEXT)", Vec::new())
.await
.expect("could not create the table");
assert_eq!(
connection.row_key(&table("notes")).await.expect("no key"),
RowKey::RowId("rowid")
);
connection.close().await;
}
#[tokio::test]
async fn execute_binds_parameters_and_reports_rows_affected() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let affected = connection
.execute(
"UPDATE items SET name = ? WHERE id = ?",
vec![Some("renamed".to_string()), Some("1".to_string())],
)
.await
.expect("the update failed");
assert_eq!(affected, 1);
let result = connection
.run_query("SELECT name FROM items WHERE id = 1")
.await
.expect("the query failed");
assert_eq!(result.rows, [[Some("renamed".to_string())]]);
connection.close().await;
}
#[tokio::test]
async fn execute_writes_sql_null() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let affected = connection
.execute(
"UPDATE items SET name = ? WHERE id = ?",
vec![None, Some("1".to_string())],
)
.await
.expect("the update failed");
assert_eq!(affected, 1);
let result = connection
.run_query("SELECT name FROM items WHERE id = 1")
.await
.expect("the query failed");
let expected: Vec<Vec<Cell>> = vec![vec![None]];
assert_eq!(result.rows, expected);
connection.close().await;
}
#[tokio::test]
async fn execute_reports_no_rows_for_a_missing_key() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let affected = connection
.execute(
"UPDATE items SET name = ? WHERE id = ?",
vec![Some("nobody".to_string()), Some("999".to_string())],
)
.await
.expect("the update failed");
assert_eq!(affected, 0);
connection.close().await;
}
async fn shared(database: &TempDatabase) -> Arc<Connection> {
Arc::new(
Connection::open(database.config(), None)
.await
.expect("could not open the test database"),
)
}
fn finished(step: Result<Step, anyhow::Error>) -> ScriptOutcome {
match step.expect("the script could not run") {
Step::Finished(outcome) => outcome,
Step::Paused(_, failure) => panic!("the script paused on {failure:?}"),
}
}
fn paused(step: Result<Step, anyhow::Error>) -> (ScriptRun, ScriptFailure) {
match step.expect("the script could not run") {
Step::Paused(run, failure) => (*run, failure),
Step::Finished(outcome) => panic!("the script finished: {outcome:?}"),
}
}
async fn item_count(connection: &Connection) -> String {
connection
.run_query("SELECT COUNT(*) FROM items")
.await
.expect("the count failed")
.rows[0][0]
.clone()
.expect("a count is never NULL")
}
const FAILS_IN_THE_MIDDLE: &str = "INSERT INTO items VALUES (3, 'gamma', 3.0, NULL);\n\
INSERT INTO not_a_table VALUES (1);\n\
INSERT INTO items VALUES (4, 'delta', 4.0, NULL);";
#[tokio::test]
async fn a_script_runs_every_statement_in_one_transaction() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let outcome = finished(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
"INSERT INTO items VALUES (3, 'gamma', 3.0, NULL);\n\
SELECT name FROM items WHERE id = 3;",
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
assert_eq!(outcome.results.len(), 2, "one result per statement");
assert_eq!(outcome.results[1].rows, [[Some("gamma".to_string())]]);
assert!(outcome.failures.is_empty());
assert_eq!(outcome.summary(), None, "a clean run has nothing to add");
connection.close().await;
}
#[tokio::test]
async fn a_failure_pauses_the_script_and_rolling_back_undoes_it() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let (run, failure) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
FAILS_IN_THE_MIDDLE,
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
assert_eq!((failure.index, failure.total, failure.line), (1, 3, 2));
assert!(
failure.message.contains("not_a_table"),
"the failure should carry the server's message: {}",
failure.message
);
let outcome = finished(run.resume(Decision::Abort).await);
assert!(outcome.rolled_back);
assert_eq!(
item_count(&connection).await,
"2",
"the insert before the failure should have been rolled back too"
);
connection.close().await;
}
#[tokio::test]
async fn skipping_a_failure_keeps_the_rest_of_the_transaction() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let (run, _) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
FAILS_IN_THE_MIDDLE,
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
let outcome = finished(run.resume(Decision::Skip).await);
assert!(!outcome.rolled_back);
assert_eq!(outcome.results.len(), 2);
assert_eq!(outcome.failures.len(), 1);
assert_eq!(
item_count(&connection).await,
"4",
"both inserts should have been committed"
);
connection.close().await;
}
#[tokio::test]
async fn skipping_every_error_stops_asking() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let (run, _) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
"INSERT INTO not_a_table VALUES (1);\n\
INSERT INTO items VALUES (3, 'gamma', 3.0, NULL);\n\
INSERT INTO nor_this VALUES (1);",
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
let outcome = finished(run.resume(Decision::SkipAll).await);
assert_eq!(
outcome
.failures
.iter()
.map(|failure| failure.index)
.collect::<Vec<_>>(),
[0, 2]
);
assert_eq!(item_count(&connection).await, "3");
let failures = outcome.failures_result().expect("the failures are listed");
assert_eq!(failures.columns, ["#", "line", "statement", "error"]);
assert_eq!(failures.rows.len(), 2);
connection.close().await;
}
#[tokio::test]
async fn ignoring_errors_without_a_transaction_keeps_what_worked() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let outcome = finished(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
FAILS_IN_THE_MIDDLE,
ScriptMode::Autocommit,
OnFailure::Skip,
)
.await,
);
assert_eq!(outcome.failures.len(), 1);
assert_eq!(item_count(&connection).await, "4");
connection.close().await;
}
#[tokio::test]
async fn stopping_without_a_transaction_keeps_what_ran_before() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let (run, _) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
FAILS_IN_THE_MIDDLE,
ScriptMode::Autocommit,
OnFailure::Ask,
)
.await,
);
let outcome = finished(run.resume(Decision::Abort).await);
assert!(outcome.stopped && !outcome.rolled_back);
assert_eq!(item_count(&connection).await, "3");
connection.close().await;
}
#[tokio::test]
async fn closing_a_paused_script_rolls_its_transaction_back() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let (run, _) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
FAILS_IN_THE_MIDDLE,
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
run.close().await;
assert_eq!(item_count(&connection).await, "2");
connection.close().await;
}
#[tokio::test]
async fn a_scripts_statements_are_logged_one_by_one() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
connection.query_log().clear();
finished(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
FAILS_IN_THE_MIDDLE,
ScriptMode::Transaction,
OnFailure::Skip,
)
.await,
);
let log = connection.query_log().snapshot();
let sql: Vec<&str> = log.iter().map(|entry| entry.sql.as_str()).collect();
assert_eq!(sql.len(), 3, "the savepoints are not the user's: {sql:?}");
assert!(
log.iter()
.any(|entry| matches!(entry.outcome, QueryOutcome::Error(_))
&& entry.sql.contains("not_a_table")),
"the failure is logged with its error"
);
assert!(log.iter().all(|entry| entry.source == QuerySource::User));
connection.close().await;
}
#[tokio::test]
async fn a_read_only_connection_refuses_a_script_before_it_runs() {
let database = TempDatabase::new().await;
let connection = Arc::new(
Connection::open(
ConnectionConfig {
safety: SafetyMode::ReadOnly,
..database.config()
},
None,
)
.await
.expect("could not open the test database"),
);
let error = match ScriptRun::start(
PinnedConnection::new(connection.clone()),
"SELECT 1; DELETE FROM items",
ScriptMode::Transaction,
OnFailure::Ask,
)
.await
{
Ok(_) => panic!("a read-only connection should refuse the delete"),
Err(error) => format!("{error:#}"),
};
assert!(
error.contains("read-only") && error.contains("DELETE"),
"{error}"
);
connection.close().await;
}
#[tokio::test]
async fn a_pinned_connection_keeps_a_transaction_between_runs() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let pinned = PinnedConnection::new(connection.clone());
assert_eq!(pinned.state(), TxnState::Idle, "nothing is checked out yet");
pinned.run_query("BEGIN").await.expect("begin");
assert_eq!(pinned.state(), TxnState::Open);
pinned
.run_query("INSERT INTO items VALUES (3, 'gamma', 3.0, NULL)")
.await
.expect("insert");
assert_eq!(pinned.state(), TxnState::Open);
let inside = pinned
.run_query("SELECT COUNT(*) FROM items")
.await
.expect("count");
assert_eq!(
inside.rows[0][0].as_deref(),
Some("3"),
"the tab sees its own insert"
);
pinned.run_query("ROLLBACK").await.expect("rollback");
assert_eq!(pinned.state(), TxnState::Idle);
assert_eq!(
item_count(&connection).await,
"2",
"the insert was rolled back"
);
drop(pinned);
connection.close().await;
}
#[tokio::test]
async fn a_pinned_connection_keeps_session_state_between_runs() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let pinned = PinnedConnection::new(connection.clone());
pinned
.run_query("CREATE TEMP TABLE scratch (a)")
.await
.expect("create temp");
pinned
.run_query("INSERT INTO scratch VALUES (1)")
.await
.expect("insert temp");
pinned
.run_query("PRAGMA foreign_keys = ON")
.await
.expect("pragma");
let rows = pinned
.run_query("SELECT a FROM scratch")
.await
.expect("read temp");
assert_eq!(rows.rows, [[Some("1".to_string())]]);
let keys = pinned
.run_query("PRAGMA foreign_keys")
.await
.expect("read pragma");
assert_eq!(keys.rows[0][0].as_deref(), Some("1"));
assert!(connection.run_query("SELECT a FROM scratch").await.is_err());
drop(pinned);
connection.close().await;
}
#[tokio::test]
async fn a_committed_transaction_is_visible_to_the_pool() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let pinned = PinnedConnection::new(connection.clone());
pinned.run_query("BEGIN").await.expect("begin");
pinned
.run_query("DELETE FROM items WHERE id = 2")
.await
.expect("delete");
pinned.run_query("COMMIT").await.expect("commit");
assert_eq!(pinned.state(), TxnState::Idle);
assert_eq!(item_count(&connection).await, "1");
drop(pinned);
connection.close().await;
}
#[tokio::test]
async fn a_read_only_pin_runs_transaction_control_but_refuses_writes() {
let database = TempDatabase::new().await;
let connection = Arc::new(
Connection::open(
ConnectionConfig {
safety: SafetyMode::ReadOnly,
..database.config()
},
None,
)
.await
.expect("could not open the test database"),
);
let pinned = PinnedConnection::new(connection.clone());
pinned
.run_query("BEGIN")
.await
.expect("begin is not a write");
assert_eq!(pinned.state(), TxnState::Open);
let error = pinned
.run_query("DELETE FROM items")
.await
.expect_err("a read-only connection refuses a delete");
assert!(format!("{error:#}").contains("DELETE"), "{error:#}");
pinned.run_query("ROLLBACK").await.expect("rollback");
assert_eq!(pinned.state(), TxnState::Idle);
drop(pinned);
connection.close().await;
}
#[tokio::test]
async fn a_script_inside_an_open_transaction_rolls_back_only_itself() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let pinned = PinnedConnection::new(connection.clone());
pinned.run_query("BEGIN").await.expect("begin");
pinned
.run_query("INSERT INTO items VALUES (10, 'mine', 1.0, NULL)")
.await
.expect("insert");
let (run, _) = paused(
ScriptRun::start(
pinned.clone(),
FAILS_IN_THE_MIDDLE,
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
let outcome = finished(run.resume(Decision::Abort).await);
assert!(outcome.rolled_back);
assert_eq!(
pinned.state(),
TxnState::Open,
"the tab's transaction is still open"
);
let count = pinned
.run_query("SELECT COUNT(*) FROM items")
.await
.expect("count");
assert_eq!(
count.rows[0][0].as_deref(),
Some("3"),
"only the tab's insert is left"
);
pinned.run_query("COMMIT").await.expect("commit");
assert_eq!(item_count(&connection).await, "3");
drop(pinned);
connection.close().await;
}
#[tokio::test]
async fn a_scripts_own_begin_leaves_the_tab_in_a_transaction() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let pinned = PinnedConnection::new(connection.clone());
finished(
ScriptRun::start(
pinned.clone(),
"BEGIN;\nINSERT INTO items VALUES (3, 'gamma', 3.0, NULL);",
ScriptMode::Autocommit,
OnFailure::Ask,
)
.await,
);
assert_eq!(pinned.state(), TxnState::Open);
pinned.run_query("ROLLBACK").await.expect("rollback");
assert_eq!(item_count(&connection).await, "2");
drop(pinned);
connection.close().await;
}
#[tokio::test]
async fn pinned_connections_are_capped_per_connection() {
let database = TempDatabase::new().await;
let connection = shared(&database).await;
let mut pins = Vec::new();
for _ in 0..super::connection::PINNED_MAX {
let pinned = PinnedConnection::new(connection.clone());
pinned.run_query("SELECT 1").await.expect("within the cap");
pins.push(pinned);
}
let one_more = PinnedConnection::new(connection.clone());
let error = one_more
.run_query("SELECT 1")
.await
.expect_err("past the cap");
assert!(
format!("{error:#}").contains("too many query tabs"),
"{error:#}"
);
assert_eq!(item_count(&connection).await, "2");
drop(pins.pop());
one_more
.run_query("SELECT 1")
.await
.expect("a slot was freed");
drop((pins, one_more));
connection.close().await;
}
#[tokio::test]
async fn a_failing_schema_statement_rolls_the_whole_change_back() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let error = connection
.execute_script(&[
"ALTER TABLE items ADD COLUMN note TEXT".to_string(),
"ALTER TABLE not_a_table ADD COLUMN note TEXT".to_string(),
])
.await
.expect_err("the second statement should fail");
assert!(
format!("{error:#}").contains("statement 2"),
"the error should name which statement failed: {error:#}"
);
let result = connection
.run_query("SELECT name FROM pragma_table_info('items') WHERE name = 'note'")
.await
.expect("the query failed");
assert!(
result.rows.is_empty(),
"the column added before the failure should have been rolled back too"
);
connection.close().await;
}
#[test]
fn a_connection_saved_before_safety_modes_reads_back_as_staged() {
let saved = r#"{
"id": "00000000-0000-0000-0000-000000000001",
"name": "old",
"engine": "Postgres",
"host": "localhost",
"port": 5432,
"username": "postgres",
"database": "app"
}"#;
let config: ConnectionConfig =
serde_json::from_str(saved).expect("an older connection should still load");
assert_eq!(config.safety, SafetyMode::Staged);
let written = serde_json::to_string(&ConnectionConfig {
safety: SafetyMode::AutoApply,
..ConnectionConfig::new(Engine::Postgres)
})
.expect("the config should serialize");
assert!(
written.contains("\"safety\":\"AutoApply\""),
"the mode belongs in the file: {written}"
);
}
#[test]
fn a_connection_saved_before_colouring_reads_back_uncoloured() {
let saved = r#"{
"id": "00000000-0000-0000-0000-000000000001",
"name": "old",
"engine": "Postgres",
"host": "localhost",
"port": 5432,
"username": "postgres",
"database": "app",
"safety": "Staged"
}"#;
let config: ConnectionConfig =
serde_json::from_str(saved).expect("an older connection should still load");
assert_eq!(config.color, None);
assert_eq!(config.last_connected, None);
}
#[test]
fn a_coloured_connection_survives_a_round_trip_through_the_file() {
let config = ConnectionConfig {
name: "Prod DB".into(),
color: Some(TagColor::Red),
last_connected: Some(
"2024-01-02T03:04:05Z"
.parse::<chrono::DateTime<chrono::Utc>>()
.expect("a fixed timestamp"),
),
..ConnectionConfig::new(Engine::Postgres)
};
let written = serde_json::to_string(&config).expect("the config should serialize");
let read: ConnectionConfig = serde_json::from_str(&written).expect("the config should parse");
assert_eq!(read.color, Some(TagColor::Red));
assert_eq!(read.last_connected, config.last_connected);
}
#[test]
fn a_connection_saved_with_a_label_still_loads() {
let saved = r#"{
"id": "7d0f3f3e-5b1a-4c55-9f4b-0f1d2c3b4a59",
"name": "app",
"engine": "Postgres",
"host": "localhost",
"port": 5432,
"username": "postgres",
"database": "app",
"tag": "Production",
"color": "red"
}"#;
let config: ConnectionConfig =
serde_json::from_str(saved).expect("a labelled connection should still load");
assert_eq!(config.color, Some(TagColor::Red));
}
#[test]
fn tag_color_keys_serialise_as_kebab_case() {
for color in TagColor::ALL {
let written = serde_json::to_string(&color).expect("the colour should serialize");
assert_eq!(written, format!("\"{}\"", color.key()));
let read: TagColor = serde_json::from_str(&written).expect("the colour should parse");
assert_eq!(read, color);
}
}
#[test]
fn only_postgres_casts_its_placeholders() {
assert_eq!(
typed_placeholder(Engine::Postgres, 2, "INT4"),
"cast($2 as INT4)"
);
assert_eq!(typed_placeholder(Engine::MySql, 2, "INT"), "?");
assert_eq!(typed_placeholder(Engine::Sqlite, 2, "INTEGER"), "?");
assert_eq!(typed_placeholder(Engine::Postgres, 1, ""), "$1");
}
#[test]
fn keyword_literals_are_recognized_up_to_case_and_whitespace() {
assert_eq!(keyword_literal("now()"), Some("NOW()"));
assert_eq!(
keyword_literal(" Current_Timestamp "),
Some("CURRENT_TIMESTAMP")
);
assert_eq!(keyword_literal("current_date"), Some("CURRENT_DATE"));
assert_eq!(keyword_literal("current_time"), Some("CURRENT_TIME"));
assert_eq!(keyword_literal("now"), None);
assert_eq!(keyword_literal("it's now()"), None);
assert_eq!(keyword_literal(""), None);
}
async fn open_with(database: &TempDatabase, safety: SafetyMode) -> Connection {
let config = ConnectionConfig {
safety,
..database.config()
};
Connection::open(config, None)
.await
.expect("could not open the test database")
}
#[tokio::test]
async fn a_read_only_connection_refuses_writes_and_names_them() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::ReadOnly).await;
let error = connection
.run_query("update items set name = 'x' where id = 1")
.await
.expect_err("a read-only connection should refuse an UPDATE");
let error = format!("{error:#}");
assert!(
error.contains("read-only") && error.contains("UPDATE"),
"the error should say what was refused: {error}"
);
connection
.execute("update items set name = ? where id = ?", vec![None, None])
.await
.expect_err("a read-only connection should refuse an inline edit too");
let result = connection
.run_query("select name from items where id = 1")
.await
.expect("a read-only connection should still read");
assert_eq!(result.rows, [[Some("alpha".to_string())]]);
connection.close().await;
}
#[tokio::test]
async fn a_read_only_pool_refuses_a_write_the_client_did_not_catch() {
let database = TempDatabase::new().await;
let config = ConnectionConfig {
safety: SafetyMode::ReadOnly,
..database.config()
};
let pool = super::sqlite::connect(&config)
.await
.expect("could not open the test database");
let error = sqlx::query("insert into items values (9, 'nine', 9.0, NULL)")
.execute(&pool)
.await
.expect_err("the file should be open read-only");
pool.close().await;
let error = format!("{error}").to_ascii_lowercase();
assert!(
error.contains("readonly") || error.contains("read-only"),
"the server should be the one refusing here: {error}"
);
}
#[tokio::test]
async fn the_other_modes_still_write() {
let database = TempDatabase::new().await;
for safety in [
SafetyMode::ConfirmWrites,
SafetyMode::Staged,
SafetyMode::AutoApply,
] {
let connection = open_with(&database, safety).await;
let affected = connection
.execute(
"update items set name = ? where id = ?",
vec![Some(format!("{safety:?}")), Some("1".to_string())],
)
.await
.expect("the update failed");
assert_eq!(affected, 1, "{safety:?} should write");
connection.close().await;
}
}
#[tokio::test]
async fn table_schema_reads_columns_and_marks_the_primary_key() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
let schema = connection
.table_schema(&table("items"))
.await
.expect("could not read the schema");
let names: Vec<&str> = schema.columns.iter().map(|c| c.name.as_str()).collect();
assert_eq!(names, ["id", "name", "score", "payload"]);
assert!(schema.columns[0].is_primary_key);
assert!(!schema.columns[1].is_primary_key);
assert!(schema.indexes.is_empty());
assert!(schema.foreign_keys.is_empty());
connection.close().await;
}
#[tokio::test]
async fn table_schema_reads_a_unique_index_and_a_foreign_key() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None)
.await
.expect("could not open the test database");
connection
.execute(
"CREATE TABLE tags (id INTEGER PRIMARY KEY, label TEXT NOT NULL)",
Vec::new(),
)
.await
.expect("could not create tags");
connection
.execute(
"CREATE UNIQUE INDEX tags_label_idx ON tags(label)",
Vec::new(),
)
.await
.expect("could not create the index");
connection
.execute(
"CREATE TABLE tagged_items ( \
item_id INTEGER, \
tag_id INTEGER, \
FOREIGN KEY(item_id) REFERENCES items(id) ON DELETE CASCADE, \
FOREIGN KEY(tag_id) REFERENCES tags(id) \
)",
Vec::new(),
)
.await
.expect("could not create tagged_items");
let tags = connection
.table_schema(&table("tags"))
.await
.expect("could not read the tags schema");
let index = tags
.indexes
.iter()
.find(|index| index.name == "tags_label_idx")
.expect("the unique index should be listed");
assert!(index.unique);
assert!(!index.is_primary_key);
assert_eq!(index.columns, ["label"]);
let tagged = connection
.table_schema(&table("tagged_items"))
.await
.expect("could not read the tagged_items schema");
assert_eq!(tagged.foreign_keys.len(), 2);
let to_items = tagged
.foreign_keys
.iter()
.find(|fk| fk.referenced_table == "items")
.expect("the foreign key to items should be listed");
assert_eq!(to_items.columns, ["item_id"]);
assert_eq!(to_items.referenced_columns, ["id"]);
assert_eq!(to_items.on_delete, ReferentialAction::Cascade);
assert_eq!(to_items.on_update, ReferentialAction::NoAction);
connection.close().await;
}
#[test]
fn a_width_free_cast_stands_in_for_one_that_would_truncate() {
assert_eq!(
typed_placeholder(Engine::Postgres, 1, "CHAR"),
"cast($1 as text)"
);
assert_eq!(
typed_placeholder(Engine::Postgres, 2, "BIT"),
"cast($2 as varbit)"
);
assert_eq!(
typed_placeholder(Engine::Postgres, 3, "VARBIT"),
"cast($3 as varbit)"
);
assert_eq!(
typed_placeholder(Engine::Postgres, 1, "TIMESTAMPTZ"),
"cast($1 as TIMESTAMPTZ)"
);
assert_eq!(
typed_placeholder(Engine::Postgres, 1, "INT4[]"),
"cast($1 as INT4[])"
);
assert_eq!(typed_placeholder(Engine::Postgres, 1, ""), "$1");
}
#[test]
fn mysql_bit_digits_are_converted_rather_than_stored_as_text() {
assert_eq!(
typed_placeholder(Engine::MySql, 1, "BIT"),
"cast(conv(?, 2, 10) as unsigned)"
);
assert_eq!(typed_placeholder(Engine::MySql, 1, "DATETIME"), "?");
assert_eq!(typed_placeholder(Engine::Sqlite, 1, "DATETIME"), "?");
}
#[test]
fn mysql_metadata_escapes_a_backslash_the_way_the_server_reads_it() {
let name = r"odd\name";
let queries = [
super::mysql::primary_key_sql(name),
super::mysql::columns_sql(name),
super::mysql::indexes_sql(name),
super::mysql::foreign_keys_sql(name),
];
for sql in queries {
assert!(
sql.contains(r"'odd\\name'"),
"the name should be escaped for MySQL: {sql}"
);
}
}
#[tokio::test]
async fn explain_reads_a_sqlite_plan_tree() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::default()).await;
let explained = connection
.explain("select * from items where id = 1", false)
.await
.expect("a select should explain");
let super::plan::Explained::Plan(plan) = explained else {
panic!("SQLite should come back as a tree");
};
assert!(plan.parsed, "the rows are the shape SQLite documents");
let labels = labels(&plan.root);
assert!(
labels.iter().any(|label| label.contains("items")),
"the plan should name the table: {labels:?}"
);
connection.close().await;
}
fn labels(node: &super::plan::PlanNode) -> Vec<String> {
let mut all = vec![node.label.clone()];
for child in &node.children {
all.extend(labels(child));
}
all
}
#[tokio::test]
async fn explain_refuses_to_analyze_a_write() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::AutoApply).await;
for sql in [
"insert into items values (3, 'gamma', 3.0, NULL)",
"with gone as (delete from items returning *) select * from gone",
"explain analyze delete from items",
] {
let error = connection
.explain(sql, true)
.await
.expect_err("analyzing a write should be refused");
assert!(
format!("{error:#}").contains("changes data"),
"{sql}: {error:#}"
);
}
connection.close().await;
}
#[tokio::test]
async fn explain_refuses_to_analyze_a_read_on_sqlite() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::default()).await;
let error = connection
.explain("select * from items", true)
.await
.expect_err("SQLite has no EXPLAIN ANALYZE");
assert!(
format!("{error:#}").contains("EXPLAIN ANALYZE"),
"the refusal should name the missing form: {error:#}"
);
let explained = connection
.explain("select * from items", false)
.await
.expect("a plain explain should still be allowed");
assert!(matches!(explained, super::plan::Explained::Plan(_)));
connection.close().await;
}
#[tokio::test]
async fn plain_explain_of_a_write_does_not_run_it() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::AutoApply).await;
connection
.explain("delete from items", false)
.await
.expect("explaining a write without ANALYZE is safe");
let count = connection
.run_query("select count(*) from items")
.await
.expect("could not count the rows");
assert_eq!(
count.rows[0][0].as_deref(),
Some("2"),
"nothing should be deleted"
);
connection.close().await;
}
#[tokio::test]
async fn a_read_only_connection_can_still_explain() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::ReadOnly).await;
connection
.explain("select * from items", false)
.await
.expect("a read-only connection should explain a read");
let _ = connection.explain("delete from items", true).await;
let count = connection
.run_query("select count(*) from items")
.await
.expect("could not count the rows");
assert_eq!(count.rows[0][0].as_deref(), Some("2"));
connection.close().await;
}
#[tokio::test]
async fn explain_needs_a_statement() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::default()).await;
let error = connection
.explain(" ; ", false)
.await
.expect_err("an empty buffer has nothing to explain");
assert!(format!("{error:#}").contains("Select one statement"));
let error = connection
.explain("select 1; select 2", false)
.await
.expect_err("a script has no single plan");
assert!(format!("{error:#}").contains("Select one statement"));
connection.close().await;
}
#[tokio::test]
async fn a_failed_table_rebuild_is_rolled_back() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::default()).await;
let error = connection
.rebuild_table(vec![
"CREATE TABLE scratch (x INTEGER)".to_string(),
"this is not sql".to_string(),
])
.await
.expect_err("a broken statement should fail the rebuild");
assert!(format!("{error:#}").contains("statement 2 of 2"));
let leftover = connection
.run_query("select name from sqlite_master where name = 'scratch'")
.await
.expect("could not read the schema");
assert!(leftover.rows.is_empty(), "the scratch table should be gone");
connection.close().await;
}
#[tokio::test]
async fn a_read_only_connection_refuses_a_table_rebuild() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::ReadOnly).await;
let error = connection
.rebuild_table(vec!["CREATE TABLE scratch (x INTEGER)".to_string()])
.await
.expect_err("a read-only connection must not rebuild");
assert!(format!("{error:#}").contains("read-only"));
connection.close().await;
}
#[tokio::test]
async fn a_rebuild_leaves_the_data_and_the_schema_behind() {
let database = TempDatabase::new().await;
let connection = open_with(&database, SafetyMode::default()).await;
connection
.rebuild_table(vec![
"CREATE TABLE new_items (id INTEGER, name TEXT, score TEXT, payload BLOB)".to_string(),
"INSERT INTO new_items (id, name, score, payload) \
SELECT id, name, score, payload FROM items"
.to_string(),
"DROP TABLE items".to_string(),
"ALTER TABLE new_items RENAME TO items".to_string(),
])
.await
.expect("a plain rebuild should run");
let score = connection
.run_query("select score from items where id = 1")
.await
.expect("could not read the rebuilt table");
assert_eq!(score.rows[0][0].as_deref(), Some("1.5"));
let shape = connection
.run_query("select name from pragma_table_info('items') order by cid")
.await
.expect("could not read the rebuilt columns");
let columns: Vec<String> = shape
.rows
.iter()
.filter_map(|row| row.first().cloned().flatten())
.collect();
assert_eq!(columns, ["id", "name", "score", "payload"]);
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see the doc comment"]
async fn live_mysql_routines_carry_their_argument_types() {
let (connection, pool) = live_mysql().await;
for sql in [
"DROP FUNCTION IF EXISTS zippa_add_one",
"DROP FUNCTION IF EXISTS zippa_no_args",
"DROP PROCEDURE IF EXISTS zippa_do_thing",
"CREATE FUNCTION zippa_add_one(x INT) RETURNS INT DETERMINISTIC RETURN x + 1",
"CREATE FUNCTION zippa_no_args() RETURNS INT DETERMINISTIC RETURN 42",
"CREATE PROCEDURE zippa_do_thing(IN p_id INT, OUT p_name VARCHAR(50)) \
SELECT p_id INTO p_name",
] {
pool.execute(AssertSqlSafe(sql.to_string()))
.await
.unwrap_or_else(|error| panic!("{sql}: {error}"));
}
let routines = connection
.stored_objects()
.await
.expect("could not read the routines");
let label = |name: &str| {
routines
.iter()
.find(|routine| routine.name == name)
.map(|routine| routine.label())
};
assert_eq!(
label("zippa_add_one").as_deref(),
Some("zippa_add_one(int)")
);
assert_eq!(label("zippa_no_args").as_deref(), Some("zippa_no_args()"));
assert_eq!(
label("zippa_do_thing").as_deref(),
Some("zippa_do_thing(int, varchar(50))"),
"a procedure's parameter types should tell it from another signature"
);
connection.close().await;
pool.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see live_mysql_routines_carry_their_argument_types"]
async fn live_mysql_table_schema_carries_auto_increment_collation_comment_and_on_update() {
let (connection, pool) = live_mysql().await;
for sql in [
"DROP TABLE IF EXISTS zippa_probe",
"CREATE TABLE zippa_probe (
id INT AUTO_INCREMENT PRIMARY KEY,
name VARCHAR(50) COLLATE utf8mb4_bin COMMENT 'the display name',
updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP \
ON UPDATE CURRENT_TIMESTAMP,
price DECIMAL(10,2),
tax DECIMAL(10,2) GENERATED ALWAYS AS (price * 0.1) STORED,
plain INT DEFAULT 5
)",
] {
pool.execute(AssertSqlSafe(sql.to_string()))
.await
.unwrap_or_else(|error| panic!("{sql}: {error}"));
}
let object = DatabaseObject {
schema: None,
name: "zippa_probe".to_string(),
kind: ObjectKind::Table,
};
let schema = connection
.table_schema(&object)
.await
.expect("could not read the table schema");
let column = |name: &str| {
schema
.columns
.iter()
.find(|column| column.name == name)
.unwrap_or_else(|| panic!("no {name} column came back"))
};
assert!(column("id").mysql_extra.auto_increment);
assert!(!column("plain").mysql_extra.auto_increment);
assert!(column("updated_at").mysql_extra.on_update_current_timestamp);
assert!(!column("plain").mysql_extra.on_update_current_timestamp);
assert_eq!(
column("name").mysql_extra.collation.as_deref(),
Some("utf8mb4_bin")
);
assert_eq!(column("plain").mysql_extra.collation, None);
assert_eq!(
column("name").mysql_extra.comment.as_deref(),
Some("the display name")
);
assert_eq!(
column("plain").mysql_extra.comment,
None,
"an unset comment is an empty string on the server, folded to None"
);
assert_eq!(
column("tax").mysql_extra.generation_expression.as_deref(),
Some("(`price` * 0.1)")
);
assert_eq!(column("plain").mysql_extra.generation_expression, None);
connection.close().await;
pool.execute(AssertSqlSafe("DROP TABLE zippa_probe".to_string()))
.await
.expect("could not drop the probe table");
pool.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server with a yen-locale database; see the doc comment"]
async fn live_postgres_money_scales_by_the_servers_locale_not_always_by_100() {
let usual = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
let result = usual
.run_query("SELECT '12.34'::money")
.await
.expect("could not read the ordinary-locale money value");
assert_eq!(
result.rows[0][0].as_deref(),
Some("12.34"),
"a two-digit locale should read back the way it was written"
);
usual.close().await;
let yen = live_postgres("ZIPPA_TEST_POSTGRES_YEN_URL", "zippa_yen").await;
let result = yen
.run_query("SELECT '1234'::money")
.await
.expect("could not read the yen-locale money value");
assert_eq!(
result.rows[0][0].as_deref(),
Some("1234"),
"a zero-digit locale's value should not be shown divided by 100"
);
yen.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see the doc comment"]
async fn live_postgres_ssl_require_is_refused_by_a_server_without_ssl() {
for mode in [SslMode::Disable, SslMode::Prefer] {
let (mut config, password) = live_postgres_config("ZIPPA_TEST_POSTGRES_URL", "app");
config.ssl.mode = mode;
let connection = Connection::open(config, Some(password))
.await
.unwrap_or_else(|error| panic!("{mode:?} should connect: {error:#}"));
connection.close().await;
}
let (mut config, password) = live_postgres_config("ZIPPA_TEST_POSTGRES_URL", "app");
config.ssl.mode = SslMode::Require;
assert!(
Connection::open(config, Some(password)).await.is_err(),
"Require should refuse a server that does not offer SSL"
);
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see the doc comment"]
async fn live_mysql_ssl_require_encrypts_and_verify_refuses_a_self_signed_server() {
let (mut config, password) = live_mysql_config();
config.ssl = SslConfig {
mode: SslMode::Require,
..SslConfig::default()
};
let encrypted = Connection::open(config.clone(), Some(password.clone()))
.await
.expect("Require should connect to a server with SSL");
let cipher = encrypted
.run_query("SHOW SESSION STATUS LIKE 'Ssl_cipher'")
.await
.expect("could not read the session's cipher");
assert!(
cipher.rows[0][1]
.as_deref()
.is_some_and(|cipher| !cipher.is_empty()),
"a Require session should be encrypted: {:?}",
cipher.rows
);
encrypted.close().await;
config.ssl.mode = SslMode::VerifyCa;
assert!(
Connection::open(config, Some(password)).await.is_err(),
"VerifyCa should refuse a self-signed certificate"
);
}
fn live_ssh() -> (SshConfig, String, String, u16) {
let url = env::var("ZIPPA_TEST_SSH_URL")
.unwrap_or_else(|_| "ssh://tunnel:secret@127.0.0.1:2222".to_string());
let rest = url.strip_prefix("ssh://").expect("an ssh:// URL");
let (credentials, authority) = rest.split_once('@').expect("user:pass@host:port");
let (username, password) = credentials.split_once(':').expect("user:pass");
let (host, port) = authority.split_once(':').expect("host:port");
let ssh = SshConfig {
enabled: true,
host: host.to_string(),
port: port.parse().expect("a port"),
username: username.to_string(),
auth: SshAuth::Password,
key_path: String::new(),
};
let target = env::var("ZIPPA_TEST_SSH_TARGET").unwrap_or_else(|_| "postgres:5432".to_string());
let (target_host, target_port) = target.split_once(':').expect("host:port");
(
ssh,
password.to_string(),
target_host.to_string(),
target_port.parse().expect("a port"),
)
}
#[tokio::test]
#[ignore = "needs a live Postgres server and SSH jump host; see the doc comment"]
async fn live_postgres_connects_through_an_ssh_tunnel() {
let (ssh, ssh_password, target_host, target_port) = live_ssh();
let (mut config, password) = live_postgres_config("ZIPPA_TEST_POSTGRES_URL", "app");
config.host = target_host.clone();
config.port = target_port;
config.ssh = ssh;
let connection = Connection::open_with(
config,
Credentials {
password: Some(password),
ssh: Some(ssh_password),
},
)
.await
.expect("should connect through the tunnel");
let result = connection
.run_query("SELECT 1 + 1")
.await
.expect("should query through the tunnel");
assert_eq!(result.rows[0][0].as_deref(), Some("2"));
assert_eq!(connection.config.host, target_host);
let other = connection
.with_database("postgres")
.await
.expect("switching database should reuse the tunnel");
other.run_query("SELECT 1").await.expect("query");
other.close().await;
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live SSH jump host; see `live_ssh`"]
async fn live_ssh_tunnel_reports_a_refused_password() {
let (ssh, _, target_host, target_port) = live_ssh();
let (mut config, password) = live_postgres_config("ZIPPA_TEST_POSTGRES_URL", "app");
config.host = target_host;
config.port = target_port;
config.ssh = ssh;
let error = Connection::open_with(
config,
Credentials {
password: Some(password),
ssh: Some("not-the-password".into()),
},
)
.await
.expect_err("a wrong SSH password should be refused");
assert!(
format!("{error:#}").contains("refused the password"),
"{error:#}"
);
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see the doc comment"]
async fn live_postgres_geometry_ranges_hstore_and_macaddr8_read_back_as_text() {
let connection = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
connection
.execute("CREATE EXTENSION IF NOT EXISTS hstore", vec![])
.await
.expect("could not install hstore");
let cases = [
("'(1.5,-2)'::point", "(1.5,-2)"),
("'{1,-1,0}'::line", "{1,-1,0}"),
("'[(0,0),(1,1)]'::lseg", "[(0,0),(1,1)]"),
("'((0,0),(2,2))'::box", "(2,2),(0,0)"),
("'((0,0),(1,1),(2,0))'::path", "((0,0),(1,1),(2,0))"),
("'[(0,0),(1,1)]'::path", "[(0,0),(1,1)]"),
("'((0,0),(1,1),(2,0))'::polygon", "((0,0),(1,1),(2,0))"),
("'<(1,2),3>'::circle", "<(1,2),3>"),
("ARRAY['(1,2)'::point, NULL]", "{\"(1,2)\",NULL}"),
("'[1,10)'::int4range", "[1,10)"),
("'(,5]'::int8range", "(,6)"),
("'empty'::int4range", "empty"),
("'1.5'::numeric", "1.5"),
("'1.50'::numeric", "1.50"),
("'[1.50,)'::numrange", "[1.50,)"),
("'[1.5,2.5]'::numrange", "[1.5,2.5]"),
(
"'[2024-01-01,2024-02-01)'::daterange",
"[2024-01-01,2024-02-01)",
),
(
"'[2024-01-01 10:00,2024-01-02)'::tsrange",
"[\"2024-01-01 10:00:00\",\"2024-01-02 00:00:00\")",
),
(
"'a=>1, \"b c\"=>NULL, q=>\"say \\\"hi\\\"\"'::hstore",
"\"a\"=>\"1\", \"b c\"=>NULL, \"q\"=>\"say \\\"hi\\\"\"",
),
(
"'08:00:2b:01:02:03:04:05'::macaddr8",
"08:00:2b:01:02:03:04:05",
),
];
for (expression, expected) in cases {
let result = connection
.run_query(&format!("SELECT {expression}"))
.await
.unwrap_or_else(|error| panic!("could not read {expression}: {error}"));
let shown = result.rows[0][0].clone();
assert_eq!(shown.as_deref(), Some(expected), "{expression}");
let type_name = &result.column_types[0];
let round_trip = connection
.run_query_with(
&format!(
"SELECT {}::text, ({expression})::text",
crate::db::sql::typed_placeholder(Engine::Postgres, 1, type_name)
),
vec![shown],
)
.await
.unwrap_or_else(|error| panic!("could not cast {expected} back: {error}"));
assert_eq!(
round_trip.rows[0][0], round_trip.rows[0][1],
"{expected} should read back as {expression}"
);
}
let result = connection
.run_query("SELECT '[2024-01-01 00:00+00,)'::tstzrange")
.await
.expect("could not read a tstzrange");
let shown = result.rows[0][0].clone().unwrap_or_default();
assert!(
shown.starts_with("[\"2024-") && shown.ends_with("\",)"),
"{shown}"
);
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see the doc comment"]
async fn live_postgres_processes_lists_and_kills_another_connection() {
let watcher = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
let victim = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
let sleeper = tokio::spawn(async move { victim.run_query("SELECT pg_sleep(30)").await });
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
let processes = watcher
.processes()
.await
.expect("could not read the process list");
let row = processes
.rows
.iter()
.find(|row| {
row.get(6)
.and_then(|cell| cell.as_deref())
.is_some_and(|query| query.contains("pg_sleep"))
})
.expect("the sleeping connection should be in the process list");
let pid = row[0].clone().expect("a pid");
watcher
.kill_process(&pid)
.await
.expect("could not end the sleeping connection");
let result = sleeper.await.expect("the task itself should not panic");
assert!(
result.is_err(),
"the killed connection's query should have been interrupted"
);
watcher.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see live_mysql_routines_carry_their_argument_types"]
async fn live_mysql_processes_lists_and_kills_another_connection() {
let (watcher, _watcher_pool) = live_mysql().await;
let (victim, _victim_pool) = live_mysql().await;
let sleeper = tokio::spawn(async move { victim.run_query("SELECT sleep(30)").await });
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
let processes = watcher
.processes()
.await
.expect("could not read the process list");
let row = processes
.rows
.iter()
.find(|row| {
row.get(7)
.and_then(|cell| cell.as_deref())
.is_some_and(|info| info.contains("sleep"))
})
.expect("the sleeping connection should be in the process list");
let id = row[0].clone().expect("an id");
watcher
.kill_process(&id)
.await
.expect("could not end the sleeping connection");
let result = sleeper.await.expect("the task itself should not panic");
assert!(
result.is_err(),
"the killed connection's query should have been interrupted"
);
watcher.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see live_postgres_processes_lists_and_kills_another_connection"]
async fn live_postgres_statement_timeout_cancels_a_long_statement() {
let connection = live_postgres_with("ZIPPA_TEST_POSTGRES_URL", "app", |config| {
config.statement_timeout = Some(1);
})
.await;
let shown = connection
.run_query("SHOW statement_timeout")
.await
.expect("could not read the setting");
assert_eq!(shown.rows[0][0].as_deref(), Some("1s"));
let error = connection
.run_query("SELECT count(*) FROM generate_series(1, 10000000000) AS zippa_timeout_test")
.await
.expect_err("the server should cancel a statement past the timeout");
assert!(
format!("{error:#}").contains("statement timeout"),
"{error:#}"
);
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see live_mysql_routines_carry_their_argument_types"]
async fn live_mysql_statement_timeout_is_set_on_the_session() {
let (connection, _pool) = live_mysql_with(|config| config.statement_timeout = Some(2)).await;
let shown = connection
.run_query("SELECT @@SESSION.max_execution_time")
.await
.expect("could not read the setting");
assert_eq!(shown.rows[0][0].as_deref(), Some("2000"));
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see live_postgres_processes_lists_and_kills_another_connection"]
async fn live_postgres_a_terminated_backend_reads_as_a_lost_connection() {
let watcher = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
let victim = Arc::new(live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await);
let sleeper = {
let victim = victim.clone();
tokio::spawn(async move {
victim
.run_query(
"SELECT count(*) FROM generate_series(1, 10000000000) AS zippa_terminate_test",
)
.await
})
};
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
watcher
.run_query(
"SELECT pg_terminate_backend(pid) FROM pg_stat_activity \
WHERE query LIKE '%zippa_terminate_test%' AND pid <> pg_backend_pid()",
)
.await
.expect("could not end the sleeping backend");
let error = sleeper
.await
.expect("the task itself should not panic")
.expect_err("the ended backend's query should fail");
assert_eq!(
crate::db::health::ConnectionTrouble::of(&error),
Some(crate::db::health::ConnectionTrouble::Lost),
"{error:#}"
);
watcher.close().await;
victim.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see live_postgres_processes_lists_and_kills_another_connection"]
async fn live_postgres_server_variables_lists_max_connections() {
let connection = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
let variables = connection
.server_variables()
.await
.expect("could not read the server variables");
let row = variables
.rows
.iter()
.find(|row| row[0].as_deref() == Some("max_connections"))
.expect("max_connections should be in pg_settings");
assert!(row[1].is_some(), "max_connections should have a value");
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see live_mysql_routines_carry_their_argument_types"]
async fn live_mysql_server_variables_lists_max_connections() {
let (connection, _pool) = live_mysql().await;
let variables = connection
.server_variables()
.await
.expect("could not read the server variables");
let row = variables
.rows
.iter()
.find(|row| {
row[0]
.as_deref()
.is_some_and(|name| name.eq_ignore_ascii_case("max_connections"))
})
.expect("max_connections should be in performance_schema.global_variables");
assert!(row[1].is_some(), "max_connections should have a value");
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see live_postgres_processes_lists_and_kills_another_connection"]
async fn live_postgres_query_digest_reports_availability() {
let connection = live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await;
let _ = connection.run_query("SELECT 1").await;
match connection
.query_digest()
.await
.expect("query_digest should not error, extension or not")
{
QueryDigest::Available(result) => {
assert!(
result.columns.contains(&"query".to_string()),
"the digest should carry pg_stat_statements' own columns"
);
}
QueryDigest::Unavailable(reason) => {
assert!(reason.contains("pg_stat_statements"));
}
}
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see live_mysql_routines_carry_their_argument_types"]
async fn live_mysql_query_digest_reports_availability() {
let (connection, _pool) = live_mysql().await;
let _ = connection.run_query("SELECT 1").await;
match connection
.query_digest()
.await
.expect("query_digest should not error, performance_schema on or off")
{
QueryDigest::Available(result) => {
assert!(
result.columns.contains(&"digest_text".to_string()),
"the digest should carry events_statements_summary_by_digest's own columns"
);
}
QueryDigest::Unavailable(reason) => {
assert!(reason.contains("performance_schema"));
}
}
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see the doc comment"]
async fn live_postgres_a_script_carries_on_past_a_skipped_failure() {
let connection = Arc::new(live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await);
let (run, failure) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
"CREATE TEMPORARY TABLE zippa_script_skip (id int PRIMARY KEY);\n\
INSERT INTO zippa_script_skip VALUES (1);\n\
INSERT INTO zippa_script_skip VALUES (1);\n\
INSERT INTO zippa_script_skip VALUES (2);\n\
SELECT COUNT(*) FROM zippa_script_skip;",
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
assert_eq!(failure.index, 2, "the duplicate key should fail");
let outcome = finished(run.resume(Decision::Skip).await);
assert_eq!(
outcome.results.last().map(|result| result.rows.clone()),
Some(vec![vec![Some("2".to_string())]]),
"the transaction should have carried on past the failure"
);
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see the doc comment"]
async fn live_mysql_a_script_of_row_changes_rolls_back() {
let (connection, pool) = live_mysql().await;
let connection = Arc::new(connection);
pool.execute("DROP TABLE IF EXISTS zippa_script_rollback")
.await
.expect("could not clear the fixture");
pool.execute("CREATE TABLE zippa_script_rollback (id int PRIMARY KEY) ENGINE=InnoDB")
.await
.expect("could not create the fixture");
let (run, _) = paused(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
"INSERT INTO zippa_script_rollback VALUES (1);\n\
INSERT INTO not_a_table VALUES (1);",
ScriptMode::Transaction,
OnFailure::Ask,
)
.await,
);
let outcome = finished(run.resume(Decision::Abort).await);
assert!(outcome.rolled_back);
let count = connection
.run_query("SELECT COUNT(*) FROM zippa_script_rollback")
.await
.expect("the count failed");
assert_eq!(count.rows, [[Some("0".to_string())]]);
pool.execute("DROP TABLE zippa_script_rollback").await.ok();
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see the doc comment"]
async fn live_mysql_an_autocommit_script_can_lock_tables() {
let (connection, pool) = live_mysql().await;
let connection = Arc::new(connection);
pool.execute("DROP TABLE IF EXISTS zippa_script_locks")
.await
.expect("could not clear the fixture");
pool.execute("CREATE TABLE zippa_script_locks (id int PRIMARY KEY) ENGINE=InnoDB")
.await
.expect("could not create the fixture");
let outcome = finished(
ScriptRun::start(
PinnedConnection::new(connection.clone()),
"LOCK TABLES zippa_script_locks WRITE;\n\
INSERT INTO zippa_script_locks VALUES (1);\n\
UNLOCK TABLES;\n\
SELECT COUNT(*) FROM zippa_script_locks;",
ScriptMode::Autocommit,
OnFailure::Ask,
)
.await,
);
assert!(outcome.failures.is_empty(), "{:?}", outcome.failures);
assert_eq!(
outcome.results.last().map(|result| result.rows.clone()),
Some(vec![vec![Some("1".to_string())]]),
);
pool.execute("DROP TABLE zippa_script_locks").await.ok();
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see the doc comment"]
async fn live_mysql_a_pinned_tab_holds_a_transaction_and_its_locks() {
let (connection, pool) = live_mysql().await;
let connection = Arc::new(connection);
pool.execute("DROP TABLE IF EXISTS zippa_pinned")
.await
.expect("could not clear the fixture");
pool.execute("CREATE TABLE zippa_pinned (id int PRIMARY KEY) ENGINE=InnoDB")
.await
.expect("could not create the fixture");
let pinned = PinnedConnection::new(connection.clone());
pinned.run_query("BEGIN").await.expect("begin");
assert_eq!(pinned.state(), TxnState::Open);
pinned
.run_query("INSERT INTO zippa_pinned VALUES (1)")
.await
.expect("insert");
let elsewhere = connection
.run_query("SELECT COUNT(*) FROM zippa_pinned")
.await
.expect("count from the pool");
assert_eq!(
elsewhere.rows,
[[Some("0".to_string())]],
"not committed yet"
);
pinned.run_query("ROLLBACK").await.expect("rollback");
assert_eq!(pinned.state(), TxnState::Idle);
pinned.run_query("USE app").await.expect("use");
pinned
.run_query("LOCK TABLES zippa_pinned WRITE")
.await
.expect("lock tables");
pinned
.run_query("INSERT INTO zippa_pinned VALUES (2)")
.await
.expect("insert under the lock, on the same connection");
pinned.run_query("UNLOCK TABLES").await.expect("unlock");
let count = connection
.run_query("SELECT COUNT(*) FROM zippa_pinned")
.await
.expect("the count failed");
assert_eq!(count.rows, [[Some("1".to_string())]]);
drop(pinned);
pool.execute("DROP TABLE zippa_pinned").await.ok();
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live Postgres server; see the doc comment"]
async fn live_postgres_a_pinned_tab_reports_a_failed_transaction() {
let connection = Arc::new(live_postgres("ZIPPA_TEST_POSTGRES_URL", "app").await);
let pinned = PinnedConnection::new(connection.clone());
pinned
.run_query("SET search_path = pg_catalog")
.await
.expect("set");
let path = pinned.run_query("SHOW search_path").await.expect("show");
assert_eq!(
path.rows,
[[Some("pg_catalog".to_string())]],
"SET lasts between runs"
);
pinned.run_query("BEGIN").await.expect("begin");
assert_eq!(pinned.state(), TxnState::Open);
pinned
.run_query("SELECT * FROM zippa_not_a_table")
.await
.expect_err("no such table");
assert_eq!(pinned.state(), TxnState::Failed);
pinned.run_query("ROLLBACK").await.expect("rollback");
assert_eq!(pinned.state(), TxnState::Idle);
drop(pinned);
connection.close().await;
}
#[tokio::test]
#[ignore = "needs a live MySQL server; see the doc comment"]
async fn live_mysql_a_mysqldump_file_imports_with_its_table_locks() {
use super::{ImportRequest, OnError};
let (connection, pool) = live_mysql().await;
pool.execute("DROP TABLE IF EXISTS zippa_import_locks")
.await
.expect("could not clear the fixture");
let path = env::temp_dir().join(format!("zippa-dump-{}.sql", Uuid::new_v4()));
std::fs::write(
&path,
"CREATE TABLE `zippa_import_locks` (`id` int NOT NULL, PRIMARY KEY (`id`)) ENGINE=InnoDB;\n\
LOCK TABLES `zippa_import_locks` WRITE;\n\
/*!40000 ALTER TABLE `zippa_import_locks` DISABLE KEYS */;\n\
INSERT INTO `zippa_import_locks` VALUES (1),(2);\n\
/*!40000 ALTER TABLE `zippa_import_locks` ENABLE KEYS */;\n\
UNLOCK TABLES;\n",
)
.expect("could not write the dump");
let (sender, _receiver) = tokio::sync::mpsc::unbounded_channel();
let summary = connection
.import_dump(
ImportRequest {
path: path.clone(),
on_error: OnError::Stop,
},
sender,
)
.await;
std::fs::remove_file(&path).ok();
let summary = summary.expect("the dump should import");
assert!(summary.errors.is_empty(), "{:?}", summary.errors);
let count = connection
.run_query("SELECT COUNT(*) FROM zippa_import_locks")
.await
.expect("the count failed");
assert_eq!(count.rows, [[Some("2".to_string())]]);
pool.execute("DROP TABLE zippa_import_locks").await.ok();
connection.close().await;
}
pub(crate) async fn live_postgres(var: &str, database: &str) -> Connection {
live_postgres_with(var, database, |_| {}).await
}
async fn live_postgres_with(
var: &str,
database: &str,
adjust: impl FnOnce(&mut ConnectionConfig),
) -> Connection {
let (mut config, password) = live_postgres_config(var, database);
adjust(&mut config);
Connection::open(config, Some(password))
.await
.expect("could not open the live Postgres connection")
}
fn live_postgres_config(var: &str, database: &str) -> (ConnectionConfig, String) {
let url = env::var(var)
.unwrap_or_else(|_| format!("postgres://postgres:secret@127.0.0.1:5433/{database}"));
let rest = url.strip_prefix("postgres://").expect("a postgres:// URL");
let (credentials, rest) = rest.split_once('@').expect("user:pass@host:port/db");
let (username, password) = credentials.split_once(':').expect("user:pass");
let (authority, database) = rest.split_once('/').expect("host:port/db");
let (host, port) = authority.split_once(':').expect("host:port");
let port: u16 = port.parse().expect("a port");
let config = ConnectionConfig {
host: host.to_string(),
port,
username: username.to_string(),
database: database.to_string(),
..ConnectionConfig::new(Engine::Postgres)
};
(config, password.to_string())
}
fn live_mysql_config() -> (ConnectionConfig, String) {
let url = env::var("ZIPPA_TEST_MYSQL_URL")
.unwrap_or_else(|_| "mysql://root:secret@127.0.0.1:3307/app".to_string());
let rest = url.strip_prefix("mysql://").expect("a mysql:// URL");
let (credentials, rest) = rest.split_once('@').expect("user:pass@host:port/db");
let (username, password) = credentials.split_once(':').expect("user:pass");
let (authority, database) = rest.split_once('/').expect("host:port/db");
let (host, port) = authority.split_once(':').expect("host:port");
let port: u16 = port.parse().expect("a port");
let config = ConnectionConfig {
host: host.to_string(),
port,
username: username.to_string(),
database: database.to_string(),
..ConnectionConfig::new(Engine::MySql)
};
(config, password.to_string())
}
async fn live_mysql() -> (Connection, sqlx::MySqlPool) {
live_mysql_with(|_| {}).await
}
async fn live_mysql_with(
adjust: impl FnOnce(&mut ConnectionConfig),
) -> (Connection, sqlx::MySqlPool) {
use sqlx::mysql::{MySqlConnectOptions, MySqlPoolOptions};
let (mut config, password) = live_mysql_config();
let options = MySqlConnectOptions::new()
.host(&config.host)
.port(config.port)
.username(&config.username)
.password(&password)
.database(&config.database);
let pool = MySqlPoolOptions::new()
.connect_with(options)
.await
.expect("could not open the live MySQL server");
adjust(&mut config);
let connection = Connection::open(config, Some(password))
.await
.expect("could not open the live MySQL connection");
(connection, pool)
}