use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use crate::dense_search::dense_search;
use crate::embedding::{Embedder, EmbedderError, embedder_with_telemetry};
use crate::trace::TraceSink;
pub(crate) trait Embeddable {
fn embed_id(&self) -> &str;
fn embed_text(&self) -> String;
}
pub(crate) struct DenseCache {
vectors: Mutex<HashMap<String, Vec<f32>>>,
embedder_override: Option<Arc<dyn Embedder>>,
}
impl DenseCache {
pub(crate) fn new() -> Self {
Self {
vectors: Mutex::new(HashMap::new()),
embedder_override: None,
}
}
#[cfg(test)]
pub(crate) fn with_embedder(embedder: Arc<dyn Embedder>) -> Self {
Self {
vectors: Mutex::new(HashMap::new()),
embedder_override: Some(embedder),
}
}
fn resolve_embedder(&self, sink: &dyn TraceSink) -> Result<Arc<dyn Embedder>, EmbedderError> {
match &self.embedder_override {
Some(e) => Ok(e.clone()),
None => embedder_with_telemetry(sink),
}
}
pub(crate) fn require_built(&self, corpus_len: usize) -> Result<(), EmbedderError> {
let cached = self
.vectors
.lock()
.expect("embeddings mutex poisoned")
.len();
if cached < corpus_len {
return Err(EmbedderError::EmbeddingsNotBuilt);
}
Ok(())
}
pub(crate) fn extend<'a, T: Embeddable + 'a>(
&self,
items: impl IntoIterator<Item = &'a T>,
sink: &dyn TraceSink,
) -> Result<(), EmbedderError> {
let mut guard = self.vectors.lock().expect("embeddings mutex poisoned");
let mut embedder: Option<Arc<dyn Embedder>> = None;
for item in items {
if guard.contains_key(item.embed_id()) {
continue;
}
if embedder.is_none() {
embedder = Some(self.resolve_embedder(sink)?);
}
let vector = embedder
.as_ref()
.expect("embedder resolved on first miss")
.embed_doc(&item.embed_text())?;
guard.insert(item.embed_id().to_string(), vector);
}
Ok(())
}
pub(crate) fn invalidate(&self, id: &str) {
self.vectors
.lock()
.expect("embeddings mutex poisoned")
.remove(id);
}
pub(crate) fn embed_query(
&self,
query: &str,
sink: &dyn TraceSink,
) -> Result<Vec<f32>, EmbedderError> {
self.resolve_embedder(sink)?.embed_query(query)
}
pub(crate) fn ranked<'a, T: Embeddable + 'a>(
&self,
items: impl IntoIterator<Item = &'a T>,
query_vec: &[f32],
depth: usize,
) -> Vec<(String, f32)> {
let guard = self.vectors.lock().expect("embeddings mutex poisoned");
let docs: Vec<(String, &[f32])> = items
.into_iter()
.filter_map(|item| {
guard
.get(item.embed_id())
.map(|v| (item.embed_id().to_string(), v.as_slice()))
})
.collect();
dense_search(docs, query_vec, depth)
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use super::*;
use crate::trace::NoopSink;
struct Doc {
id: String,
text: String,
}
impl Embeddable for Doc {
fn embed_id(&self) -> &str {
&self.id
}
fn embed_text(&self) -> String {
self.text.clone()
}
}
fn doc(id: &str, text: &str) -> Doc {
Doc {
id: id.into(),
text: text.into(),
}
}
struct CountingStub {
docs: AtomicUsize,
}
impl CountingStub {
fn new() -> Self {
Self {
docs: AtomicUsize::new(0),
}
}
fn docs(&self) -> usize {
self.docs.load(Ordering::SeqCst)
}
}
fn vec_for(text: &str) -> Vec<f32> {
if text.to_lowercase().contains("read") {
vec![1.0, 0.0]
} else {
vec![0.0, 1.0]
}
}
impl Embedder for CountingStub {
fn embed_doc(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
self.docs.fetch_add(1, Ordering::SeqCst);
Ok(vec_for(text))
}
fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbedderError> {
Ok(vec_for(text))
}
}
#[test]
fn require_built_errors_until_the_cache_covers_the_corpus() {
let cache = DenseCache::with_embedder(Arc::new(CountingStub::new()));
let items = vec![doc("a", "read"), doc("b", "write")];
assert!(matches!(
cache.require_built(items.len()),
Err(EmbedderError::EmbeddingsNotBuilt)
));
cache.extend(&items, &NoopSink).unwrap();
assert!(cache.require_built(items.len()).is_ok());
}
#[test]
fn extend_embeds_only_the_new_tail() {
let stub = Arc::new(CountingStub::new());
let cache = DenseCache::with_embedder(stub.clone());
let mut items = vec![doc("a", "read"), doc("b", "write")];
cache.extend(&items, &NoopSink).unwrap();
assert_eq!(stub.docs(), 2);
items.push(doc("c", "read"));
cache.extend(&items, &NoopSink).unwrap();
assert_eq!(stub.docs(), 3, "only the newly-appended item is embedded");
cache.extend(&items, &NoopSink).unwrap();
assert_eq!(stub.docs(), 3);
}
#[test]
fn invalidate_forces_re_embed_of_an_id() {
let stub = Arc::new(CountingStub::new());
let cache = DenseCache::with_embedder(stub.clone());
cache.extend([&doc("x", "read")], &NoopSink).unwrap();
assert_eq!(stub.docs(), 1);
cache.invalidate("x");
cache.extend([&doc("x", "write")], &NoopSink).unwrap();
assert_eq!(stub.docs(), 2, "invalidated id is re-embedded, once");
let item = doc("x", "write");
let ranked = cache.ranked([&item], &[0.0, 1.0], 10);
assert_eq!(ranked.len(), 1);
assert_eq!(ranked[0].0, "x");
assert!(ranked[0].1 > 0.9, "ranks with the re-embedded vector");
}
#[test]
fn require_built_fails_after_invalidate_until_rebuilt() {
let cache = DenseCache::with_embedder(Arc::new(CountingStub::new()));
let items = vec![doc("a", "read"), doc("b", "write")];
cache.extend(&items, &NoopSink).unwrap();
assert!(cache.require_built(items.len()).is_ok());
cache.invalidate("a");
assert!(
matches!(
cache.require_built(items.len()),
Err(EmbedderError::EmbeddingsNotBuilt)
),
"an invalidated id drops the cache below the corpus until rebuilt"
);
cache.extend(&items, &NoopSink).unwrap();
assert!(cache.require_built(items.len()).is_ok());
}
}