Skip to content
use crate::db::Database;

struct Migration {
    name: &'static str,
    sql: &'static str,
}

const MIGRATIONS: &[Migration] = &[
    Migration {
        name: "create_users_table",
        sql: r"CREATE TABLE IF NOT EXISTS users (
            id TEXT PRIMARY KEY,
            username TEXT UNIQUE NOT NULL,
            email TEXT,
            password_hash TEXT NOT NULL,
            created_at TEXT NOT NULL
        );",
    },
    Migration {
        name: "create_tokens_table",
        sql: r"CREATE TABLE IF NOT EXISTS tokens (
            token TEXT PRIMARY KEY,
            user_id TEXT NOT NULL,
            created_at TEXT NOT NULL,
            FOREIGN KEY (user_id) REFERENCES users(id)
        );",
    },
    Migration {
        name: "create_tokens_created_at_index",
        sql: r"CREATE INDEX IF NOT EXISTS idx_tokens_created_at ON tokens(created_at);",
    },
    Migration {
        name: "create_sessions_table",
        sql: r"CREATE TABLE IF NOT EXISTS sessions (
            session_key TEXT PRIMARY KEY,
            state TEXT NOT NULL,
            expires_at INTEGER NOT NULL
        );",
    },
    Migration {
        name: "create_sessions_expires_at_index",
        sql: r"CREATE INDEX IF NOT EXISTS idx_sessions_expires_at ON sessions(expires_at);",
    },
    Migration {
        name: "drop_legacy_tokens_table",
        sql: r"DROP TABLE IF EXISTS tokens;",
    },
    Migration {
        name: "create_invites_table",
        sql: r"CREATE TABLE IF NOT EXISTS invites (
            token TEXT PRIMARY KEY,
            created_at TEXT NOT NULL
        );",
    },
    Migration {
        name: "create_tickets_table",
        sql: r"CREATE TABLE IF NOT EXISTS tickets (
            id TEXT PRIMARY KEY,
            title TEXT NOT NULL,
            description TEXT NOT NULL DEFAULT '',
            creator_id TEXT NOT NULL,
            created_at TEXT NOT NULL,
            FOREIGN KEY (creator_id) REFERENCES users(id)
        );",
    },
    Migration {
        name: "create_tags_table",
        sql: r"CREATE TABLE IF NOT EXISTS tags (id TEXT PRIMARY KEY, name TEXT UNIQUE NOT NULL);",
    },
    Migration {
        name: "create_ticket_tags_table",
        sql: r"CREATE TABLE IF NOT EXISTS ticket_tags (
            ticket_id TEXT NOT NULL, tag_id TEXT NOT NULL, PRIMARY KEY (ticket_id, tag_id),
            FOREIGN KEY (ticket_id) REFERENCES tickets(id) ON DELETE CASCADE,
            FOREIGN KEY (tag_id) REFERENCES tags(id) ON DELETE CASCADE
        );",
    },
    Migration {
        name: "create_ticket_tags_index",
        sql: r"CREATE INDEX IF NOT EXISTS idx_ticket_tags_tag ON ticket_tags(tag_id);",
    },
    Migration {
        name: "add_quest_completion_fields",
        sql: r"ALTER TABLE tickets ADD COLUMN completed_at TEXT;",
    },
    Migration {
        name: "add_ticket_outcome_field",
        sql: r"ALTER TABLE tickets ADD COLUMN closed_as TEXT;",
    },
    Migration {
        name: "backfill_ticket_done_outcomes",
        sql: r"UPDATE tickets SET closed_as = 'done' WHERE completed_at IS NOT NULL;",
    },
];

impl Database {
    pub async fn create_tables(&self) -> Result<(), String> {
        let conn = self.conn().await?;

        conn.query("PRAGMA journal_mode = WAL;", ())
            .await
            .map_err(|e| format!("Failed to set WAL mode: {e}"))?;

        conn.execute(
            r"CREATE TABLE IF NOT EXISTS _migrations (
                name TEXT PRIMARY KEY,
                applied_at TEXT NOT NULL DEFAULT (datetime('now'))
            );",
            (),
        )
        .await
        .map_err(|e| format!("Failed to create migrations table: {e}"))?;

        let mut rows = conn
            .query("SELECT name FROM _migrations", ())
            .await
            .map_err(|e| format!("Failed to list applied migrations: {e}"))?;

