use super::*;
use crate::NodeId;
use khive_bm25::Bm25Index;
use khive_hnsw::HnswIndex;
use rusqlite::Connection;
use std::sync::Arc;
use tokio::sync::Mutex;
async fn setup_test_persistence() -> RetrievalPersistence {
let conn = Connection::open_in_memory().expect("open in-memory db");
conn.execute_batch("PRAGMA journal_mode=WAL; PRAGMA synchronous=NORMAL;")
.expect("set pragmas");
let persist = RetrievalPersistence::new(Arc::new(Mutex::new(conn)), "test");
persist.init_schema().await.expect("init schema");
persist
}
#[tokio::test]
async fn test_persist_and_load_bm25() {
let persist = setup_test_persistence().await;
let mut index = Bm25Index::default();
index
.index_document("doc1", "hello world")
.expect("index doc");
index
.index_document("doc2", "goodbye world")
.expect("index doc");
persist.persist_bm25_index(&index).await.expect("persist");
let loaded = persist.load_bm25_index().await.expect("load");
assert!(loaded.is_some());
let loaded = loaded.unwrap();
assert_eq!(loaded.doc_count(), 2);
}
#[tokio::test]
async fn test_persist_and_load_hnsw() {
let persist = setup_test_persistence().await;
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
index.insert(id2, vec![0.0, 1.0, 0.0, 0.0]).expect("insert");
index.insert(id3, vec![0.0, 0.0, 1.0, 0.0]).expect("insert");
assert_eq!(index.len(), 3);
persist
.persist_hnsw_snapshot(&index)
.await
.expect("persist");
let loaded = persist.load_hnsw_snapshot().await.expect("load");
assert!(loaded.is_some());
let snapshot = loaded.unwrap();
assert_eq!(snapshot.total_nodes, 3);
assert_eq!(snapshot.live_nodes, 3);
assert_eq!(snapshot.tombstone_count, 0);
assert_eq!(snapshot.indexed_ids.len(), 3);
assert!(snapshot.indexed_ids.contains(&id1));
assert!(snapshot.indexed_ids.contains(&id2));
assert!(snapshot.indexed_ids.contains(&id3));
}
#[tokio::test]
async fn test_stats() {
let persist = setup_test_persistence().await;
let stats = persist.stats().await.expect("stats");
assert_eq!(stats.hnsw_snapshot_size, 0);
assert_eq!(stats.bm25_snapshot_size, 0);
let index = Bm25Index::default();
persist.persist_bm25_index(&index).await.expect("persist");
let stats = persist.stats().await.expect("stats");
assert!(stats.bm25_snapshot_size > 0);
assert!(stats.bm25_snapshot_at.is_some());
}
#[tokio::test]
async fn test_clear() {
let persist = setup_test_persistence().await;
let index = Bm25Index::default();
persist.persist_bm25_index(&index).await.expect("persist");
persist.clear().await.expect("clear");
let loaded = persist.load_bm25_index().await.expect("load");
assert!(loaded.is_none());
}
#[tokio::test]
async fn test_shadow_validation_config_default() {
let config = ShadowValidationConfig::default();
assert!(!config.enabled);
assert!((config.sample_rate - 0.1).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_shadow_validation_config_enabled() {
let config = ShadowValidationConfig::enabled();
assert!(config.enabled);
assert!((config.sample_rate - 1.0).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_shadow_validation_config_sample_rate() {
let config = ShadowValidationConfig::with_sample_rate(0.5);
assert!(config.enabled);
assert!((config.sample_rate - 0.5).abs() < f64::EPSILON);
let config = ShadowValidationConfig::with_sample_rate(1.5);
assert!((config.sample_rate - 1.0).abs() < f64::EPSILON);
let config = ShadowValidationConfig::with_sample_rate(-0.5);
assert!((config.sample_rate - 0.0).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_bm25_shadow_validation_passes() {
let persist = setup_test_persistence().await;
let config = ShadowValidationConfig::enabled();
let mut index = Bm25Index::default();
index
.index_document("doc1", "hello world")
.expect("index doc");
index
.index_document("doc2", "goodbye world")
.expect("index doc");
let result = persist
.persist_bm25_with_validation(&index, &config)
.await
.expect("persist with validation");
assert!(result.is_some());
let validation = result.unwrap();
assert!(
validation.passed,
"validation should pass: {:?}",
validation.discrepancies
);
assert_eq!(validation.index_type, "bm25");
assert_eq!(validation.expected.item_count, 2);
assert!(validation.discrepancies.is_empty());
}
#[tokio::test]
async fn test_shadow_validation_skipped_when_disabled() {
let persist = setup_test_persistence().await;
let config = ShadowValidationConfig::default();
let index = Bm25Index::default();
let result = persist
.persist_bm25_with_validation(&index, &config)
.await
.expect("persist");
assert!(result.is_none());
let loaded = persist.load_bm25_index().await.expect("load");
assert!(loaded.is_some());
}
#[tokio::test]
async fn test_should_sample() {
use super::shadow::should_sample;
assert!(should_sample(1.0));
assert!(should_sample(1.5));
assert!(!should_sample(0.0));
assert!(!should_sample(-0.5)); }
#[tokio::test]
async fn test_hnsw_shadow_validation_passes() {
let persist = setup_test_persistence().await;
let config = ShadowValidationConfig::enabled();
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
index.insert(id2, vec![0.0, 1.0, 0.0, 0.0]).expect("insert");
let result = persist
.persist_hnsw_with_validation(&index, &config)
.await
.expect("persist with validation");
assert!(result.is_some());
let validation = result.unwrap();
assert!(
validation.passed,
"validation should pass: {:?}",
validation.discrepancies
);
assert_eq!(validation.index_type, "hnsw");
assert_eq!(validation.expected.item_count, 2);
assert!(validation.discrepancies.is_empty());
}
#[tokio::test]
async fn test_hnsw_shadow_validation_with_tombstones() {
let persist = setup_test_persistence().await;
let config = ShadowValidationConfig::enabled();
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
index.insert(id2, vec![0.0, 1.0, 0.0, 0.0]).expect("insert");
index.insert(id3, vec![0.0, 0.0, 1.0, 0.0]).expect("insert");
index.delete(id2);
let result = persist
.persist_hnsw_with_validation(&index, &config)
.await
.expect("persist with validation");
assert!(result.is_some());
let validation = result.unwrap();
assert!(
validation.passed,
"validation should pass with tombstones: {:?}",
validation.discrepancies
);
assert_eq!(validation.expected.item_count, 3); assert_eq!(validation.expected.tombstone_count, 1);
}
#[tokio::test]
async fn test_namespace_isolation() {
let conn = Connection::open_in_memory().expect("open in-memory db");
conn.execute_batch("PRAGMA journal_mode=WAL; PRAGMA synchronous=NORMAL;")
.expect("set pragmas");
let conn = Arc::new(Mutex::new(conn));
let persist_ns1 = RetrievalPersistence::new(conn.clone(), "namespace1");
let persist_ns2 = RetrievalPersistence::new(conn.clone(), "namespace2");
persist_ns1.init_schema().await.expect("init schema");
let mut index1 = Bm25Index::default();
index1
.index_document("doc1", "namespace one content")
.expect("index");
let mut index2 = Bm25Index::default();
index2
.index_document("doc2", "namespace two content")
.expect("index");
index2
.index_document("doc3", "more namespace two")
.expect("index");
persist_ns1
.persist_bm25_index(&index1)
.await
.expect("persist ns1");
persist_ns2
.persist_bm25_index(&index2)
.await
.expect("persist ns2");
let loaded1 = persist_ns1.load_bm25_index().await.expect("load ns1");
let loaded2 = persist_ns2.load_bm25_index().await.expect("load ns2");
assert!(loaded1.is_some());
assert!(loaded2.is_some());
assert_eq!(loaded1.unwrap().doc_count(), 1);
assert_eq!(loaded2.unwrap().doc_count(), 2);
persist_ns1.clear().await.expect("clear ns1");
let loaded1_after = persist_ns1
.load_bm25_index()
.await
.expect("load ns1 after clear");
let loaded2_after = persist_ns2
.load_bm25_index()
.await
.expect("load ns2 after clear");
assert!(loaded1_after.is_none(), "ns1 should be cleared");
assert!(loaded2_after.is_some(), "ns2 should still exist");
assert_eq!(loaded2_after.unwrap().doc_count(), 2);
}
#[tokio::test]
async fn test_corrupted_bm25_data_returns_error() {
let persist = setup_test_persistence().await;
{
let conn = persist.conn.clone();
let namespace = "test".to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, 'bm25', ?2, strftime('%s', 'now'))
"#,
rusqlite::params![namespace, b"not valid json {{{{"],
)
.expect("insert corrupted");
})
.await
.expect("spawn");
}
let result = persist.load_bm25_index().await;
assert!(result.is_err(), "loading corrupted data should fail");
let err = result.unwrap_err();
assert!(matches!(err, PersistError::Deserialize(_)));
}
#[tokio::test]
async fn test_corrupted_hnsw_data_returns_error() {
let persist = setup_test_persistence().await;
{
let conn = persist.conn.clone();
let namespace = "test".to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, 'hnsw', ?2, strftime('%s', 'now'))
"#,
rusqlite::params![namespace, b"truncated json {\"total_nodes\":"],
)
.expect("insert corrupted");
})
.await
.expect("spawn");
}
let result = persist.load_hnsw_snapshot().await;
assert!(result.is_err(), "loading corrupted HNSW data should fail");
let err = result.unwrap_err();
assert!(matches!(err, PersistError::Deserialize(_)));
}
#[tokio::test]
async fn test_valid_json_wrong_schema_bm25() {
let persist = setup_test_persistence().await;
{
let conn = persist.conn.clone();
let namespace = "test".to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
let wrong_schema = br#"{"some_field": "value", "other": 123}"#;
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, 'bm25', ?2, strftime('%s', 'now'))
"#,
rusqlite::params![namespace, wrong_schema.as_slice()],
)
.expect("insert wrong schema");
})
.await
.expect("spawn");
}
let result = persist.load_bm25_index().await;
assert!(result.is_err(), "loading wrong schema should fail");
let err = result.unwrap_err();
assert!(matches!(err, PersistError::Deserialize(_)));
}
#[tokio::test]
async fn test_valid_json_wrong_schema_hnsw() {
let persist = setup_test_persistence().await;
{
let conn = persist.conn.clone();
let namespace = "test".to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
let wrong_schema = br#"{"total_nodes": 5, "wrong_field": true}"#;
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, 'hnsw', ?2, strftime('%s', 'now'))
"#,
rusqlite::params![namespace, wrong_schema.as_slice()],
)
.expect("insert wrong schema");
})
.await
.expect("spawn");
}
let result = persist.load_hnsw_snapshot().await;
assert!(result.is_err(), "loading wrong schema should fail");
let err = result.unwrap_err();
assert!(matches!(err, PersistError::Deserialize(_)));
}
#[tokio::test]
async fn test_empty_blob_returns_error() {
let persist = setup_test_persistence().await;
{
let conn = persist.conn.clone();
let namespace = "test".to_string();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, 'bm25', ?2, strftime('%s', 'now'))
"#,
rusqlite::params![namespace, &[] as &[u8]],
)
.expect("insert empty blob");
})
.await
.expect("spawn");
}
let result = persist.load_bm25_index().await;
assert!(result.is_err(), "loading empty blob should fail");
let err = result.unwrap_err();
assert!(matches!(err, PersistError::Deserialize(_)));
}
#[tokio::test]
async fn test_empty_bm25_index_persistence() {
let persist = setup_test_persistence().await;
let index = Bm25Index::default();
assert_eq!(index.doc_count(), 0);
persist
.persist_bm25_index(&index)
.await
.expect("persist empty");
let loaded = persist.load_bm25_index().await.expect("load");
assert!(loaded.is_some());
let loaded = loaded.unwrap();
assert_eq!(loaded.doc_count(), 0, "empty index should remain empty");
}
#[tokio::test]
async fn test_empty_hnsw_index_persistence() {
let persist = setup_test_persistence().await;
let index = HnswIndex::new(4);
assert_eq!(index.len(), 0);
persist
.persist_hnsw_snapshot(&index)
.await
.expect("persist empty");
let loaded = persist.load_hnsw_snapshot().await.expect("load");
assert!(loaded.is_some());
let snapshot = loaded.unwrap();
assert_eq!(
snapshot.total_nodes, 0,
"empty index snapshot should have 0 nodes"
);
assert_eq!(snapshot.live_nodes, 0);
assert!(snapshot.indexed_ids.is_empty());
}
#[tokio::test]
async fn test_empty_hnsw_shadow_validation() {
let persist = setup_test_persistence().await;
let config = ShadowValidationConfig::enabled();
let index = HnswIndex::new(4);
let result = persist
.persist_hnsw_with_validation(&index, &config)
.await
.expect("persist empty with validation");
assert!(result.is_some());
let validation = result.unwrap();
assert!(validation.passed, "empty index validation should pass");
assert_eq!(validation.expected.item_count, 0);
}
#[tokio::test]
async fn test_hnsw_shadow_validation_calls_verify() {
let persist = setup_test_persistence().await;
let config = ShadowValidationConfig::enabled();
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
index.insert(id2, vec![0.0, 1.0, 0.0, 0.0]).expect("insert");
index.insert(id3, vec![0.0, 0.0, 1.0, 0.0]).expect("insert");
index.delete(id2);
let result = persist
.persist_hnsw_with_validation(&index, &config)
.await
.expect("persist with validation");
assert!(result.is_some());
let validation = result.unwrap();
assert!(
validation.passed,
"valid snapshot should pass verify(): {:?}",
validation.discrepancies
);
assert_eq!(validation.expected.item_count, 3);
assert_eq!(validation.expected.tombstone_count, 1);
}
async fn inject_raw_hnsw_snapshot(persist: &RetrievalPersistence, data: &[u8]) {
let conn = persist.conn.clone();
let namespace = persist.namespace.clone();
let data = data.to_vec();
tokio::task::spawn_blocking(move || {
let conn = conn.blocking_lock();
conn.execute(
r#"
INSERT OR REPLACE INTO retrieval_snapshots
(namespace, index_type, snapshot, created_at)
VALUES
(?1, 'hnsw', ?2, strftime('%s', 'now'))
"#,
rusqlite::params![&*namespace, data],
)
.expect("inject raw snapshot");
})
.await
.expect("spawn");
}
async fn build_and_persist_hnsw(persist: &RetrievalPersistence) -> HnswIndex {
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
index.insert(id2, vec![0.0, 1.0, 0.0, 0.0]).expect("insert");
index.insert(id3, vec![0.0, 0.0, 1.0, 0.0]).expect("insert");
persist
.persist_hnsw_snapshot(&index)
.await
.expect("persist");
index
}
#[tokio::test]
async fn test_truncated_hnsw_snapshot_detected() {
let persist = setup_test_persistence().await;
build_and_persist_hnsw(&persist).await;
let valid_snapshot = persist
.load_hnsw_snapshot()
.await
.expect("load valid")
.expect("snapshot exists");
let valid_json = serde_json::to_vec(&valid_snapshot).expect("serialize");
assert!(valid_json.len() > 20, "valid JSON should be non-trivial");
for truncate_at in [1, 10, valid_json.len() / 4, valid_json.len() / 2] {
let truncated = &valid_json[..truncate_at];
inject_raw_hnsw_snapshot(&persist, truncated).await;
let result = persist.load_hnsw_snapshot().await;
assert!(
result.is_err(),
"truncated snapshot (at byte {truncate_at}) should fail to load"
);
let err = result.unwrap_err();
assert!(
matches!(err, PersistError::Deserialize(_)),
"should be a Deserialize error, got: {err:?}"
);
}
}
#[tokio::test]
async fn test_corrupted_bytes_in_hnsw_snapshot_detected() {
let persist = setup_test_persistence().await;
build_and_persist_hnsw(&persist).await;
let valid_snapshot = persist
.load_hnsw_snapshot()
.await
.expect("load valid")
.expect("snapshot exists");
let mut corrupted_json = serde_json::to_vec(&valid_snapshot).expect("serialize");
let mid = corrupted_json.len() / 2;
for i in mid..mid.saturating_add(10).min(corrupted_json.len()) {
corrupted_json[i] = 0xFF;
}
inject_raw_hnsw_snapshot(&persist, &corrupted_json).await;
let result = persist.load_hnsw_snapshot().await;
match result {
Err(PersistError::Deserialize(_)) => {
}
Ok(Some(snapshot)) => {
let verify_result = snapshot.verify();
let _ = verify_result;
}
Ok(None) => {
panic!("snapshot was injected, should not return None");
}
Err(other) => {
panic!("unexpected error variant: {other:?}");
}
}
}
#[tokio::test]
async fn test_missing_hnsw_snapshot_returns_none() {
let persist = setup_test_persistence().await;
let result = persist
.load_hnsw_snapshot()
.await
.expect("load should not error");
assert!(
result.is_none(),
"missing snapshot should return None, not error"
);
}
#[tokio::test]
async fn test_missing_hnsw_snapshot_after_clear_returns_none() {
let persist = setup_test_persistence().await;
build_and_persist_hnsw(&persist).await;
let loaded = persist.load_hnsw_snapshot().await.expect("load");
assert!(loaded.is_some(), "snapshot should exist before clear");
persist.clear().await.expect("clear");
let after_clear = persist
.load_hnsw_snapshot()
.await
.expect("load after clear");
assert!(
after_clear.is_none(),
"snapshot should be None after clear, enabling rebuild from source"
);
}
#[tokio::test]
async fn test_inconsistent_hnsw_snapshot_detected_by_verify() {
use khive_hnsw::{HnswCheckpointConfig, HnswSnapshot};
let persist = setup_test_persistence().await;
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let bad_snapshot = HnswSnapshot {
vector_count: 0,
total_nodes: 5, live_nodes: 5,
tombstone_count: 0,
max_layer: 0,
entry_point: Some(id1),
config: HnswCheckpointConfig {
m: 16,
ef_construction: 200,
metric: "cosine".to_string(),
},
indexed_ids: vec![id1, id2], tombstoned_ids: vec![],
layers: vec![vec![(id1, vec![id2]), (id2, vec![id1])]],
vectors: vec![],
};
let data = serde_json::to_vec(&bad_snapshot).expect("serialize");
inject_raw_hnsw_snapshot(&persist, &data).await;
let load_err = persist
.load_hnsw_snapshot()
.await
.expect_err("load should fail verification at deserialization time");
let err_msg = load_err.to_string();
assert!(
err_msg.contains("indexed_ids count mismatch") || err_msg.contains("inconsistent counts"),
"error should describe total_nodes != indexed_ids.len(), got: {err_msg}"
);
}
#[tokio::test]
async fn test_tombstone_inconsistency_detected_by_verify() {
use khive_hnsw::{HnswCheckpointConfig, HnswSnapshot};
let persist = setup_test_persistence().await;
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
let id_phantom = NodeId::new([99; 16]);
let bad_snapshot = HnswSnapshot {
vector_count: 0,
total_nodes: 3,
live_nodes: 2,
tombstone_count: 1,
max_layer: 0,
entry_point: Some(id1),
config: HnswCheckpointConfig {
m: 16,
ef_construction: 200,
metric: "cosine".to_string(),
},
indexed_ids: vec![id1, id2, id3],
tombstoned_ids: vec![id_phantom], layers: vec![],
vectors: vec![],
};
let data = serde_json::to_vec(&bad_snapshot).expect("serialize");
inject_raw_hnsw_snapshot(&persist, &data).await;
let load_err = persist
.load_hnsw_snapshot()
.await
.expect_err("load should fail verification at deserialization time");
let err_msg = load_err.to_string();
assert!(
err_msg.contains("tombstoned id") || err_msg.contains("not found in indexed_ids"),
"error should describe tombstoned ID not in indexed_ids, got: {err_msg}"
);
}
#[tokio::test]
async fn test_shadow_validation_detects_corruption() {
let persist = setup_test_persistence().await;
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
index.insert(id2, vec![0.0, 1.0, 0.0, 0.0]).expect("insert");
persist
.persist_hnsw_snapshot(&index)
.await
.expect("persist");
inject_raw_hnsw_snapshot(&persist, b"not valid json at all {{{").await;
let expected = ShadowMetrics {
item_count: 2,
tombstone_count: 0,
snapshot_size: 0,
};
let result = persist.validate_hnsw_snapshot(expected).await;
assert!(
!result.passed,
"shadow validation should fail on corrupted data"
);
assert!(
!result.discrepancies.is_empty(),
"should report discrepancies"
);
}
#[tokio::test]
async fn test_full_recovery_workflow_corrupt_then_rebuild() {
let persist = setup_test_persistence().await;
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
let vectors = vec![
(id1, vec![1.0, 0.0, 0.0, 0.0]),
(id2, vec![0.0, 1.0, 0.0, 0.0]),
(id3, vec![0.0, 0.0, 1.0, 0.0]),
];
{
let mut index = HnswIndex::new(4);
for (id, vec) in &vectors {
index.insert(*id, vec.clone()).expect("insert");
}
persist
.persist_hnsw_snapshot(&index)
.await
.expect("persist");
}
inject_raw_hnsw_snapshot(&persist, b"corrupted snapshot data").await;
let load_result = persist.load_hnsw_snapshot().await;
assert!(
load_result.is_err(),
"loading corrupted snapshot should fail"
);
persist.clear().await.expect("clear corrupted data");
let after_clear = persist
.load_hnsw_snapshot()
.await
.expect("load after clear");
assert!(after_clear.is_none(), "snapshot should be gone after clear");
let mut rebuilt_index = HnswIndex::new(4);
for (id, vec) in &vectors {
rebuilt_index.insert(*id, vec.clone()).expect("re-insert");
}
assert_eq!(
rebuilt_index.len(),
3,
"rebuilt index should have 3 vectors"
);
persist
.persist_hnsw_snapshot(&rebuilt_index)
.await
.expect("persist rebuilt");
let new_snapshot = persist
.load_hnsw_snapshot()
.await
.expect("load rebuilt")
.expect("snapshot exists");
assert_eq!(new_snapshot.total_nodes, 3);
assert_eq!(new_snapshot.live_nodes, 3);
assert!(
new_snapshot.verify().is_ok(),
"rebuilt snapshot should pass verification"
);
}
#[tokio::test]
async fn test_recovery_from_inconsistent_snapshot_via_verify() {
use khive_hnsw::{HnswCheckpointConfig, HnswSnapshot};
let persist = setup_test_persistence().await;
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let id3 = NodeId::new([3; 16]);
let bad_snapshot = HnswSnapshot {
vector_count: 0,
total_nodes: 3,
live_nodes: 1,
tombstone_count: 2, max_layer: 0,
entry_point: Some(id1),
config: HnswCheckpointConfig {
m: 16,
ef_construction: 200,
metric: "cosine".to_string(),
},
indexed_ids: vec![id1, id2, id3],
tombstoned_ids: vec![id2], layers: vec![],
vectors: vec![],
};
let data = serde_json::to_vec(&bad_snapshot).expect("serialize");
inject_raw_hnsw_snapshot(&persist, &data).await;
let load_err = persist
.load_hnsw_snapshot()
.await
.expect_err("load should fail verification at deserialization time");
let err_msg = load_err.to_string();
assert!(
err_msg.contains("tombstoned_ids count mismatch"),
"should report tombstone count mismatch, got: {err_msg}"
);
persist.clear().await.expect("clear");
let mut rebuilt = HnswIndex::new(4);
rebuilt
.insert(id1, vec![1.0, 0.0, 0.0, 0.0])
.expect("insert");
rebuilt
.insert(id2, vec![0.0, 1.0, 0.0, 0.0])
.expect("insert");
rebuilt
.insert(id3, vec![0.0, 0.0, 1.0, 0.0])
.expect("insert");
persist
.persist_hnsw_snapshot(&rebuilt)
.await
.expect("persist rebuilt");
let new_snapshot = persist
.load_hnsw_snapshot()
.await
.expect("load")
.expect("snapshot exists");
assert!(
new_snapshot.verify().is_ok(),
"rebuilt snapshot should be valid"
);
}
#[tokio::test]
async fn test_restore_from_corrupt_snapshot_fails() {
use khive_hnsw::{HnswCheckpointConfig, HnswSnapshot};
let id1 = NodeId::new([1; 16]);
let id2 = NodeId::new([2; 16]);
let mut index = HnswIndex::new(4);
let bad_snapshot = HnswSnapshot {
vector_count: 0,
total_nodes: 10, live_nodes: 10,
tombstone_count: 0,
max_layer: 0,
entry_point: Some(id1),
config: HnswCheckpointConfig {
m: 16,
ef_construction: 200,
metric: "cosine".to_string(),
},
indexed_ids: vec![id1, id2], tombstoned_ids: vec![],
layers: vec![],
vectors: vec![],
};
let vectors: std::collections::HashMap<NodeId, Vec<f32>> = [
(id1, vec![1.0, 0.0, 0.0, 0.0]),
(id2, vec![0.0, 1.0, 0.0, 0.0]),
]
.into_iter()
.collect();
let result = index.restore_from_snapshot(&bad_snapshot, &vectors);
assert!(
result.is_err(),
"restore_from_snapshot should reject corrupt snapshot"
);
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("Invalid snapshot"),
"error should mention invalid snapshot, got: {err_msg}"
);
}
#[tokio::test]
async fn test_binary_garbage_hnsw_snapshot_detected() {
let persist = setup_test_persistence().await;
let garbage: Vec<u8> = (0..256).map(|i| i as u8).collect();
inject_raw_hnsw_snapshot(&persist, &garbage).await;
let result = persist.load_hnsw_snapshot().await;
assert!(result.is_err(), "binary garbage should fail to deserialize");
let err = result.unwrap_err();
assert!(
matches!(err, PersistError::Deserialize(_)),
"should be Deserialize error, got: {err:?}"
);
}
#[tokio::test]
async fn test_overwrite_corrupt_snapshot_with_valid() {
let persist = setup_test_persistence().await;
inject_raw_hnsw_snapshot(&persist, b"this is not valid json").await;
assert!(persist.load_hnsw_snapshot().await.is_err());
let mut index = HnswIndex::new(4);
let id1 = NodeId::new([1; 16]);
index.insert(id1, vec![1.0, 0.0, 0.0, 0.0]).expect("insert");
persist
.persist_hnsw_snapshot(&index)
.await
.expect("persist should overwrite corrupt entry");
let loaded = persist
.load_hnsw_snapshot()
.await
.expect("load should succeed after overwrite")
.expect("snapshot should exist");
assert_eq!(loaded.total_nodes, 1);
assert!(loaded.verify().is_ok());
}