pub(crate) mod reader;
pub(crate) mod splitter;
pub use reader::Compression;
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
use anyhow::{Result, bail};
use tokio::sync::mpsc::UnboundedSender;
use super::config::Engine;
use super::dedicated::Dedicated;
use splitter::Chunk;
pub use splitter::Dialect;
pub(crate) const HEAD: usize = 64 * 1024;
const MAX_ERRORS: usize = 200;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum OnError {
#[default]
Stop,
Rollback,
Continue,
}
impl OnError {
pub const ALL: [OnError; 3] = [OnError::Stop, OnError::Rollback, OnError::Continue];
pub fn label(self) -> &'static str {
match self {
OnError::Stop => "Stop at the first error",
OnError::Rollback => "Roll back everything",
OnError::Continue => "Continue and log errors",
}
}
pub fn description(self) -> &'static str {
match self {
OnError::Stop => "Keeps what ran before the failure",
OnError::Rollback => "Undoes the whole import on any failure",
OnError::Continue => "Skips failed statements and reports them at the end",
}
}
}
#[derive(Debug, Clone)]
pub struct ImportRequest {
pub path: PathBuf,
pub on_error: OnError,
}
#[derive(Debug, Clone, Default)]
pub struct ImportProgress {
pub bytes: u64,
pub total: u64,
pub statements: u64,
pub errors: u64,
pub line: u64,
}
impl ImportProgress {
pub fn percent(&self) -> f32 {
if self.total == 0 {
return 0.0;
}
(self.bytes as f64 / self.total as f64 * 100.0).clamp(0.0, 100.0) as f32
}
}
#[derive(Debug, Clone)]
pub struct ImportError {
pub line: u64,
pub statement: String,
pub message: String,
}
#[derive(Debug, Clone, Default)]
pub struct ImportSummary {
pub statements: u64,
pub errors: Vec<ImportError>,
pub dropped_errors: u64,
pub elapsed: Duration,
pub rolled_back: bool,
pub dump_transaction: bool,
pub mysql_rollback_fallback: bool,
pub total_bytes: u64,
}
impl ImportSummary {
pub fn summary(&self) -> String {
let unit = if self.statements == 1 {
"statement"
} else {
"statements"
};
let mut text = format!(
"Imported {} {unit} in {}",
self.statements,
elapsed(self.elapsed)
);
if self.rolled_back {
text.push_str(" · rolled back");
}
let failures = self.errors.len() as u64 + self.dropped_errors;
if failures > 0 {
let unit = if failures == 1 { "error" } else { "errors" };
text.push_str(&format!(" · {failures} {unit}"));
if self.dropped_errors > 0 {
text.push_str(&format!(" (first {} kept)", self.errors.len()));
}
}
text
}
}
pub(crate) async fn run(
mut session: Dedicated,
dialect: Dialect,
request: &ImportRequest,
sender: UnboundedSender<ImportProgress>,
) -> Result<ImportSummary> {
let started = Instant::now();
let engine = session.engine();
let opened = reader::open(&request.path)?;
let mut progress = Progress {
sender,
consumed: opened.consumed.clone(),
total: opened.total_bytes,
last: Instant::now(),
};
progress.emit(0, 0, 0);
let mut summary = ImportSummary {
total_bytes: opened.total_bytes,
..ImportSummary::default()
};
let policy = match (request.on_error, engine) {
(OnError::Rollback, Engine::MySql) => {
summary.mysql_rollback_fallback = true;
OnError::Stop
}
(policy, _) => policy,
};
let mut mysql_settings = false;
if engine == Engine::MySql {
let head = reader::inspect(&request.path, HEAD)
.map(|inspection| inspection.head)
.unwrap_or_default();
if !mentions(&head, "FOREIGN_KEY_CHECKS") && !mentions(&head, "UNIQUE_CHECKS") {
session.execute("SET FOREIGN_KEY_CHECKS=0").await?;
session.execute("SET UNIQUE_CHECKS=0").await?;
mysql_settings = true;
}
}
let mut chunks = splitter::new(opened.reader, dialect).peekable();
let mut errors = Vec::new();
let mut dropped_errors = 0u64;
let mut began = false;
let mut rolled_back = false;
let mut stopped = false;
let mut last_line = 0u64;
while !stopped {
let Some(chunk) = chunks.next() else {
break;
};
let chunk = chunk?;
let (statement, copying) = match chunk {
Chunk::Statement(statement) => (statement, false),
Chunk::Copy(statement) => (statement, true),
Chunk::CopyData(_) => continue,
};
if transaction_control(&statement.sql) {
summary.dump_transaction = true;
continue;
}
if policy == OnError::Rollback && !began && !is_pragma(&statement.sql) {
session.execute("BEGIN").await?;
began = true;
}
summary.statements += 1;
last_line = statement.line;
progress.tick(summary.statements, errors.len() as u64, statement.line);
let result = if copying {
copy_block(
&mut session,
&statement.sql,
&mut chunks,
&mut progress,
summary.statements,
errors.len() as u64,
statement.line,
)
.await
} else {
session.execute(&statement.sql).await
};
if let Err(error) = result {
let entry = ImportError {
line: statement.line,
statement: excerpt(&statement.sql),
message: format!("{error:#}"),
};
match policy {
OnError::Stop => {
errors.push(entry);
stopped = true;
}
OnError::Rollback => {
session.execute("ROLLBACK").await.ok();
rolled_back = true;
errors.push(entry);
stopped = true;
}
OnError::Continue => {
if errors.len() < MAX_ERRORS {
errors.push(entry);
} else {
dropped_errors += 1;
}
}
}
}
}
if policy == OnError::Rollback && began && !rolled_back {
session.execute("COMMIT").await?;
}
if mysql_settings {
session.execute("SET FOREIGN_KEY_CHECKS=1").await.ok();
session.execute("SET UNIQUE_CHECKS=1").await.ok();
}
summary.errors = errors;
summary.dropped_errors = dropped_errors;
summary.rolled_back = rolled_back;
summary.elapsed = started.elapsed();
progress.emit(summary.statements, summary.errors.len() as u64, last_line);
Ok(summary)
}
async fn copy_block<I>(
session: &mut Dedicated,
statement: &str,
chunks: &mut std::iter::Peekable<I>,
progress: &mut Progress,
statements: u64,
errors: u64,
line: u64,
) -> Result<()>
where
I: Iterator<Item = Result<Chunk>>,
{
match session {
Dedicated::Postgres(connection, _) => {
let mut writer = connection.copy_in_raw(statement).await?;
while matches!(chunks.peek(), Some(Ok(Chunk::CopyData(_)))) {
match chunks.next() {
Some(Ok(Chunk::CopyData(data))) => {
writer.send(data).await?;
progress.tick(statements, errors, line);
}
_ => break,
}
}
writer.finish().await?;
Ok(())
}
_ => bail!("COPY data can only be read on PostgreSQL"),
}
}
fn is_pragma(sql: &str) -> bool {
leading_word(sql).as_deref() == Some("PRAGMA")
}
fn transaction_control(sql: &str) -> bool {
matches!(
leading_word(sql).as_deref(),
Some("BEGIN" | "START" | "COMMIT" | "END" | "ROLLBACK" | "ABORT")
)
}
fn leading_word(sql: &str) -> Option<String> {
let characters: Vec<char> = sql.chars().collect();
let mut index = 0;
loop {
while index < characters.len() && characters[index].is_whitespace() {
index += 1;
}
match (characters.get(index), characters.get(index + 1)) {
(Some('-'), Some('-')) | (Some('#'), _) => {
while index < characters.len() && characters[index] != '\n' {
index += 1;
}
}
(Some('/'), Some('*')) => {
index += 2;
while index < characters.len() {
if characters[index] == '*' && characters.get(index + 1) == Some(&'/') {
index += 2;
break;
}
index += 1;
}
}
_ => break,
}
}
let mut word = String::new();
while let Some(character) = characters.get(index) {
if character.is_alphabetic() || *character == '_' {
word.push(*character);
index += 1;
} else {
break;
}
}
(!word.is_empty()).then(|| word.to_ascii_uppercase())
}
fn mentions(head: &str, needle: &str) -> bool {
head.to_ascii_uppercase()
.contains(&needle.to_ascii_uppercase())
}
pub(crate) fn excerpt(sql: &str) -> String {
let flat = sql.split_whitespace().collect::<Vec<_>>().join(" ");
if flat.chars().count() > 120 {
let mut shortened: String = flat.chars().take(119).collect();
shortened.push('…');
shortened
} else {
flat
}
}
fn elapsed(duration: Duration) -> String {
if duration.as_secs_f64() >= 1.0 {
format!("{:.1} s", duration.as_secs_f64())
} else {
format!("{} ms", duration.as_millis())
}
}
struct Progress {
sender: UnboundedSender<ImportProgress>,
consumed: Arc<AtomicU64>,
total: u64,
last: Instant,
}
impl Progress {
fn tick(&mut self, statements: u64, errors: u64, line: u64) {
if self.last.elapsed() < Duration::from_millis(100) {
return;
}
self.emit(statements, errors, line);
}
fn emit(&mut self, statements: u64, errors: u64, line: u64) {
self.last = Instant::now();
let _ = self.sender.send(ImportProgress {
bytes: self.consumed.load(Ordering::Relaxed),
total: self.total,
statements,
errors,
line,
});
}
}
pub(crate) struct Preflight {
pub compression: Compression,
pub total_bytes: u64,
pub dialect_mismatch: Option<String>,
pub destructive: usize,
pub error: Option<String>,
}
pub(crate) fn preflight(path: &std::path::Path, engine: Engine) -> Preflight {
match reader::inspect(path, HEAD) {
Ok(inspection) => {
let hint = dialect_hint(&inspection.head);
let dialect_mismatch = match hint {
Some(hint) if hint != engine => Some(format!(
"This looks like a {} dump, but the connection is {}.",
hint.label(),
engine.label()
)),
_ => None,
};
Preflight {
compression: inspection.compression,
total_bytes: inspection.total_bytes,
dialect_mismatch,
destructive: destructive_count(&inspection.head, engine),
error: None,
}
}
Err(error) => Preflight {
compression: Compression::None,
total_bytes: 0,
dialect_mismatch: None,
destructive: 0,
error: Some(format!("{error:#}")),
},
}
}
fn dialect_hint(head: &str) -> Option<Engine> {
let upper = head.to_ascii_uppercase();
if upper.contains("POSTGRESQL DATABASE DUMP") || upper.contains("PG_DUMP") {
return Some(Engine::Postgres);
}
if upper.contains("MYSQL DUMP") || upper.contains("MARIADB DUMP") || upper.contains("MYSQLDUMP")
{
return Some(Engine::MySql);
}
let trimmed = head.trim_start();
if trimmed.starts_with("PRAGMA") || trimmed.starts_with("BEGIN TRANSACTION") {
return Some(Engine::Sqlite);
}
None
}
fn destructive_count(head: &str, engine: Engine) -> usize {
super::statement::split(head, engine)
.iter()
.filter(|statement| {
let upper = statement.text.to_ascii_uppercase();
match upper.split_whitespace().next().unwrap_or_default() {
"DROP" | "TRUNCATE" => true,
"DELETE" => !contains_word(&upper, "WHERE"),
_ => false,
}
})
.count()
}
fn contains_word(haystack: &str, needle: &str) -> bool {
haystack.match_indices(needle).any(|(start, _)| {
let is_word_char = |c: char| c.is_alphanumeric() || c == '_';
let before = haystack[..start].chars().next_back();
let after = haystack[start + needle.len()..].chars().next();
!before.is_some_and(is_word_char) && !after.is_some_and(is_word_char)
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn transaction_control_is_recognised_behind_comments() {
assert!(transaction_control("BEGIN"));
assert!(transaction_control(" start transaction;"));
assert!(transaction_control("-- a note\nCOMMIT"));
assert!(transaction_control("/* c */ END"));
assert!(!transaction_control("SELECT 1"));
assert!(!transaction_control("DROP TABLE t"));
}
#[test]
fn the_destructive_count_names_drop_truncate_and_unbounded_delete() {
let head =
"DROP TABLE a;\nTRUNCATE b;\nDELETE FROM c WHERE id = 1;\nDELETE FROM d;\nSELECT 1;";
assert_eq!(destructive_count(head, Engine::MySql), 3);
}
#[test]
fn a_where_clause_with_no_space_before_it_still_bounds_the_delete() {
let head = "DELETE FROM c WHERE(id > 0);\nDELETE FROM d WHERE\nid = 1;";
assert_eq!(destructive_count(head, Engine::MySql), 0);
}
#[test]
fn an_identifier_starting_with_where_is_not_mistaken_for_the_keyword() {
let head = "DELETE FROM wherefore;";
assert_eq!(destructive_count(head, Engine::MySql), 1);
}
#[test]
fn a_header_tells_the_engine_apart() {
assert_eq!(
dialect_hint("--\n-- PostgreSQL database dump\n--"),
Some(Engine::Postgres)
);
assert_eq!(
dialect_hint("-- MySQL dump 10.13 Distrib 8.0"),
Some(Engine::MySql)
);
assert_eq!(
dialect_hint("PRAGMA foreign_keys=OFF;"),
Some(Engine::Sqlite)
);
assert_eq!(dialect_hint("SELECT 1;"), None);
}
#[test]
fn an_excerpt_is_flattened_and_shortened() {
assert_eq!(excerpt("SELECT\n 1"), "SELECT 1");
let long = "a".repeat(200);
let short = excerpt(&long);
assert!(short.chars().count() <= 120);
assert!(short.ends_with('…'));
}
#[test]
fn progress_has_a_percentage() {
let progress = ImportProgress {
bytes: 50,
total: 200,
..ImportProgress::default()
};
assert_eq!(progress.percent(), 25.0);
assert_eq!(ImportProgress::default().percent(), 0.0);
}
}
#[cfg(test)]
mod database_tests {
use std::io::Write as _;
use std::path::{Path, PathBuf};
use uuid::Uuid;
use super::*;
use crate::db::tests::TempDatabase;
use crate::db::{Connection, ConnectionConfig, SafetyMode};
struct Dump {
path: PathBuf,
}
impl Dump {
fn text(contents: &str) -> Self {
let path = std::env::temp_dir().join(format!("zippa-dump-{}.sql", Uuid::new_v4()));
std::fs::write(&path, contents).expect("could not write the test dump");
Self { path }
}
fn gzip(contents: &str) -> Self {
let mut encoder =
flate2::write::GzEncoder::new(Vec::new(), flate2::Compression::default());
encoder
.write_all(contents.as_bytes())
.expect("could not compress the test dump");
let bytes = encoder.finish().expect("could not compress the test dump");
let path = std::env::temp_dir().join(format!("zippa-dump-{}.sql.gz", Uuid::new_v4()));
std::fs::write(&path, bytes).expect("could not write the test dump");
Self { path }
}
fn zstd(contents: &str) -> Self {
let bytes =
zstd::encode_all(contents.as_bytes(), 0).expect("could not compress the test dump");
let path = std::env::temp_dir().join(format!("zippa-dump-{}.sql.zst", Uuid::new_v4()));
std::fs::write(&path, bytes).expect("could not write the test dump");
Self { path }
}
}
impl Drop for Dump {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.path);
}
}
async fn import(
connection: &Connection,
path: &Path,
on_error: OnError,
) -> Result<ImportSummary> {
let (sender, _receiver) = tokio::sync::mpsc::unbounded_channel();
connection
.import_dump(
ImportRequest {
path: path.to_path_buf(),
on_error,
},
sender,
)
.await
}
async fn count(connection: &Connection, table: &str) -> Option<i64> {
let result = connection
.run_query(&format!("SELECT count(*) FROM {table}"))
.await
.ok()?;
result.rows.first()?.first()?.as_ref()?.parse().ok()
}
const BROKEN: &str = "CREATE TABLE t (id int);\nINSERT INTO t VALUES (1);\nTHIS IS NOT SQL;\nINSERT INTO t VALUES (2);\n";
#[tokio::test]
async fn a_dump_of_statements_is_imported() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::text(
"CREATE TABLE t (id int);\nINSERT INTO t VALUES (1);\nINSERT INTO t VALUES (2);\n",
);
let summary = import(&connection, &dump.path, OnError::Stop)
.await
.expect("the dump should import");
assert_eq!(summary.statements, 3);
assert!(summary.errors.is_empty());
assert!(!summary.rolled_back);
assert_eq!(count(&connection, "t").await, Some(2));
connection.close().await;
}
#[tokio::test]
async fn stopping_keeps_what_ran_before_the_failure() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::text(BROKEN);
let summary = import(&connection, &dump.path, OnError::Stop)
.await
.expect("the import should finish");
assert_eq!(summary.errors.len(), 1);
assert_eq!(summary.errors[0].line, 3, "the failing statement's line");
assert!(!summary.rolled_back);
assert_eq!(count(&connection, "t").await, Some(1));
connection.close().await;
}
#[tokio::test]
async fn rolling_back_undoes_the_whole_import() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::text(BROKEN);
let summary = import(&connection, &dump.path, OnError::Rollback)
.await
.expect("the import should finish");
assert!(summary.rolled_back);
assert_eq!(summary.errors.len(), 1);
assert_eq!(
count(&connection, "t").await,
None,
"the table created earlier in the dump should be gone"
);
connection.close().await;
}
#[tokio::test]
async fn a_pragma_runs_outside_the_wrapping_transaction() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::text(
"PRAGMA foreign_keys=OFF;\n\
CREATE TABLE parent (id int primary key);\n\
CREATE TABLE child (id int primary key, parent_id int references parent(id));\n\
INSERT INTO child VALUES (1, 99);\n",
);
let summary = import(&connection, &dump.path, OnError::Rollback)
.await
.expect("the dump should import");
assert!(summary.errors.is_empty(), "{:?}", summary.errors);
assert_eq!(count(&connection, "child").await, Some(1));
connection.close().await;
}
#[tokio::test]
async fn continuing_past_a_failure_finishes_the_dump() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::text(BROKEN);
let summary = import(&connection, &dump.path, OnError::Continue)
.await
.expect("the import should finish");
assert_eq!(summary.errors.len(), 1);
assert_eq!(summary.statements, 4);
assert_eq!(count(&connection, "t").await, Some(2));
connection.close().await;
}
#[tokio::test]
async fn a_read_only_connection_refuses_a_dump() {
let database = TempDatabase::new().await;
let config = ConnectionConfig {
safety: SafetyMode::ReadOnly,
..database.config()
};
let connection = Connection::open(config, None).await.unwrap();
let dump = Dump::text("CREATE TABLE t (id int);\n");
let error = import(&connection, &dump.path, OnError::Stop)
.await
.expect_err("a read-only connection should refuse");
assert!(format!("{error:#}").contains("read-only"), "{error:#}");
connection.close().await;
}
#[tokio::test]
async fn a_dump_that_wraps_itself_is_not_nested() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::text(
"BEGIN TRANSACTION;\nCREATE TABLE t (id int);\nINSERT INTO t VALUES (1);\nCOMMIT;\n",
);
let summary = import(&connection, &dump.path, OnError::Stop)
.await
.expect("the dump should import");
assert!(summary.dump_transaction, "the dump's transaction was noted");
assert_eq!(summary.statements, 2, "BEGIN and COMMIT are ours to manage");
assert_eq!(count(&connection, "t").await, Some(1));
connection.close().await;
}
#[tokio::test]
async fn a_gzipped_dump_is_decompressed_on_the_way_in() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::gzip("CREATE TABLE t (id int);\nINSERT INTO t VALUES (7);\n");
let summary = import(&connection, &dump.path, OnError::Stop)
.await
.expect("the dump should import");
assert_eq!(summary.statements, 2);
assert_eq!(count(&connection, "t").await, Some(1));
connection.close().await;
}
#[tokio::test]
async fn a_zstd_dump_is_decompressed_on_the_way_in() {
let database = TempDatabase::new().await;
let connection = Connection::open(database.config(), None).await.unwrap();
let dump = Dump::zstd("CREATE TABLE t (id int);\nINSERT INTO t VALUES (7);\n");
let summary = import(&connection, &dump.path, OnError::Stop)
.await
.expect("the dump should import");
assert_eq!(summary.statements, 2);
assert_eq!(count(&connection, "t").await, Some(1));
connection.close().await;
}
}