diff --git a/src/db/migrations.rs b/src/db/migrations.rs new file mode 100644 index 0000000..5059bb5 --- /dev/null +++ b/src/db/migrations.rs @@ -0,0 +1,218 @@ +use anyhow::Result; +use tokio_rusqlite::Connection; + +const CURRENT_VERSION: u32 = 1; + +/// Run migrations to set up or upgrade the database schema +pub async fn run_migrations(conn: &Connection) -> Result<()> { + conn.call(|c| { + check_and_migrate(c).map_err(|e| tokio_rusqlite::Error::Rusqlite(e)) + }) + .await + .map_err(|e| anyhow::anyhow!(e))?; + + Ok(()) +} + +/// Check schema version and apply migrations if needed +fn check_and_migrate(conn: &rusqlite::Connection) -> rusqlite::Result<()> { + // Check if schema_version table exists + let mut stmt = conn.prepare( + "SELECT name FROM sqlite_master WHERE type='table' AND name='schema_version'", + )?; + + let exists = stmt.exists([])?; + + if !exists { + // Fresh database - create full schema + create_schema_v1(conn)?; + return Ok(()); + } + + // Read current version + let mut stmt = conn.prepare("SELECT version FROM schema_version LIMIT 1")?; + let db_version: u32 = stmt.query_row([], |row| row.get(0))?; + + if db_version == CURRENT_VERSION { + // Already at current version + return Ok(()); + } + + if db_version > CURRENT_VERSION { + // Database is newer than this binary + // Return an error to signal this condition + return Err(rusqlite::Error::SqliteFailure( + rusqlite::ffi::Error::new(1), + Some("database newer than binary".to_string()), + )); + } + + // Handle future migrations if needed (e.g., db_version < CURRENT_VERSION) + // For now, this is a fresh v1 implementation + Ok(()) +} + +/// Create the initial database schema (version 1) +fn create_schema_v1(conn: &rusqlite::Connection) -> rusqlite::Result<()> { + // Schema version table + conn.execute( + "CREATE TABLE schema_version ( + version INTEGER NOT NULL, + migrated_at TEXT NOT NULL + )", + [], + )?; + + conn.execute( + "INSERT INTO schema_version (version, migrated_at) VALUES (?, ?)", + [ + "1", + &chrono::Utc::now().to_rfc3339(), + ], + )?; + + // Projects table + conn.execute( + "CREATE TABLE projects ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL UNIQUE, + path TEXT NOT NULL, + registered_at TEXT NOT NULL, + config_overrides TEXT, + metadata TEXT NOT NULL DEFAULT '{}' + )", + [], + )?; + + // Nodes table (unified work graph) + conn.execute( + "CREATE TABLE nodes ( + id TEXT PRIMARY KEY, + project_id TEXT NOT NULL REFERENCES projects(id), + node_type TEXT NOT NULL, + title TEXT NOT NULL, + description TEXT NOT NULL, + status TEXT NOT NULL DEFAULT 'pending', + priority TEXT, + assigned_to TEXT, + created_by TEXT, + labels TEXT NOT NULL DEFAULT '[]', + created_at TEXT NOT NULL, + started_at TEXT, + completed_at TEXT, + blocked_reason TEXT, + metadata TEXT NOT NULL DEFAULT '{}' + )", + [], + )?; + + conn.execute( + "CREATE INDEX idx_nodes_project ON nodes(project_id)", + [], + )?; + conn.execute( + "CREATE INDEX idx_nodes_type ON nodes(node_type)", + [], + )?; + conn.execute( + "CREATE INDEX idx_nodes_status ON nodes(status)", + [], + )?; + + // Edges table (unified relationships) + conn.execute( + "CREATE TABLE edges ( + id TEXT PRIMARY KEY, + edge_type TEXT NOT NULL, + from_node TEXT NOT NULL REFERENCES nodes(id), + to_node TEXT NOT NULL REFERENCES nodes(id), + label TEXT, + created_at TEXT NOT NULL + )", + [], + )?; + + conn.execute( + "CREATE INDEX idx_edges_from ON edges(from_node)", + [], + )?; + conn.execute( + "CREATE INDEX idx_edges_to ON edges(to_node)", + [], + )?; + conn.execute( + "CREATE INDEX idx_edges_type ON edges(edge_type)", + [], + )?; + + // Sessions table (temporal) + conn.execute( + "CREATE TABLE sessions ( + id TEXT PRIMARY KEY, + project_id TEXT NOT NULL REFERENCES projects(id), + goal_id TEXT NOT NULL REFERENCES nodes(id), + started_at TEXT NOT NULL, + ended_at TEXT, + handoff_notes TEXT, + agent_ids TEXT NOT NULL DEFAULT '[]', + summary TEXT + )", + [], + )?; + + // Full-text search virtual table + conn.execute( + "CREATE VIRTUAL TABLE nodes_fts USING fts5( + title, + description, + content='nodes', + content_rowid='rowid' + )", + [], + )?; + + // FTS sync triggers + conn.execute( + "CREATE TRIGGER nodes_ai AFTER INSERT ON nodes BEGIN + INSERT INTO nodes_fts(rowid, title, description) + VALUES (new.rowid, new.title, new.description); + END", + [], + )?; + + conn.execute( + "CREATE TRIGGER nodes_ad AFTER DELETE ON nodes BEGIN + INSERT INTO nodes_fts(nodes_fts, rowid, title, description) + VALUES ('delete', old.rowid, old.title, old.description); + END", + [], + )?; + + conn.execute( + "CREATE TRIGGER nodes_au AFTER UPDATE ON nodes BEGIN + INSERT INTO nodes_fts(nodes_fts, rowid, title, description) + VALUES ('delete', old.rowid, old.title, old.description); + INSERT INTO nodes_fts(rowid, title, description) + VALUES (new.rowid, new.title, new.description); + END", + [], + )?; + + // Worker conversations table + conn.execute( + "CREATE TABLE worker_conversations ( + id TEXT PRIMARY KEY, + session_id TEXT NOT NULL REFERENCES sessions(id), + agent_id TEXT NOT NULL, + task_ids TEXT NOT NULL DEFAULT '[]', + messages TEXT NOT NULL, + total_input_tokens INTEGER NOT NULL DEFAULT 0, + total_output_tokens INTEGER NOT NULL DEFAULT 0, + started_at TEXT NOT NULL, + completed_at TEXT + )", + [], + )?; + + Ok(()) +} diff --git a/src/db/mod.rs b/src/db/mod.rs new file mode 100644 index 0000000..495cab5 --- /dev/null +++ b/src/db/mod.rs @@ -0,0 +1,67 @@ +use anyhow::Result; +use std::path::Path; +use tokio_rusqlite::Connection; + +pub mod migrations; + +/// Database wrapper providing async access to SQLite +#[derive(Clone)] +pub struct Database { + conn: Connection, +} + +impl Database { + /// Open a database at the given path + pub async fn open(path: &Path) -> Result { + // Create parent directories if they don't exist + if let Some(parent) = path.parent() { + tokio::fs::create_dir_all(parent).await?; + } + + // Open the connection + let conn = Connection::open(path).await?; + + // Initialize pragmas + Self::init_pragmas(&conn).await?; + + // Run migrations + migrations::run_migrations(&conn).await?; + + Ok(Self { conn }) + } + + /// Open an in-memory database (useful for testing) + pub async fn open_in_memory() -> Result { + let conn = Connection::open_in_memory().await?; + + // Initialize pragmas + Self::init_pragmas(&conn).await?; + + // Run migrations + migrations::run_migrations(&conn).await?; + + Ok(Self { conn }) + } + + /// Get a reference to the connection + pub fn connection(&self) -> &Connection { + &self.conn + } + + /// Initialize pragma settings for WAL mode and safety + async fn init_pragmas(conn: &Connection) -> Result<()> { + conn.call(|c| { + c.execute_batch( + "PRAGMA journal_mode = WAL; + PRAGMA foreign_keys = ON; + PRAGMA busy_timeout = 5000; + PRAGMA wal_autocheckpoint = 1000;", + )?; + + Ok(()) + }) + .await?; + + Ok(()) + } +} diff --git a/src/lib.rs b/src/lib.rs index 33ec1cf..d4978a3 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,4 +1,5 @@ pub mod config; +pub mod db; pub mod llm; pub mod logging; pub mod planning; diff --git a/tests/db_test.rs b/tests/db_test.rs new file mode 100644 index 0000000..ada9db4 --- /dev/null +++ b/tests/db_test.rs @@ -0,0 +1,318 @@ +use rustagent::db::Database; +use tempfile::TempDir; + +#[tokio::test] +async fn test_database_open_creates_file() { + let temp_dir = TempDir::new().unwrap(); + let db_path = temp_dir.path().join("test.db"); + + let db = Database::open(&db_path).await.unwrap(); + + // Verify the file was created + assert!(db_path.exists()); + assert!(db_path.is_file()); + + // Connection should be accessible + let conn = db.connection(); + conn.call(|c| { + // Verify basic pragma + let mut stmt = c.prepare("PRAGMA database_list")?; + let _db_name: String = stmt.query_row([], |row| row.get(1))?; + Ok(()) + }) + .await + .unwrap(); +} + +#[tokio::test] +async fn test_database_pragma_wal_mode() { + let temp_dir = TempDir::new().unwrap(); + let db_path = temp_dir.path().join("test.db"); + let db = Database::open(&db_path).await.unwrap(); + + let conn = db.connection(); + let journal_mode = conn + .call(|c| { + let mut stmt = c.prepare("PRAGMA journal_mode")?; + let mode: String = stmt.query_row([], |row| row.get::<_, String>(0))?; + Ok(mode) + }) + .await + .unwrap(); + + assert_eq!(journal_mode.to_lowercase(), "wal"); +} + +#[tokio::test] +async fn test_database_pragma_foreign_keys() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let foreign_keys = conn + .call(|c| { + let mut stmt = c.prepare("PRAGMA foreign_keys")?; + let value: i32 = stmt.query_row([], |row| row.get::<_, i32>(0))?; + Ok(value) + }) + .await + .unwrap(); + + assert_eq!(foreign_keys, 1); +} + +#[tokio::test] +async fn test_database_schema_tables_exist() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let tables = conn + .call(|c| { + let mut stmt = c.prepare( + "SELECT name FROM sqlite_master WHERE type='table' ORDER BY name", + )?; + let tables = stmt + .query_map([], |row| row.get::<_, String>(0))? + .collect::, _>>()?; + Ok(tables) + }) + .await + .unwrap(); + + let expected_tables = vec![ + "edges", + "nodes", + "nodes_fts", + "projects", + "schema_version", + "sessions", + "worker_conversations", + ]; + + for expected in expected_tables { + assert!( + tables.contains(&expected.to_string()), + "Missing table: {}", + expected + ); + } +} + +#[tokio::test] +async fn test_schema_version_after_fresh_init() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let version = conn + .call(|c| { + let mut stmt = c.prepare("SELECT version FROM schema_version LIMIT 1")?; + let v: u32 = stmt.query_row([], |row| row.get::<_, u32>(0))?; + Ok(v) + }) + .await + .unwrap(); + + assert_eq!(version, 1); +} + +#[tokio::test] +async fn test_projects_table_has_correct_schema() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let columns = conn + .call(|c| { + let mut stmt = c.prepare("PRAGMA table_info(projects)")?; + let cols = stmt + .query_map([], |row| { + Ok(( + row.get::<_, String>(1)?, + row.get::<_, String>(2)?, + )) + })? + .collect::, _>>()?; + Ok(cols) + }) + .await + .unwrap(); + + let expected_cols = vec!["id", "name", "path", "registered_at", "config_overrides", "metadata"]; + + for (col_name, _) in columns { + assert!( + expected_cols.contains(&col_name.as_str()), + "Unexpected column: {}", + col_name + ); + } +} + +#[tokio::test] +async fn test_nodes_table_has_correct_schema() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let columns = conn + .call(|c| { + let mut stmt = c.prepare("PRAGMA table_info(nodes)")?; + let cols = stmt + .query_map([], |row| row.get::<_, String>(1))? + .collect::, _>>()?; + Ok(cols) + }) + .await + .unwrap(); + + let expected_cols = vec![ + "id", + "project_id", + "node_type", + "title", + "description", + "status", + "priority", + "assigned_to", + "created_by", + "labels", + "created_at", + "started_at", + "completed_at", + "blocked_reason", + "metadata", + ]; + + for expected in expected_cols { + assert!( + columns.contains(&expected.to_string()), + "Missing column in nodes table: {}", + expected + ); + } +} + +#[tokio::test] +async fn test_fts_table_exists() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let exists = conn + .call(|c| { + let mut stmt = c.prepare( + "SELECT name FROM sqlite_master WHERE type='table' AND name='nodes_fts'", + )?; + let result = stmt.exists([])?; + Ok(result) + }) + .await + .unwrap(); + + assert!(exists, "nodes_fts virtual table does not exist"); +} + +#[tokio::test] +async fn test_triggers_exist() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let triggers = conn + .call(|c| { + let mut stmt = c.prepare( + "SELECT name FROM sqlite_master WHERE type='trigger' ORDER BY name", + )?; + let triggers = stmt + .query_map([], |row| row.get::<_, String>(0))? + .collect::, _>>()?; + Ok(triggers) + }) + .await + .unwrap(); + + let expected_triggers = vec!["nodes_ai", "nodes_ad", "nodes_au"]; + + for expected in expected_triggers { + assert!( + triggers.contains(&expected.to_string()), + "Missing trigger: {}", + expected + ); + } +} + +#[tokio::test] +async fn test_opening_existing_database_does_not_error() { + let temp_dir = TempDir::new().unwrap(); + let db_path = temp_dir.path().join("test.db"); + + // Create database first time + let _db1 = Database::open(&db_path).await.unwrap(); + drop(_db1); + + // Open again - should work + let _db2 = Database::open(&db_path).await.unwrap(); +} + +#[tokio::test] +async fn test_database_clone() { + let db = Database::open_in_memory().await.unwrap(); + + let db_clone = db.clone(); + + // Both should be able to access the database + let conn1 = db.connection(); + let conn2 = db_clone.connection(); + + let result1 = conn1 + .call(|c| { + let mut stmt = c.prepare("SELECT version FROM schema_version")?; + let version: u32 = stmt.query_row([], |row| row.get::<_, u32>(0))?; + Ok(version) + }) + .await + .unwrap(); + + let result2 = conn2 + .call(|c| { + let mut stmt = c.prepare("SELECT version FROM schema_version")?; + let version: u32 = stmt.query_row([], |row| row.get::<_, u32>(0))?; + Ok(version) + }) + .await + .unwrap(); + + assert_eq!(result1, result2); +} + +#[tokio::test] +async fn test_indexes_exist() { + let db = Database::open_in_memory().await.unwrap(); + + let conn = db.connection(); + let indexes = conn + .call(|c| { + let mut stmt = c.prepare( + "SELECT name FROM sqlite_master WHERE type='index' AND name LIKE 'idx_%' ORDER BY name", + )?; + let indexes = stmt + .query_map([], |row| row.get::<_, String>(0))? + .collect::, _>>()?; + Ok(indexes) + }) + .await + .unwrap(); + + let expected_indexes = vec![ + "idx_edges_from", + "idx_edges_to", + "idx_edges_type", + "idx_nodes_project", + "idx_nodes_status", + "idx_nodes_type", + ]; + + for expected in expected_indexes { + assert!( + indexes.contains(&expected.to_string()), + "Missing index: {}", + expected + ); + } +}