#![cfg(test)]
use crate::commands::backup::handle_backup;
use crate::memory::crud::test_fake_embedder;
use crate::sqlite::{Database, vec_to_blob};
use rusqlite::Connection;
use std::path::Path;
fn create_test_db() -> (tempfile::TempDir, std::path::PathBuf) {
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("memories.db");
Database::open(&path).unwrap();
(dir, path)
}
fn insert_row(db: &Database, project_id: &str, content: &str, embedding: &[f32]) -> String {
db.insert(project_id, content, embedding, None, "fact", "active")
.unwrap()
}
fn default_destination(src: &Path) -> std::path::PathBuf {
let stem = src.file_stem().and_then(|s| s.to_str()).unwrap();
let ext = src.extension().and_then(|s| s.to_str()).unwrap_or("db");
let parent = src.parent().unwrap();
parent.join(format!("{}-backup.{}", stem, ext))
}
#[test]
fn test_backup_default_destination_is_queryable_copy() {
let (_dir, src) = create_test_db();
let db = Database::open(&src).unwrap();
let emb = test_fake_embedder("alpha").unwrap();
let id = insert_row(&db, "proj", "alpha", &emb);
let exit = handle_backup(&src, None, true).expect("handle_backup should succeed");
assert_eq!(exit, std::process::ExitCode::SUCCESS);
let expected = default_destination(&src);
assert!(expected.exists(), "default destination must exist");
let dest_conn = Connection::open(&expected).unwrap();
let count: i64 = dest_conn
.query_row("SELECT COUNT(*) FROM memories", [], |r| r.get(0))
.unwrap();
assert_eq!(count, 1);
let (id2, blob2): (String, Vec<u8>) = dest_conn
.query_row("SELECT id, embedding FROM memories LIMIT 1", [], |r| {
Ok((r.get(0)?, r.get(1)?))
})
.unwrap();
assert_eq!(id2, id);
let expected_blob: Vec<u8> = vec_to_blob(&emb).unwrap();
assert_eq!(
blob2, expected_blob,
"backup must be byte-identical for the embedding BLOB"
);
}
#[test]
fn test_backup_explicit_output_path() {
let (_dir, src) = create_test_db();
let db = Database::open(&src).unwrap();
let emb = test_fake_embedder("beta").unwrap();
let _id = insert_row(&db, "proj", "beta", &emb);
let out = _dir.keep().join("explicit-backup.db");
let exit = handle_backup(&src, Some(Path::new(out.to_str().unwrap())), true)
.expect("handle_backup should succeed");
assert_eq!(exit, std::process::ExitCode::SUCCESS);
let dest_conn = Connection::open(&out).unwrap();
let count: i64 = dest_conn
.query_row("SELECT COUNT(*) FROM memories", [], |r| r.get(0))
.unwrap();
assert_eq!(count, 1);
}
#[test]
fn test_backup_carries_corrupt_and_null_embeddings_verbatim() {
let (_dir, src) = create_test_db();
let db = Database::open(&src).unwrap();
let normal_emb = test_fake_embedder("normal").unwrap();
let normal_id = insert_row(&db, "proj", "normal", &normal_emb);
let corrupt_blob: Vec<u8> = vec![0xABu8; 1535];
let corrupt_id = "corrupt-1".to_string();
db.conn().execute(
"INSERT INTO memories (id, project_id, content, embedding, metadata, created_at, updated_at, type, status)
VALUES (?1, 'proj', 'corrupt', ?2, NULL, '2024-01-01T00:00:00Z', '2024-01-01T00:00:00Z', 'fact', 'active')",
rusqlite::params![&corrupt_id, &corrupt_blob],
).unwrap();
let empty_id = "empty-2".to_string();
db.conn().execute(
"INSERT INTO memories (id, project_id, content, embedding, metadata, created_at, updated_at, type, status)
VALUES (?1, 'proj', 'empty', ?2, NULL, '2024-01-01T00:00:00Z', '2024-01-01T00:00:00Z', 'fact', 'active')",
rusqlite::params![&empty_id, Vec::<u8>::new()],
)
.unwrap();
let exit = handle_backup(&src, None, true).expect("handle_backup should succeed");
assert_eq!(exit, std::process::ExitCode::SUCCESS);
let expected = default_destination(&src);
let dest_conn = Connection::open(&expected).unwrap();
let count: i64 = dest_conn
.query_row("SELECT COUNT(*) FROM memories", [], |r| r.get(0))
.unwrap();
assert_eq!(count, 3, "all three rows must be in the backup");
let cases: Vec<(&str, Vec<u8>)> = vec![
(&normal_id, vec_to_blob(&normal_emb).unwrap()),
(&corrupt_id, vec![0xABu8; 1535]),
(&empty_id, Vec::new()),
];
for (id, expected_blob) in cases {
let blob: Vec<u8> = dest_conn
.query_row("SELECT embedding FROM memories WHERE id = ?", [id], |r| {
r.get(0)
})
.unwrap();
assert_eq!(
blob, expected_blob,
"row {id} embedding must carry through byte-for-byte"
);
}
}
#[test]
fn test_backup_fails_fast_when_source_locked() {
let (_dir, src) = create_test_db();
let db = Database::open(&src).unwrap();
let emb = test_fake_embedder("locked").unwrap();
let _id = insert_row(&db, "proj", "locked", &emb);
let lock_conn = Connection::open(&src).unwrap();
lock_conn.execute("BEGIN EXCLUSIVE", []).unwrap();
let result = handle_backup(&src, None, true);
match result {
Err(e) => {
let msg = e.to_string();
assert!(
msg.contains("locked"),
"Expected 'locked' in error message, got: {}",
msg
);
assert!(
msg.contains("MCP server"),
"Expected 'MCP server' in error message, got: {}",
msg
);
}
Ok(_) => panic!("Expected error when database is locked, got Ok"),
}
lock_conn.execute("ROLLBACK", []).unwrap();
}
#[test]
fn test_backup_replaces_existing_destination() {
let (_dir, src) = create_test_db();
let db = Database::open(&src).unwrap();
let emb1 = test_fake_embedder("first").unwrap();
let _id = insert_row(&db, "proj", "first", &emb1);
let expected = default_destination(&src);
let exit1 = handle_backup(&src, None, true).expect("first backup should succeed");
assert_eq!(exit1, std::process::ExitCode::SUCCESS);
assert!(expected.exists());
let emb2 = test_fake_embedder("second").unwrap();
let _id2 = insert_row(&db, "proj", "second", &emb2);
let exit2 = handle_backup(&src, None, true).expect("second backup should succeed");
assert_eq!(exit2, std::process::ExitCode::SUCCESS);
let dest_conn = Connection::open(&expected).unwrap();
let count: i64 = dest_conn
.query_row("SELECT COUNT(*) FROM memories", [], |r| r.get(0))
.unwrap();
assert_eq!(count, 2, "second backup must replace the first");
}
#[test]
fn test_backup_preserves_fts_index() {
let (_dir, src) = create_test_db();
let db = Database::open(&src).unwrap();
let emb = test_fake_embedder("fts probe").unwrap();
let _id = insert_row(&db, "proj", "fts probe", &emb);
let exit = handle_backup(&src, None, true).expect("handle_backup should succeed");
assert_eq!(exit, std::process::ExitCode::SUCCESS);
let expected = default_destination(&src);
let dest_conn = Connection::open(&expected).unwrap();
let fts_rows: i64 = dest_conn
.query_row("SELECT count(*) FROM memories_fts", [], |r| r.get(0))
.expect("FTS5 content must be queryable (a corrupted index would error)");
assert_eq!(fts_rows, 1, "FTS5 content must reflect the seeded row");
}