use std::collections::BTreeMap;
use std::path::{Path, PathBuf};
use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use trusty_common::bm25::BM25Index;
use crate::bm25_lane::BM25Hit;
pub const SNAPSHOT_FILENAME: &str = "bm25_index.json";
#[derive(Debug, Clone, Serialize, Deserialize)]
struct Document {
doc_id: String,
text: String,
}
pub struct PalaceBm25Index {
inner: BM25Index,
snapshot_path: PathBuf,
docs: BTreeMap<String, String>,
dirty: bool,
}
impl PalaceBm25Index {
pub fn load_or_create(data_dir: &Path) -> Result<Self> {
std::fs::create_dir_all(data_dir)
.with_context(|| format!("create palace bm25 dir {}", data_dir.display()))?;
let snapshot_path = data_dir.join(SNAPSHOT_FILENAME);
let mut inner = BM25Index::new();
let mut docs = BTreeMap::new();
match std::fs::read(&snapshot_path) {
Ok(bytes) => match serde_json::from_slice::<Vec<Document>>(&bytes) {
Ok(rows) => {
for row in rows {
inner.upsert_document(&row.doc_id, &row.text);
docs.insert(row.doc_id, row.text);
}
tracing::info!(
path = %snapshot_path.display(),
doc_count = docs.len(),
"loaded BM25 snapshot"
);
}
Err(e) => {
tracing::warn!(
path = %snapshot_path.display(),
"corrupt BM25 snapshot ({e}); starting with empty index"
);
}
},
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
tracing::debug!(
path = %snapshot_path.display(),
"no BM25 snapshot found — starting with empty index"
);
}
Err(e) => {
return Err(anyhow::Error::new(e)
.context(format!("read BM25 snapshot at {}", snapshot_path.display())));
}
}
Ok(Self {
inner,
snapshot_path,
docs,
dirty: false,
})
}
pub fn index_doc(&mut self, doc_id: &str, text: &str) {
self.inner.upsert_document(doc_id, text);
self.docs.insert(doc_id.to_string(), text.to_string());
self.dirty = true;
}
pub fn search(&self, query: &str, top_k: usize) -> Vec<BM25Hit> {
let top_k = top_k.max(1);
self.inner
.score_query_all(query, top_k)
.into_iter()
.map(|(doc_id, score)| BM25Hit { doc_id, score })
.collect()
}
pub fn delete_doc(&mut self, doc_id: &str) -> bool {
let was_present = self.docs.remove(doc_id).is_some();
if was_present {
self.inner.remove_document(doc_id);
self.dirty = true;
}
was_present
}
pub fn doc_count(&self) -> usize {
self.inner.len()
}
pub fn total_text_bytes(&self) -> u64 {
self.docs.values().map(|t| t.len() as u64).sum()
}
pub fn missing_docs(&self, doc_ids: &[String]) -> Vec<String> {
doc_ids
.iter()
.filter(|id| !self.docs.contains_key(*id))
.cloned()
.collect()
}
pub fn snapshot_path(&self) -> &Path {
&self.snapshot_path
}
pub fn is_dirty(&self) -> bool {
self.dirty
}
pub fn flush(&mut self) -> Result<()> {
if !self.dirty {
return Ok(());
}
let rows: Vec<Document> = self
.docs
.iter()
.map(|(doc_id, text)| Document {
doc_id: doc_id.clone(),
text: text.clone(),
})
.collect();
let json = serde_json::to_vec(&rows).context("serialise BM25 snapshot")?;
let tmp_path = self.snapshot_path.with_extension("json.tmp");
std::fs::write(&tmp_path, &json)
.with_context(|| format!("write BM25 snapshot tmp file {}", tmp_path.display()))?;
std::fs::rename(&tmp_path, &self.snapshot_path).with_context(|| {
format!(
"atomic rename {} → {}",
tmp_path.display(),
self.snapshot_path.display()
)
})?;
self.dirty = false;
tracing::debug!(
path = %self.snapshot_path.display(),
doc_count = rows.len(),
"flushed BM25 snapshot"
);
Ok(())
}
}
#[cfg(test)]
#[path = "bm25_index_tests.rs"]
mod tests;