use anyhow::{anyhow, Result};
use parking_lot::Mutex;
use serde_json::Value as JsonValue;
use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use vecstore::{Metadata, Query, VecStore};
pub use vecstore::Neighbor as SearchResult;
use crate::storage::embedding::EmbeddingEngine;
use crate::storage::fingerprint::json_fingerprint;
const MAX_SYNC_RETRIES: u32 = 5;
fn sync_retry_step(succeeded: bool, failures: u32) -> (u32, Option<std::time::Duration>, bool) {
if succeeded {
return (0, None, false);
}
let f = failures + 1;
if f >= MAX_SYNC_RETRIES {
(f, None, true)
} else {
(f, Some(std::time::Duration::from_millis(100 * f as u64)), false)
}
}
#[derive(Clone)]
pub struct VectorEngine {
path: String,
store: Arc<Mutex<Option<VecStore>>>,
embedding: Option<Arc<EmbeddingEngine>>,
dirty: Arc<AtomicBool>,
sync_in_flight: Arc<AtomicBool>,
}
impl VectorEngine {
pub fn embedding_is_loaded(&self) -> bool {
self.embedding.as_ref().map(|e| e.is_loaded()).unwrap_or(false)
}
pub fn with_embedding(path: &str, engine: EmbeddingEngine) -> Result<Self> {
Ok(Self {
path: path.to_string(),
store: Arc::new(Mutex::new(None)),
embedding: Some(Arc::new(engine)),
dirty: Arc::new(AtomicBool::new(false)),
sync_in_flight: Arc::new(AtomicBool::new(false)),
})
}
pub fn store_document(&self, id: &str, document: JsonValue) -> Result<()> {
let Some(engine) = &self.embedding else {
return Ok(());
};
let fingerprint = json_fingerprint(&document);
let vector = engine.embed(&fingerprint)?;
let meta = json_to_metadata(document);
let dirty = self.dirty.clone();
self.with_store(|s| {
s.upsert(id.to_string(), vector, meta)
.map_err(|e| anyhow!("failed to store document {id:?}: {e}"))?;
dirty.store(true, Ordering::Release);
Ok(())
})
}
pub fn store_documents_batch(&self, entries: &[(&str, JsonValue)]) -> Result<()> {
let Some(engine) = &self.embedding else {
return Ok(());
};
if entries.is_empty() {
return Ok(());
}
let fingerprints: Vec<String> = entries
.iter()
.map(|(_, doc)| json_fingerprint(doc))
.collect();
let fp_refs: Vec<&str> = fingerprints.iter().map(String::as_str).collect();
let vectors = engine.embed_batch(&fp_refs)?;
let dirty = self.dirty.clone();
self.with_store(|s| {
for ((id, doc), vector) in entries.iter().zip(vectors) {
let meta = json_to_metadata(doc.clone());
s.upsert(id.to_string(), vector, meta)
.map_err(|e| anyhow!("failed to store document {id:?}: {e}"))?;
}
dirty.store(true, Ordering::Release);
Ok(())
})
}
pub fn delete_vector(&self, id: &str) -> Result<()> {
let dirty = self.dirty.clone();
self.with_store(|s| {
match s.remove(id) {
Ok(()) => {
dirty.store(true, Ordering::Release);
Ok(())
}
Err(e) if e.to_string().to_lowercase().contains("not found") => Ok(()),
Err(e) => Err(anyhow!("failed to remove vector {id:?}: {e}")),
}
})
}
pub fn search(&self, query_vector: Vec<f32>, limit: usize) -> Result<Vec<SearchResult>> {
let q = Query::new(query_vector).with_limit(limit);
let mut results = self
.with_store(|s| s.query(q).map_err(|e| anyhow!("vector search failed: {e}")))?;
distance_to_similarity(&mut results);
Ok(results)
}
pub fn search_json(&self, query: &JsonValue, limit: usize) -> Result<Vec<SearchResult>> {
let engine = self
.embedding
.clone()
.ok_or_else(|| anyhow!("search_json requires an EmbeddingEngine"))?;
let fingerprint = json_fingerprint(query);
let vector = engine.embed(&fingerprint)?;
self.search(vector, limit)
}
pub fn count(&self) -> Result<usize> {
self.with_store(|s| Ok(s.count()))
}
pub fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
let Some(engine) = &self.embedding else {
return Err(anyhow!("vector engine has no embedding model configured"));
};
engine.embed_batch(texts)
}
pub fn sync(&self) -> Result<()> {
if !self.dirty.load(Ordering::Acquire) {
return Ok(());
}
Self::drain_dirty(&self.store, &self.dirty)
}
pub fn sync_in_background(&self) {
if !self.dirty.load(Ordering::Acquire) {
return;
}
if self.sync_in_flight.swap(true, Ordering::AcqRel) {
return;
}
let store = self.store.clone();
let dirty = self.dirty.clone();
let in_flight = self.sync_in_flight.clone();
std::thread::spawn(move || {
let mut failures: u32 = 0;
loop {
let (next_failures, backoff, give_up) =
match Self::drain_dirty(&store, &dirty) {
Ok(()) => sync_retry_step(true, failures),
Err(e) => {
tracing::warn!(
target: "inkhaven::storage::vector",
"background vector sync failed (attempt {}): {e}",
failures + 1,
);
sync_retry_step(false, failures)
}
};
failures = next_failures;
if let Some(delay) = backoff {
std::thread::sleep(delay);
}
in_flight.store(false, Ordering::Release);
if give_up {
break;
}
if !dirty.load(Ordering::Acquire) {
break;
}
if in_flight.swap(true, Ordering::AcqRel) {
break;
}
}
});
}
fn drain_dirty(store: &Mutex<Option<VecStore>>, dirty: &AtomicBool) -> Result<()> {
let mut guard = store.lock();
if !dirty.load(Ordering::Acquire) {
return Ok(());
}
let Some(s) = guard.as_mut() else {
dirty.store(false, Ordering::Release);
return Ok(());
};
match s.save() {
Ok(()) => {
dirty.store(false, Ordering::Release);
Ok(())
}
Err(e) => Err(anyhow!("failed to sync vector store: {e}")),
}
}
fn with_store<R, F: FnOnce(&mut VecStore) -> Result<R>>(&self, f: F) -> Result<R> {
let mut guard = self.store.lock();
if guard.is_none() {
*guard = Some(
VecStore::open(&self.path)
.map_err(|e| anyhow!("failed to open vector store at {:?}: {e}", self.path))?,
);
}
let store = guard.as_mut().expect("set immediately above when None");
f(store)
}
}
fn distance_to_similarity(results: &mut [SearchResult]) {
for r in results.iter_mut() {
r.score = 1.0 - r.score;
}
}
fn json_to_metadata(json: JsonValue) -> Metadata {
let fields = match json {
JsonValue::Object(map) => map.into_iter().collect(),
other => {
let mut m = HashMap::new();
m.insert("value".to_string(), other);
m
}
};
Metadata { fields }
}
#[cfg(test)]
mod tests_sync_retry {
use super::{sync_retry_step, MAX_SYNC_RETRIES};
#[test]
fn success_resets_and_never_backs_off() {
let (failures, backoff, give_up) = sync_retry_step(true, 4);
assert_eq!(failures, 0);
assert!(backoff.is_none());
assert!(!give_up);
}
#[test]
fn failures_back_off_linearly_then_give_up() {
let mut failures = 0u32;
let mut gave_up = false;
for step in 1..=MAX_SYNC_RETRIES {
let (f, backoff, give_up) = sync_retry_step(false, failures);
failures = f;
assert_eq!(f, step);
if step < MAX_SYNC_RETRIES {
assert_eq!(backoff.unwrap().as_millis() as u64, 100 * step as u64);
assert!(!give_up);
} else {
assert!(backoff.is_none());
assert!(give_up);
gave_up = true;
}
}
assert!(gave_up, "must give up at the retry ceiling");
}
}