mail4agent_server/
schema_guard.rs1use rusqlite::Connection;
18
19pub const SCHEMA_VERSION: i64 = 2;
20
21#[derive(Debug)]
22pub struct SchemaTooNew {
23 pub found: i64,
24 pub known: i64,
25}
26
27impl std::fmt::Display for SchemaTooNew {
28 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
29 write!(
30 f,
31 "the database has schema version {} but this build only knows up to {}: refusing to open it (use a newer build, or restore the snapshot taken before the upgrade)",
32 self.found, self.known
33 )
34 }
35}
36impl std::error::Error for SchemaTooNew {}
37
38fn has_data(conn: &Connection) -> bool {
39 conn.query_row("SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'rooms'", [], |_| Ok(())).is_ok()
40}
41
42pub fn enter(conn: &Connection) -> rusqlite::Result<()> {
44 let found: i64 = conn.query_row("PRAGMA user_version", [], |r| r.get(0))?;
45 if found > SCHEMA_VERSION {
46 return Err(rusqlite::Error::ToSqlConversionFailure(Box::new(SchemaTooNew { found, known: SCHEMA_VERSION })));
47 }
48 if found < SCHEMA_VERSION && has_data(conn) {
49 if let Err(e) = snapshot(conn) {
50 return Err(rusqlite::Error::ToSqlConversionFailure(format!("could not take the pre-migration snapshot ({e}); set M4A_DB_SNAPSHOT=off to migrate without one").into()));
51 }
52 }
53 Ok(())
54}
55
56pub fn leave(conn: &Connection) -> rusqlite::Result<()> {
58 let found: i64 = conn.query_row("PRAGMA user_version", [], |r| r.get(0))?;
59 if found < SCHEMA_VERSION {
60 conn.execute_batch(&format!("PRAGMA user_version = {SCHEMA_VERSION}"))?;
61 }
62 Ok(())
63}
64
65fn snapshot(conn: &Connection) -> std::io::Result<()> {
66 if std::env::var("M4A_DB_SNAPSHOT").map(|v| v == "off").unwrap_or(false) {
67 return Ok(());
68 }
69 let Some(path) = conn.path().filter(|p| !p.is_empty()).map(str::to_string) else { return Ok(()) };
70 let _ = conn.query_row("PRAGMA wal_checkpoint(TRUNCATE)", [], |_| Ok(()));
71 let src = std::path::Path::new(&path);
72 let name = src.file_name().map(|n| n.to_string_lossy().to_string()).unwrap_or_else(|| "db".into());
73 let dir = std::env::var("M4A_DB_SNAPSHOT_DIR").map(std::path::PathBuf::from).unwrap_or_else(|_| src.parent().map(|p| p.to_path_buf()).unwrap_or_default());
74 std::fs::create_dir_all(&dir)?;
75 let secs = std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).map(|d| d.as_secs()).unwrap_or(0);
76 std::fs::copy(src, dir.join(format!("{name}.pre-schema-v{SCHEMA_VERSION}-{secs}")))?;
77 Ok(())
78}
79
80#[cfg(test)]
81mod tests {
82 use super::*;
83
84 #[test]
85 fn a_newer_database_is_refused_and_an_older_one_is_stamped_after_a_snapshot() {
86 let dir = std::env::temp_dir().join(format!("m4a-schema-{}-{}", std::process::id(), std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos()));
87 std::fs::create_dir_all(&dir).unwrap();
88 let file = dir.join("t.db");
89 {
90 let c = Connection::open(&file).unwrap();
91 crate::store::create_matrix_schema(&c).unwrap();
92 assert_eq!(c.query_row("PRAGMA user_version", [], |r| r.get::<_, i64>(0)).unwrap(), SCHEMA_VERSION);
93 c.execute_batch("PRAGMA user_version = 0").unwrap();
95 }
96 {
97 let c = Connection::open(&file).unwrap();
98 crate::store::create_matrix_schema(&c).unwrap();
99 assert_eq!(c.query_row("PRAGMA user_version", [], |r| r.get::<_, i64>(0)).unwrap(), SCHEMA_VERSION);
100 }
101 let copies: Vec<_> = std::fs::read_dir(&dir).unwrap().flatten().filter(|e| e.file_name().to_string_lossy().contains(".pre-schema-v")).collect();
102 assert_eq!(copies.len(), 1, "one snapshot of the old file");
103 {
104 let c = Connection::open(&file).unwrap();
105 c.execute_batch(&format!("PRAGMA user_version = {}", SCHEMA_VERSION + 1)).unwrap();
106 }
107 let c = Connection::open(&file).unwrap();
108 let err = crate::store::create_matrix_schema(&c).unwrap_err().to_string();
109 assert!(err.contains("refusing to open"), "{err}");
110 let _ = std::fs::remove_dir_all(&dir);
111 }
112}