use std::collections::HashSet;
use std::fmt::Write as _;
use std::time::{SystemTime, UNIX_EPOCH};
use rusqlite::{Connection, params};
use sha2::{Digest, Sha256};
use super::embedded::{INIT_SQL, MIGRATIONS};
use crate::error::{Error, Result};
const TOLERATED_ERROR_MARKERS: [&str; 2] = ["duplicate column", "already exists"];
pub fn apply_all(conn: &Connection) -> Result<u32> {
if !schema_versions_exists(conn)? {
conn.execute_batch(INIT_SQL)
.map_err(|source| Error::Migration {
version: 1,
message: format!("0001_init.sql: {source}"),
})?;
}
let supported = MIGRATIONS.last().map_or(1, |migration| migration.version);
let found = read_current_version(conn)?;
if found > supported {
return Err(Error::SchemaAhead {
found: i64::from(found),
supported: i64::from(supported),
});
}
let mut known = known_versions(conn)?;
for migration in MIGRATIONS {
if known.contains(&migration.version) {
continue;
}
if let Err(source) = conn.execute_batch(migration.sql) {
let message = source.to_string();
let tolerated = TOLERATED_ERROR_MARKERS
.iter()
.any(|marker| message.to_ascii_lowercase().contains(marker));
if !tolerated {
return Err(Error::Migration {
version: i64::from(migration.version),
message: format!("{}: {message}", migration.filename),
});
}
}
record_version(
conn,
migration.version,
checksum_hex(migration.sql.as_bytes()),
)?;
known.insert(migration.version);
}
read_current_version(conn)
}
fn read_current_version(conn: &Connection) -> Result<u32> {
let max: i64 = conn.query_row(
"SELECT COALESCE(MAX(version), 0) FROM schema_versions",
[],
|row| row.get(0),
)?;
u32::try_from(max)
.map_err(|_| Error::unknown(format!("schema_versions.version out of range: {max}")))
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn dump_sqlite_master(conn: &Connection) -> Result<String> {
let mut stmt =
conn.prepare("SELECT type, name, tbl_name, sql FROM sqlite_master ORDER BY type, name")?;
let rows = stmt.query_map([], |row| {
let kind: String = row.get(0)?;
let name: String = row.get(1)?;
let tbl_name: String = row.get(2)?;
let sql: Option<String> = row.get(3)?;
Ok((kind, name, tbl_name, sql))
})?;
let mut out = String::new();
for row in rows {
let (kind, name, tbl_name, sql) = row?;
let sql = sql.unwrap_or_default().replace('\n', "\\n");
writeln!(out, "{kind}\t{name}\t{tbl_name}\t{sql}")
.expect("writing to a String cannot fail");
}
Ok(out)
}
fn schema_versions_exists(conn: &Connection) -> Result<bool> {
let exists: bool = conn.query_row(
"SELECT EXISTS(SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'schema_versions')",
[],
|row| row.get(0),
)?;
Ok(exists)
}
fn known_versions(conn: &Connection) -> Result<HashSet<u32>> {
if !schema_versions_exists(conn)? {
return Ok(HashSet::new());
}
let mut stmt = conn.prepare("SELECT version FROM schema_versions")?;
let rows = stmt.query_map([], |row| row.get::<_, i64>(0))?;
let mut set = HashSet::new();
for row in rows {
let version = row?;
set.insert(
u32::try_from(version).map_err(|_| {
Error::unknown(format!("negative schema_versions.version: {version}"))
})?,
);
}
Ok(set)
}
fn record_version(conn: &Connection, version: u32, checksum: String) -> Result<()> {
let applied_at = i64::try_from(
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0),
)
.unwrap_or(i64::MAX);
conn.execute(
"INSERT OR IGNORE INTO schema_versions (version, applied_at, checksum) VALUES (?1, ?2, ?3)",
params![i64::from(version), applied_at, checksum],
)?;
Ok(())
}
fn checksum_hex(bytes: &[u8]) -> String {
let digest = Sha256::digest(bytes);
let mut out = String::with_capacity(digest.len() * 2);
for byte in digest.as_slice() {
write!(out, "{byte:02x}").expect("writing to a String cannot fail");
}
out
}