use super::passages::embed_passages_parallel_shared;
use super::sizing::entity_embed_batch_size;
use crate::embedder::{is_openrouter_initialized, LlmBackendKind};
use crate::errors::AppError;
use std::path::Path;
use std::sync::Arc;
use std::sync::OnceLock;
struct CacheEntry {
vector: Arc<Vec<f32>>,
stored_at: std::time::Instant,
}
#[derive(Default)]
pub(crate) struct EntityEmbedCacheMap {
entries: std::collections::HashMap<u64, CacheEntry>,
}
impl EntityEmbedCacheMap {
pub(crate) fn insert(&mut self, key: u64, vector: Arc<Vec<f32>>) {
self.entries.insert(
key,
CacheEntry {
vector,
stored_at: std::time::Instant::now(),
},
);
}
pub(crate) fn get(&self, key: &u64) -> Option<&Arc<Vec<f32>>> {
let ttl = std::time::Duration::from_secs(crate::constants::entity_embed_cache_ttl_secs());
let now = std::time::Instant::now();
self.entries
.get(key)
.filter(|entry| now.duration_since(entry.stored_at) < ttl)
.map(|entry| &entry.vector)
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.entries.len()
}
pub(crate) fn evict_expired_and_overflow(&mut self, incoming: usize) {
let ttl = std::time::Duration::from_secs(crate::constants::entity_embed_cache_ttl_secs());
let now = std::time::Instant::now();
self.entries
.retain(|_, entry| now.duration_since(entry.stored_at) < ttl);
let ceiling = crate::constants::entity_embed_cache_max_entries();
let target = ceiling.saturating_sub(incoming.min(ceiling));
if self.entries.len() <= target {
return;
}
let mut by_age: Vec<(u64, std::time::Instant)> = self
.entries
.iter()
.map(|(key, entry)| (*key, entry.stored_at))
.collect();
by_age.sort_by_key(|(_, stored_at)| *stored_at);
for (key, _) in by_age.into_iter().take(self.entries.len() - target) {
self.entries.remove(&key);
}
}
}
static ENTITY_EMBED_CACHE: OnceLock<parking_lot::Mutex<EntityEmbedCacheMap>> = OnceLock::new();
pub(crate) fn entity_embed_cache() -> &'static parking_lot::Mutex<EntityEmbedCacheMap> {
ENTITY_EMBED_CACHE.get_or_init(|| parking_lot::Mutex::new(EntityEmbedCacheMap::default()))
}
pub(crate) fn entity_cache_key(model: &str, text: &str) -> u64 {
let mut hasher = blake3::Hasher::new();
hasher.update(model.as_bytes());
hasher.update(b"\0");
hasher.update(text.as_bytes());
let h = hasher.finalize();
let bytes = h.as_bytes();
u64::from_le_bytes([
bytes[0], bytes[1], bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7],
])
}
pub fn embed_entity_texts_cached(
models_dir: &Path,
texts: &[String],
parallelism: usize,
backends: crate::cli::BackendChoice,
) -> Result<(Vec<Vec<f32>>, EmbedCacheStats), AppError> {
let crate::cli::BackendChoice {
llm: llm_backend,
embedding: embedding_backend,
} = backends;
if texts.is_empty() {
return Ok((Vec::new(), EmbedCacheStats::default()));
}
let chain = embedding_backend.to_chain(llm_backend);
if chain.as_slice() == [LlmBackendKind::None] {
let out: Vec<Vec<f32>> = texts.iter().map(|_| Vec::new()).collect();
return Ok((
out,
EmbedCacheStats {
requested: texts.len(),
hits: 0,
misses: texts.len(),
},
));
}
let routed_openrouter =
chain.first() == Some(&LlmBackendKind::OpenRouter) && is_openrouter_initialized();
let model = if routed_openrouter {
format!("openrouter:{}", crate::constants::embedding_dim())
} else {
format!("none:{}", crate::constants::embedding_dim())
};
let cache = entity_embed_cache();
let mut hits: Vec<Option<Arc<Vec<f32>>>> = vec![None; texts.len()];
let mut miss_indices: Vec<usize> = Vec::with_capacity(texts.len());
{
let guard = cache.lock();
for (i, text) in texts.iter().enumerate() {
let key = entity_cache_key(&model, text);
match guard.get(&key) {
Some(vector) => hits[i] = Some(Arc::clone(vector)),
None => miss_indices.push(i),
}
}
}
let miss_count = miss_indices.len();
if miss_count > 0 {
let miss_texts: Vec<String> = miss_indices.iter().map(|&i| texts[i].clone()).collect();
let mut miss_vecs = embed_passages_parallel_shared(
models_dir,
Arc::from(miss_texts),
parallelism,
entity_embed_batch_size(),
backends,
)?;
let mut guard = cache.lock();
guard.evict_expired_and_overflow(miss_count);
for (slot, &orig_idx) in miss_indices.iter().enumerate() {
let vector = Arc::new(std::mem::take(&mut miss_vecs[slot]));
let key = entity_cache_key(&model, &texts[orig_idx]);
guard.insert(key, Arc::clone(&vector));
hits[orig_idx] = Some(vector);
}
}
let mut out = Vec::with_capacity(texts.len());
for hit in hits.into_iter() {
let v = hit.ok_or_else(|| {
AppError::Embedding(crate::i18n::validation::embedding_entity_cache_null())
})?;
out.push((*v).clone());
}
Ok((
out,
EmbedCacheStats {
requested: texts.len(),
hits: texts.len() - miss_count,
misses: miss_count,
},
))
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, serde::Serialize)]
pub struct EmbedCacheStats {
pub requested: usize,
pub hits: usize,
pub misses: usize,
}
impl EmbedCacheStats {
pub fn hit_rate(&self) -> f64 {
if self.requested == 0 {
0.0
} else {
self.hits as f64 / self.requested as f64
}
}
}