        let mut applied = Vec::new();
        while let Some(row) = rows.next().await.map_err(|e| e.to_string())? {
            applied.push(row.get::<String>(0).map_err(|e| e.to_string())?);
        }

        for migration in MIGRATIONS {
            if applied.contains(&migration.name.to_string()) {
                continue;
            }

            conn.execute("BEGIN IMMEDIATE", ())
                .await
                .map_err(|e| e.to_string())?;

            let result: Result<bool, String> = async {
                let mut rows = conn
                    .query(
                        "SELECT 1 FROM _migrations WHERE name = ?1",
                        turso::params![migration.name],
                    )
                    .await
                    .map_err(|e| format!("Failed to check migration {}: {e}", migration.name))?;

                if rows.next().await.map_err(|e| e.to_string())?.is_some() {
                    return Ok(false);
                }

                conn.execute(migration.sql, ())
                    .await
                    .map_err(|e| format!("Migration {} failed: {e}", migration.name))?;

                conn.execute(
                    "INSERT INTO _migrations (name) VALUES (?1)",
                    turso::params![migration.name],
                )
                .await
                .map_err(|e| format!("Failed to record migration {}: {e}", migration.name))?;

                Ok(true)
            }
            .await;

            match result {
                Ok(true) => {
                    conn.execute("COMMIT", ()).await.map_err(|e| {
                        format!("Failed to commit migration {}: {e}", migration.name)
                    })?;
                }
                Ok(false) => {
                    conn.execute("ROLLBACK", ()).await.map_err(|e| {
                        format!("Failed to rollback migration {}: {e}", migration.name)
                    })?;
                }
                Err(e) => {
                    conn.execute("ROLLBACK", ())
                        .await
                        .map_err(|re| format!("{e} (rollback also failed: {re})"))?;
                    return Err(e);
                }
            }
        }

        Ok(())
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    #[tokio::test]
    async fn test_create_tables_runs_all_migrations() {
        let db_path = format!("/tmp/test_pear_migrations_{}.db", uuid::Uuid::new_v4());
        let db = Database::new(&db_path);

        assert!(db.create_tables().await.is_ok());

        let mut rows = db
            .conn()
            .await
            .unwrap()
            .query(
                "SELECT name FROM sqlite_master WHERE type='table' ORDER BY name",
                (),
            )
            .await
            .unwrap();

        let mut tables = Vec::new();
        while let Some(row) = rows.next().await.unwrap() {
            tables.push(row.get::<String>(0).unwrap());
        }

        for table in [
            "users",
            "sessions",
            "invites",
            "tickets",
            "tags",
            "ticket_tags",
            "_migrations",
        ] {
            assert!(tables.contains(&table.to_string()), "missing {table}");
        }
        assert!(!tables.contains(&"projects".to_string()));
        assert!(!tables.contains(&"ticket_projects".to_string()));

        let mut columns = db
            .conn()
            .await
            .unwrap()
            .query("PRAGMA table_info(tickets)", ())
            .await
            .unwrap();
        let mut ticket_columns = Vec::new();
        while let Some(row) = columns.next().await.unwrap() {
            ticket_columns.push(row.get::<String>(1).unwrap());
        }
        assert!(ticket_columns.contains(&"closed_as".to_string()));

        let _ = std::fs::remove_file(&db_path);
    }

    #[tokio::test]
    async fn test_concurrent_create_tables_runs_migrations_once() {
        let db_path = format!(
            "/tmp/test_pear_concurrent_migrations_{}.db",
            uuid::Uuid::new_v4()
        );
        let db1 = Database::new(&db_path);
        let db2 = Database::new(&db_path);

        let (result1, result2) = tokio::join!(db1.create_tables(), db2.create_tables());
        assert!(result1.is_ok(), "first create_tables failed: {result1:?}");
        assert!(result2.is_ok(), "second create_tables failed: {result2:?}");

        let mut rows = db1
            .conn()
            .await
            .unwrap()
            .query(
                "SELECT name FROM _migrations WHERE name = 'add_ticket_outcome_field'",
                (),
            )
            .await
            .unwrap();
        assert!(rows.next().await.unwrap().is_some());

        let _ = std::fs::remove_file(&db_path);
    }
}