use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use crate::blob_store::BlobStore;
use crate::index::{SparseIndex, SparseVector};
use crate::mmap_index::{self, MmapPostingData};
use crate::wand::Postings;
const MMAP_FILE: &str = "sparse.mmap";
const VECTORS_FILE: &str = "sparse_vectors.bin";
const DIMS_FILE: &str = "sparse_dims.bin";
const LEGACY_FILE: &str = "sparse.bin";
const INDEX_FILES: &[&str] = &[MMAP_FILE, VECTORS_FILE, DIMS_FILE];
const BLOB_PREFIX: &str = "Sparse_";
static CACHE_SEQ: AtomicUsize = AtomicUsize::new(0);
enum StorageBackend {
Filesystem,
Store {
store: Arc<dyn BlobStore>,
index_name: String,
},
}
struct Inner {
index: SparseIndex,
mmap: Option<MmapPostingData>,
postings_loaded: bool,
vectors_loaded: bool,
num_vectors: usize,
dirty: bool,
}
pub struct SparseHandle {
inner: Mutex<Inner>,
path: PathBuf,
backend: StorageBackend,
}
impl SparseHandle {
pub fn create(path: &str) -> Result<Self, String> {
std::fs::create_dir_all(Path::new(path))
.map_err(|e| format!("cannot create directory {path}: {e}"))?;
let handle = Self {
inner: Mutex::new(Inner {
index: SparseIndex::new(),
mmap: None,
postings_loaded: true,
vectors_loaded: true,
num_vectors: 0,
dirty: false,
}),
path: PathBuf::from(path),
backend: StorageBackend::Filesystem,
};
handle.commit_inner()?;
Ok(handle)
}
pub fn open(path: &str) -> Result<Self, String> {
let base = Path::new(path);
let mmap_path = base.join(MMAP_FILE);
if mmap_path.exists() {
Self::open_mmap(base, StorageBackend::Filesystem)
} else {
Self::open_legacy(base)
}
}
pub fn create_with_store(
store: Arc<dyn BlobStore>,
index_name: &str,
cache_base: &Path,
) -> Result<Self, String> {
let blob_name = format!("{BLOB_PREFIX}{index_name}");
let cache_dir = Self::make_cache_dir(cache_base, &blob_name)?;
let handle = Self {
inner: Mutex::new(Inner {
index: SparseIndex::new(),
mmap: None,
postings_loaded: true,
vectors_loaded: true,
num_vectors: 0,
dirty: false,
}),
path: cache_dir,
backend: StorageBackend::Store {
store,
index_name: blob_name,
},
};
handle.commit_inner()?;
Ok(handle)
}
pub fn open_with_store(
store: Arc<dyn BlobStore>,
index_name: &str,
cache_base: &Path,
) -> Result<Self, String> {
let blob_name = format!("{BLOB_PREFIX}{index_name}");
let cache_dir = Self::make_cache_dir(cache_base, &blob_name)?;
let files = store
.list(&blob_name)
.map_err(|e| format!("cannot list blobs for {blob_name}: {e}"))?;
for file_name in &files {
let data = store
.load(&blob_name, file_name)
.map_err(|e| format!("cannot load {blob_name}/{file_name}: {e}"))?;
std::fs::write(cache_dir.join(file_name), data)
.map_err(|e| format!("cannot write cache {file_name}: {e}"))?;
}
let backend = StorageBackend::Store {
store,
index_name: blob_name,
};
if cache_dir.join(MMAP_FILE).exists() {
Self::open_mmap(&cache_dir, backend)
} else {
let handle = Self {
inner: Mutex::new(Inner {
index: SparseIndex::new(),
mmap: None,
postings_loaded: true,
vectors_loaded: true,
num_vectors: 0,
dirty: false,
}),
path: cache_dir,
backend,
};
handle.commit_inner()?;
Ok(handle)
}
}
fn make_cache_dir(base: &Path, index_name: &str) -> Result<PathBuf, String> {
let seq = CACHE_SEQ.fetch_add(1, Ordering::Relaxed);
let pid = std::process::id();
let dir = base
.join(format!("{pid}"))
.join(format!("{index_name}_{seq}"));
std::fs::create_dir_all(&dir)
.map_err(|e| format!("cannot create cache dir {}: {e}", dir.display()))?;
Ok(dir)
}
fn open_mmap(base: &Path, backend: StorageBackend) -> Result<Self, String> {
let mmap = MmapPostingData::open(&base.join(MMAP_FILE))?;
let dims_data = std::fs::read(base.join(DIMS_FILE))
.map_err(|e| format!("cannot read {DIMS_FILE}: {e}"))?;
let (dim_map, dim_reverse): (HashMap<u32, usize>, Vec<u32>) =
bincode::deserialize(&dims_data)
.map_err(|e| format!("cannot deserialize dims: {e}"))?;
let num_dims = mmap.num_dims();
let num_vectors = mmap.num_vectors();
let empty_postings: Vec<Postings> = (0..num_dims).map(|_| Postings::new()).collect();
let index =
SparseIndex::from_parts(dim_map, dim_reverse, empty_postings, HashMap::new());
Ok(Self {
inner: Mutex::new(Inner {
index,
mmap: Some(mmap),
postings_loaded: false,
vectors_loaded: false,
num_vectors,
dirty: false,
}),
path: base.to_path_buf(),
backend,
})
}
fn open_legacy(base: &Path) -> Result<Self, String> {
let data_path = base.join(LEGACY_FILE);
let data = std::fs::read(&data_path)
.map_err(|e| format!("cannot read {}: {e}", data_path.display()))?;
let index: SparseIndex = bincode::deserialize(&data)
.map_err(|e| format!("cannot deserialize sparse index: {e}"))?;
let num_vectors = index.len();
Ok(Self {
inner: Mutex::new(Inner {
index,
mmap: None,
postings_loaded: true,
vectors_loaded: true,
num_vectors,
dirty: false,
}),
path: base.to_path_buf(),
backend: StorageBackend::Filesystem,
})
}
fn ensure_postings_loaded(inner: &mut Inner) {
if inner.postings_loaded {
return;
}
if let Some(ref mmap) = inner.mmap {
let postings = inner.index.postings_mut();
for (i, pl) in postings.iter_mut().enumerate() {
*pl = mmap.load_postings(i);
}
}
inner.postings_loaded = true;
}
fn ensure_vectors_loaded(inner: &mut Inner, path: &Path) -> Result<(), String> {
if inner.vectors_loaded {
return Ok(());
}
let vectors_path = path.join(VECTORS_FILE);
let data = std::fs::read(&vectors_path)
.map_err(|e| format!("cannot read {}: {e}", vectors_path.display()))?;
let vectors: HashMap<u64, SparseVector> = bincode::deserialize(&data)
.map_err(|e| format!("cannot deserialize vectors: {e}"))?;
inner.index.set_vectors(vectors);
inner.vectors_loaded = true;
Ok(())
}
pub fn insert(&self, node_id: u64, vector: &SparseVector) -> Result<(), String> {
let mut inner = self.inner.lock().map_err(|_| "lock poisoned".to_string())?;
Self::ensure_vectors_loaded(&mut inner, &self.path)?;
Self::ensure_postings_loaded(&mut inner);
inner.index.insert(node_id, vector);
inner.num_vectors = inner.index.len();
inner.dirty = true;
Ok(())
}
pub fn remove(&self, node_id: u64) -> Result<bool, String> {
let mut inner = self.inner.lock().map_err(|_| "lock poisoned".to_string())?;
Self::ensure_vectors_loaded(&mut inner, &self.path)?;
Self::ensure_postings_loaded(&mut inner);
let removed = inner.index.remove(node_id);
if removed {
inner.num_vectors = inner.index.len();
inner.dirty = true;
}
Ok(removed)
}
pub fn search(&self, query: &SparseVector, limit: usize) -> Vec<(u64, f32)> {
let inner = self.inner.lock().unwrap();
if !inner.dirty {
if let Some(ref mmap) = inner.mmap {
return mmap_index::search_mmap(
mmap,
inner.index.dim_map(),
query,
limit,
&|_| true,
);
}
}
inner.index.search(query, limit)
}
pub fn search_filtered(
&self,
query: &SparseVector,
limit: usize,
allowed_ids: &[u64],
) -> Vec<(u64, f32)> {
let inner = self.inner.lock().unwrap();
if !inner.dirty {
if let Some(ref mmap) = inner.mmap {
return mmap_index::search_mmap_allowed(
mmap,
inner.index.dim_map(),
query,
limit,
allowed_ids,
);
}
}
inner.index.search_filtered(query, limit, allowed_ids)
}
pub fn len(&self) -> usize {
self.inner.lock().unwrap().num_vectors
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn commit_inner(&self) -> Result<(), String> {
let mut inner = self.inner.lock().map_err(|_| "lock poisoned".to_string())?;
if !inner.postings_loaded && !inner.dirty && inner.mmap.is_some() {
return Ok(());
}
Self::ensure_postings_loaded(&mut inner);
Self::ensure_vectors_loaded(&mut inner, &self.path)?;
mmap_index::write_mmap_file(
&self.path.join(MMAP_FILE),
inner.index.postings(),
inner.num_vectors as u32,
)?;
let vectors_data = bincode::serialize(inner.index.vectors())
.map_err(|e| format!("cannot serialize vectors: {e}"))?;
std::fs::write(self.path.join(VECTORS_FILE), &vectors_data)
.map_err(|e| format!("cannot write {VECTORS_FILE}: {e}"))?;
let dims_data =
bincode::serialize(&(inner.index.dim_map(), inner.index.dim_reverse()))
.map_err(|e| format!("cannot serialize dims: {e}"))?;
std::fs::write(self.path.join(DIMS_FILE), &dims_data)
.map_err(|e| format!("cannot write {DIMS_FILE}: {e}"))?;
if let StorageBackend::Store { ref store, ref index_name } = self.backend {
for &file in INDEX_FILES {
let data = std::fs::read(self.path.join(file))
.map_err(|e| format!("cannot read cache {file}: {e}"))?;
store
.save(index_name, file, &data)
.map_err(|e| format!("cannot save {index_name}/{file} to store: {e}"))?;
}
}
let mmap = MmapPostingData::open(&self.path.join(MMAP_FILE))?;
inner.mmap = Some(mmap);
inner.dirty = false;
let legacy = self.path.join(LEGACY_FILE);
if legacy.exists() {
let _ = std::fs::remove_file(&legacy);
if let StorageBackend::Store { ref store, ref index_name } = self.backend {
let _ = store.delete(index_name, LEGACY_FILE);
}
}
Ok(())
}
}
impl Drop for SparseHandle {
fn drop(&mut self) {
if let StorageBackend::Store { .. } = &self.backend {
let _ = std::fs::remove_dir_all(&self.path);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::blob_store::MemBlobStore;
fn tmp_path(name: &str) -> PathBuf {
std::env::temp_dir().join(name)
}
fn cleanup(path: &Path) {
let _ = std::fs::remove_dir_all(path);
}
#[test]
fn create_writes_mmap_format() {
let p = tmp_path("sparse_mmap_create_test");
cleanup(&p);
let path = p.to_str().unwrap();
let _handle = SparseHandle::create(path).unwrap();
assert!(p.join(MMAP_FILE).exists());
assert!(p.join(VECTORS_FILE).exists());
assert!(p.join(DIMS_FILE).exists());
let handle2 = SparseHandle::open(path).unwrap();
assert_eq!(handle2.len(), 0);
cleanup(&p);
}
#[test]
fn persistence_roundtrip_mmap() {
let p = tmp_path("sparse_mmap_roundtrip_test");
cleanup(&p);
let path = p.to_str().unwrap();
let handle = SparseHandle::create(path).unwrap();
handle
.insert(42, &SparseVector::new(vec![1, 2], vec![0.5, 0.3]))
.unwrap();
handle
.insert(99, &SparseVector::new(vec![2, 3], vec![0.8, 0.2]))
.unwrap();
handle.commit_inner().unwrap();
drop(handle);
let handle2 = SparseHandle::open(path).unwrap();
assert_eq!(handle2.len(), 2);
let results = handle2.search(&SparseVector::new(vec![2], vec![1.0]), 10);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 99);
assert!((results[0].1 - 0.8).abs() < 1e-6);
assert_eq!(results[1].0, 42);
assert!((results[1].1 - 0.3).abs() < 1e-6);
cleanup(&p);
}
#[test]
fn mmap_search_filtered() {
let p = tmp_path("sparse_mmap_filtered_test");
cleanup(&p);
let path = p.to_str().unwrap();
let handle = SparseHandle::create(path).unwrap();
handle
.insert(1, &SparseVector::new(vec![1, 2], vec![0.5, 0.3]))
.unwrap();
handle
.insert(2, &SparseVector::new(vec![1, 3], vec![0.9, 0.1]))
.unwrap();
handle
.insert(3, &SparseVector::new(vec![1], vec![0.7]))
.unwrap();
handle.commit_inner().unwrap();
drop(handle);
let handle2 = SparseHandle::open(path).unwrap();
let results = handle2.search_filtered(&SparseVector::new(vec![1], vec![1.0]), 10, &[1, 3]);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 3); assert_eq!(results[1].0, 1);
cleanup(&p);
}
#[test]
fn mutation_after_mmap_open() {
let p = tmp_path("sparse_mmap_mutation_test");
cleanup(&p);
let path = p.to_str().unwrap();
let handle = SparseHandle::create(path).unwrap();
handle
.insert(1, &SparseVector::new(vec![10], vec![1.0]))
.unwrap();
handle.commit_inner().unwrap();
drop(handle);
let handle2 = SparseHandle::open(path).unwrap();
handle2
.insert(2, &SparseVector::new(vec![10], vec![2.0]))
.unwrap();
let results = handle2.search(&SparseVector::new(vec![10], vec![1.0]), 10);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 2); assert_eq!(results[1].0, 1);
handle2.commit_inner().unwrap();
drop(handle2);
let handle3 = SparseHandle::open(path).unwrap();
let results = handle3.search(&SparseVector::new(vec![10], vec![1.0]), 10);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 2);
cleanup(&p);
}
#[test]
fn legacy_fallback() {
let p = tmp_path("sparse_mmap_legacy_test");
cleanup(&p);
let path = p.to_str().unwrap();
std::fs::create_dir_all(&p).unwrap();
let mut index = SparseIndex::new();
index.insert(7, &SparseVector::new(vec![1], vec![0.42]));
let data = bincode::serialize(&index).unwrap();
std::fs::write(p.join(LEGACY_FILE), data).unwrap();
let handle = SparseHandle::open(path).unwrap();
assert_eq!(handle.len(), 1);
let results = handle.search(&SparseVector::new(vec![1], vec![1.0]), 10);
assert_eq!(results[0].0, 7);
handle.commit_inner().unwrap();
assert!(p.join(MMAP_FILE).exists());
assert!(!p.join(LEGACY_FILE).exists());
cleanup(&p);
}
#[test]
fn many_docs_mmap_roundtrip() {
let p = tmp_path("sparse_mmap_many_docs_test");
cleanup(&p);
let path = p.to_str().unwrap();
let handle = SparseHandle::create(path).unwrap();
for i in 0..500u64 {
let token = (i % 50) as u32;
let weight = (i as f32) / 500.0;
handle
.insert(
i,
&SparseVector::new(vec![token, token + 50], vec![weight, weight * 0.5]),
)
.unwrap();
}
handle.commit_inner().unwrap();
drop(handle);
let handle2 = SparseHandle::open(path).unwrap();
assert_eq!(handle2.len(), 500);
let results = handle2.search(&SparseVector::new(vec![0, 50], vec![1.0, 1.0]), 5);
assert_eq!(results.len(), 5);
assert_eq!(results[0].0, 450);
cleanup(&p);
}
fn test_cache_base() -> PathBuf {
std::env::temp_dir().join("sparse_test_cache")
}
#[test]
fn blob_store_create_and_search() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
let handle = SparseHandle::create_with_store(store.clone(), "test_idx", &cb).unwrap();
handle
.insert(42, &SparseVector::new(vec![1, 2], vec![0.5, 0.3]))
.unwrap();
handle
.insert(99, &SparseVector::new(vec![2, 3], vec![0.8, 0.2]))
.unwrap();
handle.commit_inner().unwrap();
assert!(store.exists("Sparse_test_idx", MMAP_FILE).unwrap());
assert!(store.exists("Sparse_test_idx", VECTORS_FILE).unwrap());
assert!(store.exists("Sparse_test_idx", DIMS_FILE).unwrap());
assert_eq!(store.list("Sparse_test_idx").unwrap().len(), 3);
let results = handle.search(&SparseVector::new(vec![2], vec![1.0]), 10);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 99);
}
#[test]
fn blob_store_close_reopen() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
{
let handle = SparseHandle::create_with_store(store.clone(), "reopen_idx", &cb).unwrap();
handle
.insert(1, &SparseVector::new(vec![10], vec![1.0]))
.unwrap();
handle
.insert(2, &SparseVector::new(vec![10, 20], vec![0.5, 0.8]))
.unwrap();
handle.commit_inner().unwrap();
}
let handle2 = SparseHandle::open_with_store(store.clone(), "reopen_idx", &cb).unwrap();
assert_eq!(handle2.len(), 2);
let results = handle2.search(&SparseVector::new(vec![10], vec![1.0]), 10);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 1); assert_eq!(results[1].0, 2); }
#[test]
fn blob_store_mutation_after_reopen() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
{
let handle = SparseHandle::create_with_store(store.clone(), "mut_idx", &cb).unwrap();
handle
.insert(1, &SparseVector::new(vec![5], vec![1.0]))
.unwrap();
handle.commit_inner().unwrap();
}
let handle2 = SparseHandle::open_with_store(store.clone(), "mut_idx", &cb).unwrap();
handle2
.insert(2, &SparseVector::new(vec![5], vec![2.0]))
.unwrap();
handle2.commit_inner().unwrap();
drop(handle2);
let handle3 = SparseHandle::open_with_store(store.clone(), "mut_idx", &cb).unwrap();
assert_eq!(handle3.len(), 2);
let results = handle3.search(&SparseVector::new(vec![5], vec![1.0]), 10);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 2); assert_eq!(results[1].0, 1); }
#[test]
fn blob_store_delete_and_reopen() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
{
let handle = SparseHandle::create_with_store(store.clone(), "del_idx", &cb).unwrap();
handle
.insert(1, &SparseVector::new(vec![1], vec![1.0]))
.unwrap();
handle
.insert(2, &SparseVector::new(vec![1], vec![2.0]))
.unwrap();
handle.commit_inner().unwrap();
}
let handle2 = SparseHandle::open_with_store(store.clone(), "del_idx", &cb).unwrap();
assert_eq!(handle2.len(), 2);
handle2.remove(1).unwrap();
assert_eq!(handle2.len(), 1);
handle2.commit_inner().unwrap();
drop(handle2);
let handle3 = SparseHandle::open_with_store(store.clone(), "del_idx", &cb).unwrap();
assert_eq!(handle3.len(), 1);
let results = handle3.search(&SparseVector::new(vec![1], vec![1.0]), 10);
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, 2);
}
#[test]
fn blob_store_multiple_indexes_isolated() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
let h1 = SparseHandle::create_with_store(store.clone(), "idx_a", &cb).unwrap();
let h2 = SparseHandle::create_with_store(store.clone(), "idx_b", &cb).unwrap();
h1.insert(1, &SparseVector::new(vec![1], vec![1.0]))
.unwrap();
h1.insert(2, &SparseVector::new(vec![1], vec![0.5]))
.unwrap();
h2.insert(10, &SparseVector::new(vec![1], vec![3.0]))
.unwrap();
h1.commit_inner().unwrap();
h2.commit_inner().unwrap();
assert_eq!(h1.len(), 2);
assert_eq!(h2.len(), 1);
assert_eq!(store.list("Sparse_idx_a").unwrap().len(), 3);
assert_eq!(store.list("Sparse_idx_b").unwrap().len(), 3);
}
#[test]
fn blob_store_survives_cache_cleanup() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
{
let handle = SparseHandle::create_with_store(store.clone(), "surv_idx", &cb).unwrap();
for i in 0..50u64 {
handle
.insert(i, &SparseVector::new(vec![(i % 10) as u32], vec![i as f32]))
.unwrap();
}
handle.commit_inner().unwrap();
}
let handle2 = SparseHandle::open_with_store(store.clone(), "surv_idx", &cb).unwrap();
assert_eq!(handle2.len(), 50);
let results = handle2.search(&SparseVector::new(vec![0], vec![1.0]), 5);
assert_eq!(results.len(), 4);
assert_eq!(results[0].0, 40);
}
#[test]
fn blob_store_search_filtered_after_reopen() {
let store = Arc::new(MemBlobStore::new());
let cb = test_cache_base();
{
let handle = SparseHandle::create_with_store(store.clone(), "filt_idx", &cb).unwrap();
handle
.insert(1, &SparseVector::new(vec![1, 2], vec![0.5, 0.3]))
.unwrap();
handle
.insert(2, &SparseVector::new(vec![1, 3], vec![0.9, 0.1]))
.unwrap();
handle
.insert(3, &SparseVector::new(vec![1], vec![0.7]))
.unwrap();
handle.commit_inner().unwrap();
}
let handle2 = SparseHandle::open_with_store(store.clone(), "filt_idx", &cb).unwrap();
let results =
handle2.search_filtered(&SparseVector::new(vec![1], vec![1.0]), 10, &[1, 3]);
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 3); assert_eq!(results[1].0, 1); }
}