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<Vec<Vec<f32>>>,
embedder_override: Option<Arc<dyn Embedder>>,
}
impl DenseCache {
pub(crate) fn new() -> Self {
Self {
vectors: Mutex::new(Vec::new()),
embedder_override: None,
}
}
#[cfg(test)]
pub(crate) fn with_embedder(embedder: Arc<dyn Embedder>) -> Self {
Self {
vectors: Mutex::new(Vec::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<T: Embeddable>(
&self,
items: &[T],
sink: &dyn TraceSink,
) -> Result<(), EmbedderError> {
let mut guard = self.vectors.lock().expect("embeddings mutex poisoned");
if guard.len() >= items.len() {
return Ok(());
}
let embedder = self.resolve_embedder(sink)?;
for item in &items[guard.len()..] {
guard.push(embedder.embed_doc(&item.embed_text())?);
}
Ok(())
}
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<T: Embeddable>(
&self,
items: &[T],
query_vec: &[f32],
depth: usize,
) -> Vec<(String, f32)> {
let guard = self.vectors.lock().expect("embeddings mutex poisoned");
let mut latest: HashMap<&str, &[f32]> = HashMap::new();
for (item, embedding) in items.iter().zip(guard.iter()) {
latest.insert(item.embed_id(), embedding.as_slice());
}
dense_search(
latest.into_iter().map(|(id, v)| (id.to_string(), v)),
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 ranked_dedups_duplicate_ids_last_wins() {
let cache = DenseCache::with_embedder(Arc::new(CountingStub::new()));
let items = vec![doc("x", "read"), doc("x", "write")];
cache.extend(&items, &NoopSink).unwrap();
let ranked = cache.ranked(&items, &[0.0, 1.0], 10); assert_eq!(ranked.len(), 1, "a duplicate id collapses to one entry");
assert_eq!(ranked[0].0, "x");
assert!(ranked[0].1 > 0.9, "ranks with the last-registered vector");
}
}