use std::future::Future;
use std::pin::Pin;
use dashmap::DashMap;
use super::{VectorMatch, VectorMetadata, VectorStore, tenant_matches};
use crate::error::{LiterLlmError, Result};
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm_a == 0.0 || norm_b == 0.0 {
return 0.0;
}
dot / (norm_a * norm_b)
}
struct Entry {
vec: Vec<f32>,
metadata: VectorMetadata,
}
pub struct InMemoryVectorStore {
entries: DashMap<String, Entry>,
dim: usize,
}
impl InMemoryVectorStore {
#[must_use]
pub fn new(dim: usize) -> Self {
Self {
entries: DashMap::new(),
dim,
}
}
}
impl VectorStore for InMemoryVectorStore {
fn search<'a>(
&'a self,
query_vec: &'a [f32],
k: usize,
threshold: f32,
tenant_id: Option<&'a str>,
) -> Pin<Box<dyn Future<Output = Vec<VectorMatch>> + Send + 'a>> {
let mut matches: Vec<VectorMatch> = self
.entries
.iter()
.filter_map(|entry| {
if !tenant_matches(entry.metadata.tenant_id.as_deref(), tenant_id) {
return None;
}
let sim = cosine_similarity(query_vec, &entry.vec);
if sim >= threshold {
Some(VectorMatch {
id: entry.key().clone(),
similarity: sim,
metadata: entry.metadata.clone(),
})
} else {
None
}
})
.collect();
matches.sort_by(|a, b| {
b.similarity
.partial_cmp(&a.similarity)
.unwrap_or(std::cmp::Ordering::Equal)
});
matches.truncate(k);
Box::pin(std::future::ready(matches))
}
fn upsert<'a>(
&'a self,
id: String,
vec: Vec<f32>,
metadata: VectorMetadata,
) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
if vec.len() != self.dim {
return Box::pin(std::future::ready(Err(LiterLlmError::InternalError {
message: format!(
"vector dimension mismatch: store expects {} but received {}",
self.dim,
vec.len()
),
})));
}
self.entries.insert(id, Entry { vec, metadata });
Box::pin(std::future::ready(Ok(())))
}
fn delete<'a>(&'a self, id: &'a str) -> Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>> {
self.entries.remove(id);
Box::pin(std::future::ready(Ok(())))
}
fn dim(&self) -> usize {
self.dim
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::time::SystemTime;
use super::*;
fn meta(cache_key: u64) -> VectorMetadata {
VectorMetadata {
cache_key,
original_request_body: String::new(),
image_url: None,
tenant_id: None,
inserted_at: SystemTime::now(),
extra: HashMap::new(),
}
}
#[test]
fn cosine_similarity_identical_vectors() {
let v = vec![1.0_f32, 0.0, 0.0];
assert!((cosine_similarity(&v, &v) - 1.0).abs() < 1e-6);
}
#[test]
fn cosine_similarity_orthogonal_vectors() {
let a = vec![1.0_f32, 0.0, 0.0];
let b = vec![0.0_f32, 1.0, 0.0];
assert!(cosine_similarity(&a, &b).abs() < 1e-6);
}
#[test]
fn cosine_similarity_zero_vector_returns_zero() {
let a = vec![0.0_f32, 0.0, 0.0];
let b = vec![1.0_f32, 0.0, 0.0];
assert_eq!(cosine_similarity(&a, &b), 0.0);
}
#[test]
fn cosine_similarity_length_mismatch_returns_zero() {
let a = vec![1.0_f32, 0.0];
let b = vec![1.0_f32, 0.0, 0.0];
assert_eq!(cosine_similarity(&a, &b), 0.0);
}
#[tokio::test]
async fn upsert_and_search_returns_match_above_threshold() {
let store = InMemoryVectorStore::new(3);
store.upsert("v1".into(), vec![1.0, 0.0, 0.0], meta(42)).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.99, None).await;
assert_eq!(results.len(), 1, "should find the identical vector");
assert_eq!(results[0].id, "v1");
assert!((results[0].similarity - 1.0).abs() < 1e-5);
assert_eq!(results[0].metadata.cache_key, 42);
}
#[tokio::test]
async fn image_metadata_converts_to_chat_content_part() {
let store = InMemoryVectorStore::new(2);
let mut metadata = meta(7);
metadata.image_url = Some(crate::types::ImageUrl {
url: "https://example.com/result.png".into(),
detail: None,
});
store.upsert("image".into(), vec![1.0, 0.0], metadata).await.unwrap();
let result = store.search(&[1.0, 0.0], 1, 0.99, None).await.remove(0);
assert_eq!(
result.metadata.image_content_part(),
Some(crate::types::ContentPart::image_url("https://example.com/result.png"))
);
}
#[tokio::test]
async fn search_filters_below_threshold() {
let store = InMemoryVectorStore::new(3);
store.upsert("v1".into(), vec![1.0, 0.0, 0.0], meta(1)).await.unwrap();
store.upsert("v2".into(), vec![0.0, 1.0, 0.0], meta(2)).await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.9, None).await;
assert_eq!(results.len(), 1, "orthogonal vector should be filtered out");
assert_eq!(results[0].id, "v1");
}
#[tokio::test]
async fn search_returns_k_nearest() {
let store = InMemoryVectorStore::new(2);
store.upsert("a".into(), vec![1.0, 0.0], meta(1)).await.unwrap();
store.upsert("b".into(), vec![0.9, 0.1], meta(2)).await.unwrap();
store.upsert("c".into(), vec![0.8, 0.2], meta(3)).await.unwrap();
let results = store.search(&[1.0, 0.0], 2, 0.0, None).await;
assert_eq!(results.len(), 2, "should return exactly k results");
assert!(results[0].similarity >= results[1].similarity);
}
#[tokio::test]
async fn search_returns_empty_when_store_is_empty() {
let store = InMemoryVectorStore::new(3);
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.0, None).await;
assert!(results.is_empty());
}
#[tokio::test]
async fn delete_removes_entry() {
let store = InMemoryVectorStore::new(3);
store.upsert("v1".into(), vec![1.0, 0.0, 0.0], meta(1)).await.unwrap();
store.delete("v1").await.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.0, None).await;
assert!(results.is_empty(), "deleted entry must not appear in search results");
}
#[tokio::test]
async fn delete_nonexistent_is_noop() {
let store = InMemoryVectorStore::new(3);
let result = store.delete("does-not-exist").await;
assert!(result.is_ok(), "deleting a non-existent entry should not error");
}
#[tokio::test]
async fn upsert_replaces_existing_entry() {
let store = InMemoryVectorStore::new(3);
store.upsert("v1".into(), vec![1.0, 0.0, 0.0], meta(1)).await.unwrap();
store.upsert("v1".into(), vec![0.0, 1.0, 0.0], meta(99)).await.unwrap();
let results = store.search(&[0.0, 1.0, 0.0], 5, 0.99, None).await;
assert_eq!(results.len(), 1);
assert_eq!(results[0].metadata.cache_key, 99, "upsert should replace metadata");
}
#[tokio::test]
async fn upsert_dimension_mismatch_returns_error() {
let store = InMemoryVectorStore::new(3);
let result = store.upsert("bad".into(), vec![1.0, 0.0], meta(1)).await;
assert!(result.is_err(), "dimension mismatch must return an error");
}
fn meta_for_tenant(cache_key: u64, tenant_id: &str) -> VectorMetadata {
VectorMetadata {
tenant_id: Some(tenant_id.to_owned()),
..meta(cache_key)
}
}
#[tokio::test]
async fn search_excludes_entries_from_other_tenants() {
let store = InMemoryVectorStore::new(3);
store
.upsert("a".into(), vec![1.0, 0.0, 0.0], meta_for_tenant(1, "tenant-a"))
.await
.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.99, Some("tenant-b")).await;
assert!(
results.is_empty(),
"tenant-b must not see tenant-a's semantically-matched entry"
);
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.99, Some("tenant-a")).await;
assert_eq!(results.len(), 1, "tenant-a must still see its own entry");
}
#[tokio::test]
async fn search_tenant_none_is_not_a_wildcard() {
let store = InMemoryVectorStore::new(3);
store
.upsert("scoped".into(), vec![1.0, 0.0, 0.0], meta_for_tenant(1, "tenant-a"))
.await
.unwrap();
store
.upsert("unscoped".into(), vec![1.0, 0.0, 0.0], meta(2))
.await
.unwrap();
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.99, None).await;
assert_eq!(results.len(), 1, "None query must only match the None-tenant entry");
assert_eq!(results[0].id, "unscoped");
let results = store.search(&[1.0, 0.0, 0.0], 5, 0.99, Some("tenant-a")).await;
assert_eq!(
results.len(),
1,
"tenant-scoped query must only match that tenant's entry"
);
assert_eq!(results[0].id, "scoped");
}
#[test]
fn dim_returns_configured_dimension() {
let store = InMemoryVectorStore::new(512);
assert_eq!(store.dim(), 512);
}
}