Skip to main content

mail4agent_server/
schema_guard.rs

1//! Database schema versioning: forward-only, with a snapshot before a migration.
2//!
3//! The version lives in SQLite's `user_version` pragma. A build knows one version
4//! ([`SCHEMA_VERSION`]). On open:
5//!
6//! - a file whose version is NEWER than the build's is refused with a clear error (a downgrade
7//!   would run old code against tables it does not understand; restore the snapshot instead);
8//! - a file with data and an OLDER version (0 = written before versioning, i.e. 0.4.x) is first
9//!   copied, then migrated, then marked with the new version. The copy is the downgrade path.
10//!
11//! Environment: `M4A_DB_SNAPSHOT=off` disables the copy; `M4A_DB_SNAPSHOT_DIR` puts it in a
12//! directory (default: next to the database file). The copy is named
13//! `<file>.pre-schema-v<new>-<unix seconds>`.
14//!
15//! History: 1 = up to 0.4.x (no marker); 2 = 0.5.0 (presence, remote media cache, login tokens).
16
17use 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
42/// Called before the schema is created. Refuses a newer database; snapshots an older one.
43pub 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
56/// Called after the schema is in place: stamps the version (forward only).
57pub 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            // Pretend the file was written by 0.4.x: no marker.
94            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}