use std::fs::File;
use std::io::Read;
use std::path::Path;
use crate::sqlite::compressed_vfs::{FLAG_ENCRYPTED, HEADER_FLAGS_OFFSET, LEGACY_MAGIC, MAGIC};
const SQLITE_MAGIC: &[u8; 16] = b"SQLite format 3\0";
const DETECT_PREFIX_LEN: usize = 16;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DatabaseFileFormat {
Missing,
PlainSQLite,
CompressedContainer { encrypted: bool },
Unrecognized,
}
impl DatabaseFileFormat {
#[must_use]
pub fn requires_key(self) -> bool {
matches!(
self,
Self::CompressedContainer { encrypted: true } | Self::Unrecognized
)
}
}
pub fn detect_database_file_format(path: &Path) -> std::io::Result<DatabaseFileFormat> {
let mut file = match File::open(path) {
Ok(file) => file,
Err(err) if err.kind() == std::io::ErrorKind::NotFound => {
return Ok(DatabaseFileFormat::Missing);
}
Err(err) => return Err(err),
};
let mut prefix = [0_u8; DETECT_PREFIX_LEN];
let mut read = 0;
while read < prefix.len() {
let n = file.read(&mut prefix[read..])?;
if n == 0 {
break;
}
read += n;
}
if read == 0 {
return Ok(DatabaseFileFormat::Missing);
}
if read < DETECT_PREFIX_LEN {
return Ok(DatabaseFileFormat::Unrecognized);
}
if &prefix == SQLITE_MAGIC {
return Ok(DatabaseFileFormat::PlainSQLite);
}
if &prefix[..MAGIC.len()] == MAGIC || &prefix[..LEGACY_MAGIC.len()] == LEGACY_MAGIC {
let flags = u32::from_le_bytes([
prefix[HEADER_FLAGS_OFFSET],
prefix[HEADER_FLAGS_OFFSET + 1],
prefix[HEADER_FLAGS_OFFSET + 2],
prefix[HEADER_FLAGS_OFFSET + 3],
]);
return Ok(DatabaseFileFormat::CompressedContainer {
encrypted: flags & FLAG_ENCRYPTED != 0,
});
}
Ok(DatabaseFileFormat::Unrecognized)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sqlite::compressed_vfs::SQLiteCompressionOptions;
use crate::sqlite::connection::ManagedConnection;
fn temp_dir() -> tempfile::TempDir {
tempfile::tempdir().expect("tempdir")
}
#[test]
fn missing_file_detects_as_missing() {
let dir = temp_dir();
let path = dir.path().join("absent.db");
assert_eq!(
detect_database_file_format(&path).unwrap(),
DatabaseFileFormat::Missing
);
}
#[test]
fn empty_file_detects_as_missing() {
let dir = temp_dir();
let path = dir.path().join("empty.db");
std::fs::write(&path, b"").unwrap();
assert_eq!(
detect_database_file_format(&path).unwrap(),
DatabaseFileFormat::Missing
);
}
#[test]
fn plaintext_sqlite_detects_as_plain() {
let dir = temp_dir();
let path = dir.path().join("plain.db");
{
let conn = ManagedConnection::open(&path).unwrap();
conn.with(|c| Ok(c.execute_batch("CREATE TABLE t (id INTEGER)")?))
.unwrap();
}
assert_eq!(
detect_database_file_format(&path).unwrap(),
DatabaseFileFormat::PlainSQLite
);
}
#[test]
fn sqlcipher_database_detects_as_unrecognized() {
let dir = temp_dir();
let path = dir.path().join("cipher.db");
{
let conn = ManagedConnection::open_encrypted(&path, "secret").unwrap();
conn.with(|c| Ok(c.execute_batch("CREATE TABLE t (id INTEGER)")?))
.unwrap();
}
assert_eq!(
detect_database_file_format(&path).unwrap(),
DatabaseFileFormat::Unrecognized
);
}
#[test]
fn compressed_container_detects_with_encryption_flag() {
let dir = temp_dir();
let plain = dir.path().join("container.db");
{
let conn =
ManagedConnection::open_compressed(&plain, SQLiteCompressionOptions::default())
.unwrap();
conn.with(|c| Ok(c.execute_batch("CREATE TABLE t (id INTEGER)")?))
.unwrap();
}
assert_eq!(
detect_database_file_format(&plain).unwrap(),
DatabaseFileFormat::CompressedContainer { encrypted: false }
);
let encrypted = dir.path().join("container-enc.db");
{
let conn = ManagedConnection::open_compressed_encrypted(
&encrypted,
"secret",
SQLiteCompressionOptions::default(),
)
.unwrap();
conn.with(|c| Ok(c.execute_batch("CREATE TABLE t (id INTEGER)")?))
.unwrap();
}
assert_eq!(
detect_database_file_format(&encrypted).unwrap(),
DatabaseFileFormat::CompressedContainer { encrypted: true }
);
}
#[test]
fn legacy_compressed_header_is_classified_for_an_explicit_migration_error() {
let dir = temp_dir();
let path = dir.path().join("legacy-container.db");
let mut prefix = [0_u8; DETECT_PREFIX_LEN];
prefix[..LEGACY_MAGIC.len()].copy_from_slice(LEGACY_MAGIC);
prefix[8..12].copy_from_slice(&1_u32.to_le_bytes());
prefix[HEADER_FLAGS_OFFSET..HEADER_FLAGS_OFFSET + 4]
.copy_from_slice(&FLAG_ENCRYPTED.to_le_bytes());
std::fs::write(&path, prefix).unwrap();
assert_eq!(
detect_database_file_format(&path).unwrap(),
DatabaseFileFormat::CompressedContainer { encrypted: true }
);
}
#[test]
fn short_or_foreign_files_detect_as_unrecognized() {
let dir = temp_dir();
let short = dir.path().join("short.bin");
std::fs::write(&short, b"abc").unwrap();
assert_eq!(
detect_database_file_format(&short).unwrap(),
DatabaseFileFormat::Unrecognized
);
let foreign = dir.path().join("foreign.bin");
std::fs::write(&foreign, vec![0xAB_u8; 64]).unwrap();
assert_eq!(
detect_database_file_format(&foreign).unwrap(),
DatabaseFileFormat::Unrecognized
);
}
#[test]
fn requires_key_reflects_format() {
assert!(!DatabaseFileFormat::Missing.requires_key());
assert!(!DatabaseFileFormat::PlainSQLite.requires_key());
assert!(!DatabaseFileFormat::CompressedContainer { encrypted: false }.requires_key());
assert!(DatabaseFileFormat::CompressedContainer { encrypted: true }.requires_key());
assert!(DatabaseFileFormat::Unrecognized.requires_key());
}
}