use std::fmt;
use std::io;
use std::path::{Path, PathBuf};
use turso::Connection;
use crate::store::lock::{self, DataDirLock, LockError};
use crate::store::schema::{self, SchemaError};
const STATE_DIR: &str = "state";
const DATABASE_FILE: &str = "firewall.db";
pub struct Opened {
pub lock: DataDirLock,
pub connection: Result<Connection, RecoveryError>,
}
pub async fn open_and_recover(data_dir: &Path) -> Result<Opened, StartupError> {
let state_dir = data_dir.join(STATE_DIR);
std::fs::create_dir_all(&state_dir).map_err(|source| StartupError::Directory {
path: state_dir.clone(),
source,
})?;
let lock = lock::acquire(data_dir).map_err(StartupError::Lock)?;
let connection = open_database(&state_dir.join(DATABASE_FILE)).await;
if let Err(err) = &connection {
tracing::error!(
data_dir = %data_dir.display(),
error = %err,
"the database could not be recovered: readiness stays false and every package \
request is refused; the database is left untouched"
);
}
Ok(Opened { lock, connection })
}
async fn open_database(path: &Path) -> Result<Connection, RecoveryError> {
let log = log_path(path);
let log_bytes = std::fs::metadata(&log).map(|meta| meta.len()).unwrap_or(0);
let database_bytes = std::fs::metadata(path).map(|meta| meta.len()).unwrap_or(0);
if log_bytes > 0 && database_bytes == 0 {
return Err(RecoveryError::OrphanedLog {
database: path.to_path_buf(),
log,
log_bytes,
});
}
let path_text = path
.to_str()
.ok_or_else(|| RecoveryError::Path(path.to_path_buf()))?;
let database = turso::Builder::new_local(path_text)
.build()
.await
.map_err(RecoveryError::Database)?;
let mut connection = database.connect().map_err(RecoveryError::Database)?;
connection
.pragma_update("journal_mode", "WAL")
.await
.map_err(RecoveryError::Database)?;
let journal_mode = read_pragma(&connection, "journal_mode").await?;
let journal_mode = journal_mode
.as_text()
.cloned()
.unwrap_or_else(|| format!("{journal_mode:?}"));
if !journal_mode.eq_ignore_ascii_case("wal") {
return Err(RecoveryError::Pragma {
pragma: "journal_mode",
expected: "wal".to_owned(),
found: journal_mode,
});
}
connection
.pragma_update("synchronous", "FULL")
.await
.map_err(RecoveryError::Database)?;
let synchronous = read_pragma(&connection, "synchronous").await?;
const SYNCHRONOUS_FULL: i64 = 2;
if synchronous.as_integer().copied() != Some(SYNCHRONOUS_FULL) {
return Err(RecoveryError::Pragma {
pragma: "synchronous",
expected: format!("{SYNCHRONOUS_FULL} (full)"),
found: format!("{synchronous:?}"),
});
}
schema::ensure(&mut connection)
.await
.map_err(RecoveryError::Schema)?;
Ok(connection)
}
fn log_path(database: &Path) -> PathBuf {
let mut name = database.as_os_str().to_owned();
name.push("-wal");
PathBuf::from(name)
}
async fn read_pragma(
connection: &Connection,
pragma: &'static str,
) -> Result<turso::Value, RecoveryError> {
let mut value = None;
connection
.pragma_query(pragma, |row| {
if value.is_none() {
value = row.get_value(0).ok();
}
Ok(())
})
.await
.map_err(RecoveryError::Database)?;
value.ok_or(RecoveryError::Pragma {
pragma,
expected: "a value".to_owned(),
found: "nothing".to_owned(),
})
}
#[derive(Debug)]
pub enum StartupError {
Directory { path: PathBuf, source: io::Error },
Lock(LockError),
}
impl fmt::Display for StartupError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
StartupError::Directory { path, source } => {
write!(f, "cannot create {}: {source}", path.display())
}
StartupError::Lock(err) => write!(f, "{err}"),
}
}
}
impl std::error::Error for StartupError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
StartupError::Directory { source, .. } => Some(source),
StartupError::Lock(err) => Some(err),
}
}
}
#[derive(Debug)]
pub enum RecoveryError {
Path(PathBuf),
OrphanedLog {
database: PathBuf,
log: PathBuf,
log_bytes: u64,
},
Database(turso::Error),
Pragma {
pragma: &'static str,
expected: String,
found: String,
},
Schema(SchemaError),
}
impl fmt::Display for RecoveryError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
RecoveryError::Path(path) => {
write!(f, "the database path {} is not valid UTF-8", path.display())
}
RecoveryError::OrphanedLog {
database,
log,
log_bytes,
} => write!(
f,
"{} is missing or empty while its write-ahead log {} still holds {log_bytes} \
bytes; refusing to open, because opening would overwrite the log and lose \
what is still in it. Restore the pair together, or move both aside to start \
fresh deliberately",
database.display(),
log.display()
),
RecoveryError::Database(err) => write!(f, "{err}"),
RecoveryError::Pragma {
pragma,
expected,
found,
} => write!(f, "{pragma} is {found}, not {expected}"),
RecoveryError::Schema(err) => write!(f, "{err}"),
}
}
}
impl std::error::Error for RecoveryError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
RecoveryError::Database(err) => Some(err),
RecoveryError::Schema(err) => Some(err),
RecoveryError::Path(_)
| RecoveryError::OrphanedLog { .. }
| RecoveryError::Pragma { .. } => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn the_connection_is_in_wal_mode_at_synchronous_full() {
let dir = tempfile::tempdir().expect("a temporary data directory");
let opened = open_and_recover(dir.path())
.await
.expect("the data directory opens");
let connection = opened.connection.expect("a fresh database recovers");
let journal_mode = read_pragma(&connection, "journal_mode")
.await
.expect("journal_mode is readable");
assert_eq!(
journal_mode.as_text().map(String::as_str),
Some("wal"),
"recovery after a crash is a property of the mode, not of the statement \
that asked for it"
);
let synchronous = read_pragma(&connection, "synchronous")
.await
.expect("synchronous is readable");
assert_eq!(
synchronous.as_integer().copied(),
Some(2),
"2 is FULL: a commit reaches the disk before it is reported committed, \
which is what lets the poller publish only after committing"
);
}
}