use std::collections::HashMap;
use std::path::Path;
use memmap2::Mmap;
use crate::index::{run_search, SparseVector};
use crate::wand::{MmapCursor, Postings};
const MAGIC: u32 = 0x53505253; const FORMAT_VERSION: u32 = 3;
const GLOBAL_DIMS_VERSION: u32 = 3;
const MIN_READABLE_VERSION: u32 = 1;
const FOOTER_SIZE: usize = 8;
#[repr(C)]
struct FileHeader {
magic: u32,
version: u32,
num_dims: u32,
num_vectors: u32,
}
#[repr(C)]
struct DimHeader {
offset: u64,
count: u32,
token_id: u32,
}
#[repr(C)]
#[derive(Clone, Copy, Debug)]
pub struct PostingEntry {
pub record_id: u64,
pub weight: f32,
pub max_next_weight: f32,
}
pub struct MmapPostingData {
mmap: Mmap,
num_dims: usize,
num_vectors: usize,
version: u32,
}
impl MmapPostingData {
pub fn open(path: &Path) -> Result<Self, String> {
let file = std::fs::File::open(path)
.map_err(|e| format!("cannot open {}: {e}", path.display()))?;
let mmap = unsafe { Mmap::map(&file) }
.map_err(|e| format!("cannot mmap {}: {e}", path.display()))?;
if mmap.len() < std::mem::size_of::<FileHeader>() {
return Err("sparse.mmap too small for header".into());
}
let header = unsafe { &*(mmap.as_ptr() as *const FileHeader) };
if header.magic != MAGIC {
return Err(format!("bad magic: {:#x}", header.magic));
}
if header.version < MIN_READABLE_VERSION || header.version > FORMAT_VERSION {
return Err(format!("unsupported version: {} (this build reads {MIN_READABLE_VERSION}..={FORMAT_VERSION})",
header.version));
}
let versioned = Self {
mmap,
num_dims: header.num_dims as usize,
num_vectors: header.num_vectors as usize,
version: header.version,
};
versioned.check_length(path)?;
if std::env::var("LUCIVY_SPARSE_VERIFY_CRC").is_ok_and(|v| v != "0") {
versioned.verify_checksum(path)?;
}
Ok(versioned)
}
fn expected_len(&self) -> usize {
let dim_headers = std::mem::size_of::<FileHeader>()
+ self.num_dims * std::mem::size_of::<DimHeader>();
let entries: usize = (0..self.num_dims)
.map(|i| {
let ptr = unsafe {
self.mmap.as_ptr().add(
std::mem::size_of::<FileHeader>() + i * std::mem::size_of::<DimHeader>(),
) as *const DimHeader
};
unsafe { (*ptr).count as usize }
})
.sum();
dim_headers + entries * std::mem::size_of::<PostingEntry>()
+ if self.version >= 2 { FOOTER_SIZE } else { 0 }
}
fn check_length(&self, path: &Path) -> Result<(), String> {
let headers_end = std::mem::size_of::<FileHeader>()
+ self.num_dims * std::mem::size_of::<DimHeader>();
if self.mmap.len() < headers_end {
return Err(format!(
"{}: truncated — {} bytes for {} dimension headers",
path.display(), self.mmap.len(), self.num_dims,
));
}
let expected = self.expected_len();
if self.mmap.len() != expected {
return Err(format!(
"{}: truncated or corrupt — {} bytes, its headers describe {expected}",
path.display(), self.mmap.len(),
));
}
Ok(())
}
pub fn verify_checksum(&self, path: &Path) -> Result<(), String> {
if self.version < 2 {
return Ok(());
}
let len = self.mmap.len();
let body = &self.mmap[..len - FOOTER_SIZE];
let stored = u32::from_le_bytes(self.mmap[len - FOOTER_SIZE..len - 4].try_into().unwrap());
let magic = u32::from_le_bytes(self.mmap[len - 4..].try_into().unwrap());
if magic != MAGIC {
return Err(format!("{}: footer magic is {magic:#x}", path.display()));
}
let mut hasher = crc32fast::Hasher::new();
hasher.update(body);
let actual = hasher.finalize();
if actual != stored {
return Err(format!("{}: checksum mismatch — {actual:#x}, the file says {stored:#x}",
path.display()));
}
Ok(())
}
pub fn num_dims(&self) -> usize {
self.num_dims
}
pub fn num_vectors(&self) -> usize {
self.num_vectors
}
pub fn version(&self) -> u32 {
self.version
}
pub fn has_global_dims(&self) -> bool {
self.version >= GLOBAL_DIMS_VERSION
}
fn dim_headers(&self) -> &[DimHeader] {
let ptr = unsafe { self.mmap.as_ptr().add(std::mem::size_of::<FileHeader>()) }
as *const DimHeader;
unsafe { std::slice::from_raw_parts(ptr, self.num_dims) }
}
pub fn dim_of_token(&self, token_id: u32) -> Option<usize> {
if !self.has_global_dims() {
return None;
}
self.dim_headers()
.binary_search_by_key(&token_id, |dh| dh.token_id)
.ok()
}
pub fn tokens(&self) -> impl Iterator<Item = (u32, usize)> + '_ {
let global = self.has_global_dims();
self.dim_headers().iter().enumerate()
.filter(move |_| global)
.map(|(i, dh)| (dh.token_id, i))
}
pub fn entries_of_token(&self, token_id: u32) -> &[PostingEntry] {
match self.dim_of_token(token_id) {
Some(i) => self.entries(i),
None => &[],
}
}
pub fn entries(&self, dim_idx: usize) -> &[PostingEntry] {
if dim_idx >= self.num_dims {
return &[];
}
let dim_headers_offset = std::mem::size_of::<FileHeader>();
let dh_ptr = unsafe {
self.mmap
.as_ptr()
.add(dim_headers_offset + dim_idx * std::mem::size_of::<DimHeader>())
} as *const DimHeader;
let dh = unsafe { &*dh_ptr };
if dh.count == 0 {
return &[];
}
let entries_ptr =
unsafe { self.mmap.as_ptr().add(dh.offset as usize) } as *const PostingEntry;
unsafe { std::slice::from_raw_parts(entries_ptr, dh.count as usize) }
}
pub fn cursor(&self, dim_idx: usize) -> Option<MmapCursor<'_>> {
MmapCursor::open(self, dim_idx as u32)
}
pub fn load_postings_of_token(&self, token_id: u32) -> Postings {
match self.dim_of_token(token_id) {
Some(i) => self.load_postings(i),
None => Postings::new(),
}
}
pub fn load_postings(&self, dim_idx: usize) -> Postings {
let pairs: Vec<(u64, f32)> = self
.entries(dim_idx)
.iter()
.map(|e| (e.record_id, e.weight))
.collect();
Postings::from_sorted_pairs(&pairs)
}
}
fn dims_for<'a>(
mmap: &'a MmapPostingData,
dim_map: &'a HashMap<u32, usize>,
query: &SparseVector,
) -> HashMap<u32, usize> {
if !mmap.has_global_dims() {
return query.indices.iter()
.filter_map(|t| dim_map.get(t).map(|&d| (*t, d)))
.collect();
}
query.indices.iter()
.filter_map(|&t| mmap.dim_of_token(t).map(|d| (t, d)))
.collect()
}
pub fn search_mmap<F: Fn(u64) -> bool>(
mmap: &MmapPostingData,
dim_map: &HashMap<u32, usize>,
query: &SparseVector,
limit: usize,
filter: &F,
) -> Vec<(u64, f32)> {
if query.is_empty() || mmap.num_vectors() == 0 {
return Vec::new();
}
let dims = dims_for(mmap, dim_map, query);
run_search(query, &dims, limit, filter, |dim| mmap.cursor(dim as usize))
}
pub fn search_mmap_allowed(
mmap: &MmapPostingData,
dim_map: &HashMap<u32, usize>,
query: &SparseVector,
limit: usize,
allowed: &[u64],
) -> Vec<(u64, f32)> {
if query.is_empty() || mmap.num_vectors() == 0 {
return Vec::new();
}
let dims = dims_for(mmap, dim_map, query);
crate::index::run_search_allowed(query, &dims, limit, allowed, |dim| mmap.cursor(dim as usize))
}
pub fn write_mmap_file(
path: &Path,
postings: &[Postings],
dim_tokens: &[u32],
num_vectors: u32,
) -> Result<(), String> {
use std::io::Write;
if dim_tokens.len() != postings.len() {
return Err(format!(
"{} posting lists for {} dimension ids", postings.len(), dim_tokens.len()));
}
let mut order: Vec<usize> = (0..postings.len()).filter(|&i| !postings[i].is_empty()).collect();
order.sort_by_key(|&i| dim_tokens[i]);
let num_dims = order.len() as u32;
let header_size = std::mem::size_of::<FileHeader>();
let dim_headers_size = num_dims as usize * std::mem::size_of::<DimHeader>();
let entries_start = header_size + dim_headers_size;
write_atomic(path, |file| {
let mut out = CrcWriter { inner: file, hasher: crc32fast::Hasher::new() };
let out = &mut out;
let header = FileHeader {
magic: MAGIC,
version: FORMAT_VERSION,
num_dims,
num_vectors,
};
out.write_all(as_bytes(&header))
.map_err(|e| format!("write header: {e}"))?;
let mut current_offset = entries_start;
for &i in &order {
let dh = DimHeader {
offset: current_offset as u64,
count: postings[i].len() as u32,
token_id: dim_tokens[i],
};
out.write_all(as_bytes(&dh))
.map_err(|e| format!("write dim header: {e}"))?;
current_offset += postings[i].len() * std::mem::size_of::<PostingEntry>();
}
for &i in &order {
let p = &postings[i];
for x in p.as_slice() {
let entry = PostingEntry {
record_id: x.id,
weight: x.weight,
max_next_weight: x.tail_max,
};
out.write_all(as_bytes(&entry))
.map_err(|e| format!("write entry: {e}"))?;
}
}
let crc = out.hasher.clone().finalize();
out.inner.write_all(&crc.to_le_bytes()).map_err(|e| format!("write checksum: {e}"))?;
out.inner.write_all(&MAGIC.to_le_bytes()).map_err(|e| format!("write footer magic: {e}"))?;
Ok(())
})
}
struct CrcWriter<'a> {
inner: &'a mut std::io::BufWriter<std::fs::File>,
hasher: crc32fast::Hasher,
}
impl std::io::Write for CrcWriter<'_> {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
let n = self.inner.write(buf)?;
self.hasher.update(&buf[..n]);
Ok(n)
}
fn flush(&mut self) -> std::io::Result<()> {
self.inner.flush()
}
}
fn write_atomic(
path: &Path,
body: impl FnOnce(&mut std::io::BufWriter<std::fs::File>) -> Result<(), String>,
) -> Result<(), String> {
use std::io::Write;
let name = path.file_name().ok_or_else(|| format!("{}: no file name", path.display()))?;
let tmp = path.with_file_name(format!("{}.tmp", name.to_string_lossy()));
let file = std::fs::File::create(&tmp)
.map_err(|e| format!("cannot create {}: {e}", tmp.display()))?;
let mut out = std::io::BufWriter::new(file);
let written = body(&mut out).and_then(|()| {
out.flush().map_err(|e| format!("flush {}: {e}", tmp.display()))?;
out.get_ref().sync_all().map_err(|e| format!("sync {}: {e}", tmp.display()))
});
drop(out);
if let Err(e) = written {
let _ = std::fs::remove_file(&tmp);
return Err(e);
}
std::fs::rename(&tmp, path)
.map_err(|e| format!("cannot rename {} onto {}: {e}", tmp.display(), path.display()))?;
if let Some(dir) = path.parent() {
if let Ok(d) = std::fs::File::open(dir) {
let _ = d.sync_all();
}
}
Ok(())
}
pub fn write_file_atomic(path: &Path, data: &[u8]) -> Result<(), String> {
use std::io::Write;
write_atomic(path, |out| out.write_all(data).map_err(|e| format!("write {}: {e}", path.display())))
}
fn as_bytes<T: Sized>(val: &T) -> &[u8] {
unsafe { std::slice::from_raw_parts(val as *const T as *const u8, std::mem::size_of::<T>()) }
}