Skip to main content

aven_core/
db.rs

1use crate::ids::WorkspaceId;
2use std::fs;
3use std::path::{Path, PathBuf};
4use std::process::Command;
5use std::str::FromStr;
6use std::time::{Duration, SystemTime, UNIX_EPOCH};
7
8use anyhow::{Context, Result, bail};
9use serde_json::Value;
10use sqlx::sqlite::{SqliteConnectOptions, SqliteJournalMode, SqlitePoolOptions, SqliteRow};
11use sqlx::{Connection as _, Row, Sqlite, SqliteConnection, SqlitePool, Transaction};
12
13use crate::choices::{TaskPriority, TaskStatus};
14use crate::ids::{new_id, now};
15use crate::types::Task;
16use crate::workspaces::ensure_default_workspace;
17
18static MIGRATOR: sqlx::migrate::Migrator = sqlx::migrate!("./migrations");
19const MIGRATION_BACKUP_KEEP: usize = 20;
20
21#[derive(Clone)]
22pub struct Database {
23    pool: SqlitePool,
24    path: PathBuf,
25}
26
27impl Database {
28    pub async fn open(path: &Path) -> Result<Self> {
29        Ok(Self {
30            pool: open_db(path).await?,
31            path: path.to_path_buf(),
32        })
33    }
34
35    pub fn path(&self) -> &Path {
36        &self.path
37    }
38
39    pub fn latest_schema_version() -> Option<i64> {
40        MIGRATOR.iter().map(|migration| migration.version).max()
41    }
42
43    pub async fn meta(&self, key: &str) -> Result<Option<String>> {
44        let mut conn = self.acquire().await?;
45        get_meta(&mut conn, key).await
46    }
47
48    pub async fn conflict_exists(
49        &self,
50        workspace_id: &WorkspaceId,
51        task_id: &crate::ids::TaskId,
52        field: &str,
53    ) -> Result<bool> {
54        let mut conn = self.acquire().await?;
55        conflict_exists(&mut conn, workspace_id, task_id, field).await
56    }
57
58    pub(crate) async fn acquire(&self) -> Result<sqlx::pool::PoolConnection<Sqlite>> {
59        Ok(self.pool.acquire().await?)
60    }
61}
62
63pub(crate) async fn open_db(path: &Path) -> Result<SqlitePool> {
64    let existed_before_open = path.exists();
65    if let Some(parent) = path.parent() {
66        fs::create_dir_all(parent)
67            .with_context(|| format!("could not create {}", parent.display()))?;
68    }
69    let options = SqliteConnectOptions::from_str(&path.display().to_string())?
70        .create_if_missing(true)
71        .foreign_keys(true)
72        .journal_mode(SqliteJournalMode::Wal)
73        .busy_timeout(Duration::from_secs(5));
74    let pool = SqlitePoolOptions::new()
75        .max_connections(1)
76        .connect_with(options)
77        .await
78        .with_context(|| format!("could not open {}", path.display()))?;
79    backup_before_pending_migrations(path, existed_before_open, &pool).await?;
80    MIGRATOR.run(&pool).await?;
81    initialize_meta(&pool).await?;
82    let mut conn = pool.acquire().await?;
83    ensure_default_workspace(&mut conn).await?;
84    Ok(pool)
85}
86
87async fn backup_before_pending_migrations(
88    path: &Path,
89    existed_before_open: bool,
90    pool: &SqlitePool,
91) -> Result<()> {
92    if !migration_backups_enabled() || !existed_before_open || !has_pending_migrations(pool).await?
93    {
94        return Ok(());
95    }
96    let backup_path = migration_backup_path(path)?;
97    run_sqlite_backup(path, &backup_path)?;
98    prune_migration_backups(path)?;
99    Ok(())
100}
101
102fn migration_backups_enabled() -> bool {
103    std::env::var_os("AVEN_DEV_MIGRATION_BACKUPS").is_some()
104}
105
106async fn has_pending_migrations(pool: &SqlitePool) -> Result<bool> {
107    let applied_versions =
108        match sqlx::query_scalar::<_, i64>("SELECT version FROM _sqlx_migrations")
109            .fetch_all(pool)
110            .await
111        {
112            Ok(versions) => versions,
113            Err(error) => {
114                let Some(db_error) = error.as_database_error() else {
115                    return Err(error.into());
116                };
117                if db_error.code().as_deref() == Some("1") {
118                    return Ok(MIGRATOR.iter().next().is_some());
119                }
120                return Err(error.into());
121            }
122        };
123    Ok(MIGRATOR
124        .iter()
125        .any(|migration| !applied_versions.contains(&migration.version)))
126}
127
128fn migration_backup_path(path: &Path) -> Result<PathBuf> {
129    default_sqlite_backup_path(path, "before-migrate")
130}
131
132pub fn default_backup_path(path: &Path, reason: &str) -> Result<PathBuf> {
133    backup_path_with_extension(path, reason, "aven-backup.tar.zst")
134}
135
136pub fn default_sqlite_backup_path(path: &Path, reason: &str) -> Result<PathBuf> {
137    backup_path_with_extension(path, reason, "sqlite")
138}
139
140fn backup_path_with_extension(path: &Path, reason: &str, extension: &str) -> Result<PathBuf> {
141    let parent = path.parent().unwrap_or_else(|| Path::new("."));
142    let backup_dir = parent.join("backups");
143    fs::create_dir_all(&backup_dir)
144        .with_context(|| format!("could not create {}", backup_dir.display()))?;
145    let stem = path
146        .file_name()
147        .and_then(|name| name.to_str())
148        .unwrap_or("db.sqlite");
149    Ok(backup_dir.join(format!(
150        "{stem}.{reason}-{}.{}",
151        backup_timestamp()?,
152        extension
153    )))
154}
155
156pub fn backup_database(source: &Path, backup: &Path) -> Result<()> {
157    if let Some(parent) = backup.parent() {
158        fs::create_dir_all(parent)
159            .with_context(|| format!("could not create {}", parent.display()))?;
160    }
161    run_sqlite_backup(source, backup)
162}
163
164pub fn wal_path(path: &Path) -> PathBuf {
165    PathBuf::from(format!("{}-wal", path.display()))
166}
167
168pub fn shm_path(path: &Path) -> PathBuf {
169    PathBuf::from(format!("{}-shm", path.display()))
170}
171
172pub async fn restore_database_file(target: &Path, source: &Path) -> Result<PathBuf> {
173    validate_sqlite_source(source).await?;
174    let safety = default_sqlite_backup_path(target, "before-restore")?;
175    backup_database(target, &safety)?;
176    let staging = target.with_extension("restore-staging");
177    if staging.exists() {
178        fs::remove_file(&staging)
179            .with_context(|| format!("could not remove {}", staging.display()))?;
180    }
181    fs::copy(source, &staging).with_context(|| {
182        format!(
183            "could not copy {} -> {}",
184            source.display(),
185            staging.display()
186        )
187    })?;
188    for sidecar in [wal_path(target), shm_path(target)] {
189        if sidecar.exists() {
190            fs::remove_file(&sidecar)
191                .with_context(|| format!("could not remove {}", sidecar.display()))?;
192        }
193    }
194    fs::rename(&staging, target)
195        .with_context(|| format!("could not replace {}", target.display()))?;
196    Ok(safety)
197}
198
199async fn validate_sqlite_source(source: &Path) -> Result<()> {
200    let mut conn = sqlx::SqliteConnection::connect_with(
201        &SqliteConnectOptions::new()
202            .filename(source)
203            .read_only(true)
204            .foreign_keys(true),
205    )
206    .await
207    .with_context(|| format!("could not open source {}", source.display()))?;
208    let quick_check: String = sqlx::query_scalar("PRAGMA quick_check")
209        .fetch_one(&mut conn)
210        .await?;
211    if quick_check != "ok" {
212        bail!("error backup-source-corrupt quick_check={quick_check}");
213    }
214    Ok(())
215}
216
217fn backup_timestamp() -> Result<u64> {
218    Ok(SystemTime::now()
219        .duration_since(UNIX_EPOCH)
220        .context("system clock is before unix epoch")?
221        .as_secs())
222}
223
224fn run_sqlite_backup(source: &Path, backup: &Path) -> Result<()> {
225    let backup_sql = format!(".backup '{}'", sqlite_single_quoted(backup));
226    let output = Command::new("sqlite3")
227        .arg(source)
228        .arg(backup_sql)
229        .output()
230        .context("could not run sqlite3 for backup")?;
231    if !output.status.success() {
232        let stderr = String::from_utf8_lossy(&output.stderr);
233        bail!(
234            "sqlite3 backup failed status={} stderr={}",
235            output.status,
236            stderr.trim()
237        );
238    }
239    Ok(())
240}
241
242fn sqlite_single_quoted(path: &Path) -> String {
243    path.display().to_string().replace('\'', "''")
244}
245
246fn prune_migration_backups(path: &Path) -> Result<()> {
247    let Some(parent) = path.parent() else {
248        return Ok(());
249    };
250    let backup_dir = parent.join("backups");
251    let Some(file_name) = path.file_name().and_then(|name| name.to_str()) else {
252        return Ok(());
253    };
254    let prefix = format!("{file_name}.before-migrate-");
255    let mut backups = fs::read_dir(&backup_dir)
256        .with_context(|| format!("could not read {}", backup_dir.display()))?
257        .filter_map(|entry| entry.ok())
258        .filter(|entry| {
259            entry
260                .file_name()
261                .to_str()
262                .is_some_and(|name| name.starts_with(&prefix) && name.ends_with(".sqlite"))
263        })
264        .collect::<Vec<_>>();
265    backups.sort_by_key(|entry| entry.file_name());
266    let remove_count = backups.len().saturating_sub(MIGRATION_BACKUP_KEEP);
267    for entry in backups.into_iter().take(remove_count) {
268        let path = entry.path();
269        fs::remove_file(&path).with_context(|| format!("could not remove {}", path.display()))?;
270    }
271    Ok(())
272}
273
274async fn initialize_meta(pool: &SqlitePool) -> Result<()> {
275    let mut conn = pool.acquire().await?;
276    insert_meta_if_missing(&mut conn, "client_id", &new_id()).await?;
277    insert_meta_if_missing(&mut conn, "sync_cursor", "0").await?;
278    insert_meta_if_missing(&mut conn, "local_seq", "0").await?;
279    Ok(())
280}
281
282pub(crate) async fn current_schema_version(conn: &mut SqliteConnection) -> Result<i64> {
283    let version: Option<i64> = sqlx::query_scalar("SELECT MAX(version) FROM _sqlx_migrations")
284        .fetch_one(conn)
285        .await?;
286    Ok(version.unwrap_or(0))
287}
288
289pub(crate) async fn get_meta(conn: &mut SqliteConnection, key: &str) -> Result<Option<String>> {
290    Ok(
291        sqlx::query_scalar!("SELECT value FROM meta WHERE key = ?", key)
292            .fetch_optional(&mut *conn)
293            .await?,
294    )
295}
296
297pub(crate) async fn set_meta(conn: &mut SqliteConnection, key: &str, value: &str) -> Result<()> {
298    sqlx::query!(
299        "INSERT INTO meta(key, value) VALUES (?, ?)
300         ON CONFLICT(key) DO UPDATE SET value = excluded.value",
301        key,
302        value,
303    )
304    .execute(&mut *conn)
305    .await?;
306    Ok(())
307}
308
309pub(crate) async fn begin_immediate(
310    conn: &mut SqliteConnection,
311) -> sqlx::Result<Transaction<'_, Sqlite>> {
312    conn.begin_with("BEGIN IMMEDIATE").await
313}
314
315async fn insert_meta_if_missing(conn: &mut SqliteConnection, key: &str, value: &str) -> Result<()> {
316    sqlx::query!(
317        "INSERT OR IGNORE INTO meta(key, value) VALUES (?, ?)",
318        key,
319        value
320    )
321    .execute(&mut *conn)
322    .await?;
323    Ok(())
324}
325
326async fn next_local_seq(conn: &mut SqliteConnection) -> Result<i64> {
327    let seq = get_meta(conn, "local_seq")
328        .await?
329        .unwrap_or_else(|| "0".to_string())
330        .parse::<i64>()?
331        + 1;
332    set_meta(conn, "local_seq", &seq.to_string()).await?;
333    Ok(seq)
334}
335
336pub(crate) async fn insert_change(
337    conn: &mut SqliteConnection,
338    entity_type: &str,
339    entity_id: &str,
340    field: Option<&str>,
341    op_type: &str,
342    payload: Value,
343    base_version: Option<&str>,
344) -> Result<String> {
345    let change_id = new_id();
346    let client_id = get_meta(conn, "client_id")
347        .await?
348        .context("missing client id")?;
349    let local_seq = next_local_seq(conn).await?;
350    let created_at = now();
351    let payload = payload.to_string();
352    sqlx::query!(
353        "INSERT INTO changes(change_id, client_id, local_seq, entity_type, entity_id, field,
354         op_type, payload, base_version, created_at)
355         VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
356        change_id,
357        client_id,
358        local_seq,
359        entity_type,
360        entity_id,
361        field,
362        op_type,
363        payload,
364        base_version,
365        created_at,
366    )
367    .execute(&mut *conn)
368    .await?;
369    Ok(change_id)
370}
371
372fn optional_task_date(value: String) -> Option<String> {
373    (!value.is_empty()).then_some(value)
374}
375
376pub(crate) fn task_from_row(row: &SqliteRow) -> Result<Task> {
377    Ok(Task {
378        id: row.try_get("id")?,
379        workspace_id: row.try_get("workspace_id")?,
380        title: row.try_get("title")?,
381        description: row.try_get("description")?,
382        project_id: row.try_get("project_id")?,
383        project_key: row.try_get("project_key")?,
384        project_prefix: row.try_get("project_prefix")?,
385        status: TaskStatus::parse(row.try_get::<String, _>("status")?.as_str())?,
386        priority: TaskPriority::parse(row.try_get::<String, _>("priority")?.as_str())?,
387        created_at: row.try_get("created_at")?,
388        updated_at: row.try_get("updated_at")?,
389        queue_activity_at: row.try_get("queue_activity_at")?,
390        available_at: optional_task_date(row.try_get("available_at")?),
391        due_on: optional_task_date(row.try_get("due_on")?),
392        deleted: row.try_get::<i64, _>("deleted")? != 0,
393        is_epic: row.try_get::<i64, _>("is_epic")? != 0,
394    })
395}
396
397pub(crate) async fn field_version(
398    conn: &mut SqliteConnection,
399    entity_id: &str,
400    field: &str,
401) -> Result<Option<String>> {
402    Ok(
403        sqlx::query_scalar!(
404            r#"SELECT version AS "version!: String" FROM field_versions WHERE entity_id = ? AND field = ?"#,
405            entity_id,
406            field
407        )
408            .fetch_optional(&mut *conn)
409            .await?,
410    )
411}
412
413pub(crate) async fn set_field_version(
414    conn: &mut SqliteConnection,
415    entity_id: &str,
416    field: &str,
417    version: &str,
418) -> Result<()> {
419    sqlx::query!(
420        "INSERT INTO field_versions(entity_id, field, version) VALUES (?, ?, ?)
421         ON CONFLICT(entity_id, field) DO UPDATE SET version = excluded.version",
422        entity_id,
423        field,
424        version,
425    )
426    .execute(&mut *conn)
427    .await?;
428    Ok(())
429}
430
431pub(crate) async fn conflict_exists(
432    conn: &mut SqliteConnection,
433    workspace_id: &WorkspaceId,
434    task_id: &crate::ids::TaskId,
435    field: &str,
436) -> Result<bool> {
437    Ok(sqlx::query_scalar::<_, i64>(
438        "SELECT count(*) FROM conflicts WHERE workspace_id = ? AND task_id = ? AND field = ? AND resolved = 0 LIMIT 1",
439    )
440    .bind(workspace_id)
441    .bind(task_id)
442    .bind(field)
443    .fetch_one(&mut *conn)
444    .await?
445        > 0)
446}
447
448#[cfg(test)]
449mod tests {
450    use super::*;
451
452    #[tokio::test]
453    async fn task_from_row_maps_empty_dates_to_absence() {
454        let mut conn = SqliteConnection::connect(":memory:")
455            .await
456            .expect("open db");
457        let row = sqlx::query(
458            "SELECT 'TASK000000000001' AS id,
459                    '0000000000000000' AS workspace_id,
460                    'optional dates' AS title,
461                    '' AS description,
462                    '0000000000000001' AS project_id,
463                    'app' AS project_key,
464                    'APP' AS project_prefix,
465                    'todo' AS status,
466                    'none' AS priority,
467                    't' AS created_at,
468                    't' AS updated_at,
469                    't' AS queue_activity_at,
470                    '' AS available_at,
471                    '' AS due_on,
472                    0 AS deleted,
473                    0 AS is_epic",
474        )
475        .fetch_one(&mut conn)
476        .await
477        .expect("row");
478
479        let task = task_from_row(&row).unwrap();
480
481        assert_eq!(task.available_at, None);
482        assert_eq!(task.due_on, None);
483    }
484
485    #[test]
486    fn task_date_boundary_preserves_present_values() {
487        assert_eq!(
488            optional_task_date("2099-01-01T00:00:00Z".to_string()).as_deref(),
489            Some("2099-01-01T00:00:00Z")
490        );
491        assert_eq!(
492            optional_task_date("2099-01-01".to_string()).as_deref(),
493            Some("2099-01-01")
494        );
495    }
496
497    #[tokio::test]
498    async fn task_from_row_rejects_invalid_status_and_priority() {
499        let mut conn = SqliteConnection::connect(":memory:")
500            .await
501            .expect("open db");
502        let row = sqlx::query(
503            "SELECT 'TASK000000000001' AS id,
504                    '0000000000000000' AS workspace_id,
505                    'bad status' AS title,
506                    '' AS description,
507                    '0000000000000001' AS project_id,
508                    'app' AS project_key,
509                    'APP' AS project_prefix,
510                    'blocked' AS status,
511                    'none' AS priority,
512                    't' AS created_at,
513                    't' AS updated_at,
514                    't' AS queue_activity_at,
515                    0 AS deleted,
516                    0 AS is_epic",
517        )
518        .fetch_one(&mut conn)
519        .await
520        .expect("row");
521        assert_eq!(
522            task_from_row(&row).unwrap_err().to_string(),
523            "error invalid-status input=blocked choices=inbox,backlog,todo,active,done,canceled"
524        );
525
526        let row = sqlx::query(
527            "SELECT 'TASK000000000001' AS id,
528                    '0000000000000000' AS workspace_id,
529                    'bad priority' AS title,
530                    '' AS description,
531                    '0000000000000001' AS project_id,
532                    'app' AS project_key,
533                    'APP' AS project_prefix,
534                    'inbox' AS status,
535                    'soon' AS priority,
536                    't' AS created_at,
537                    't' AS updated_at,
538                    't' AS queue_activity_at,
539                    0 AS deleted,
540                    0 AS is_epic",
541        )
542        .fetch_one(&mut conn)
543        .await
544        .expect("row");
545        assert_eq!(
546            task_from_row(&row).unwrap_err().to_string(),
547            "error invalid-priority input=soon choices=none,low,medium,high,urgent"
548        );
549    }
550}