use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use anyhow::Result;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
#[derive(Clone)]
pub struct VectorStoreExt(pub Arc<dyn VectorStore>);
const RRF_K: f32 = 60.0;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct KnowledgeIndexCitation {
pub id: String,
pub index_id: String,
pub document_title: Option<String>,
pub source_uri: String,
pub location: Option<serde_json::Value>,
pub snippet: String,
pub score: f32,
}
#[derive(Debug, Clone, PartialEq)]
pub struct EmbeddingCallUsage {
pub model: String,
pub provider: Option<String>,
pub total_tokens: u32,
pub actual_cost_usd: Option<f64>,
}
#[derive(Debug, Clone, Default)]
pub struct KnowledgeIndexSearchOutcome {
pub citations: Vec<KnowledgeIndexCitation>,
pub embedding_usage: Vec<EmbeddingCallUsage>,
}
#[async_trait]
pub trait KnowledgeIndexSearch: Send + Sync {
async fn search(
&self,
org_id: i64,
index_ids: &[String],
query: &str,
top_k: usize,
) -> Result<KnowledgeIndexSearchOutcome>;
}
pub fn index_namespace(org_id: i64, public_id: &str) -> String {
format!("org_{org_id}__{public_id}")
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VectorRecord {
pub id: String,
pub vector: Vec<f32>,
pub text: String,
pub document_id: String,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct VectorQuery {
pub vector: Option<Vec<f32>>,
pub text: Option<String>,
pub top_k: usize,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VectorMatch {
pub id: String,
pub document_id: String,
pub score: f32,
}
#[async_trait]
pub trait VectorStore: Send + Sync {
async fn upsert(&self, namespace: &str, records: Vec<VectorRecord>) -> Result<()>;
async fn query(&self, namespace: &str, query: VectorQuery) -> Result<Vec<VectorMatch>>;
async fn delete_by_document(&self, namespace: &str, document_id: &str) -> Result<()>;
async fn delete_namespace(&self, namespace: &str) -> Result<()>;
}
#[derive(Default)]
pub struct InMemoryVectorStore {
namespaces: Mutex<HashMap<String, Vec<VectorRecord>>>,
}
impl InMemoryVectorStore {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait]
impl VectorStore for InMemoryVectorStore {
async fn upsert(&self, namespace: &str, records: Vec<VectorRecord>) -> Result<()> {
let mut store = self
.namespaces
.lock()
.map_err(|_| anyhow::anyhow!("vector store poisoned"))?;
let entry = store.entry(namespace.to_string()).or_default();
for record in records {
if let Some(existing) = entry.iter_mut().find(|r| r.id == record.id) {
*existing = record;
} else {
entry.push(record);
}
}
Ok(())
}
async fn query(&self, namespace: &str, query: VectorQuery) -> Result<Vec<VectorMatch>> {
if query.top_k == 0 {
return Ok(Vec::new());
}
let store = self
.namespaces
.lock()
.map_err(|_| anyhow::anyhow!("vector store poisoned"))?;
let Some(records) = store.get(namespace) else {
return Ok(Vec::new());
};
let vector_ranked = query
.vector
.as_ref()
.map(|v| rank_by(records, |r| cosine_similarity(&r.vector, v)));
let text_ranked = query.text.as_ref().and_then(|t| {
let t = t.trim();
(!t.is_empty()).then(|| rank_by(records, |r| term_overlap_score(&r.text, t)))
});
let ordered_ids = match (vector_ranked, text_ranked) {
(Some(v), Some(t)) => fuse_rrf(&v, &t),
(Some(v), None) => v,
(None, Some(t)) => t,
(None, None) => return Ok(Vec::new()),
};
let by_id: HashMap<&str, &VectorRecord> =
records.iter().map(|r| (r.id.as_str(), r)).collect();
let matches = ordered_ids
.into_iter()
.take(query.top_k)
.filter_map(|(id, score)| {
by_id.get(id.as_str()).map(|r| VectorMatch {
id: r.id.clone(),
document_id: r.document_id.clone(),
score,
})
})
.collect();
Ok(matches)
}
async fn delete_by_document(&self, namespace: &str, document_id: &str) -> Result<()> {
let mut store = self
.namespaces
.lock()
.map_err(|_| anyhow::anyhow!("vector store poisoned"))?;
if let Some(entry) = store.get_mut(namespace) {
entry.retain(|r| r.document_id != document_id);
}
Ok(())
}
async fn delete_namespace(&self, namespace: &str) -> Result<()> {
let mut store = self
.namespaces
.lock()
.map_err(|_| anyhow::anyhow!("vector store poisoned"))?;
store.remove(namespace);
Ok(())
}
}
fn rank_by(records: &[VectorRecord], score: impl Fn(&VectorRecord) -> f32) -> Vec<(String, f32)> {
let mut scored: Vec<(String, f32)> = records
.iter()
.map(|r| (r.id.clone(), score(r)))
.filter(|(_, s)| *s > f32::NEG_INFINITY)
.collect();
scored.sort_by(|a, b| b.1.total_cmp(&a.1));
scored
}
fn fuse_rrf(a: &[(String, f32)], b: &[(String, f32)]) -> Vec<(String, f32)> {
let mut fused: HashMap<&str, f32> = HashMap::new();
for list in [a, b] {
for (rank, (id, _)) in list.iter().enumerate() {
*fused.entry(id.as_str()).or_insert(0.0) += 1.0 / (RRF_K + rank as f32 + 1.0);
}
}
let mut ranked: Vec<(String, f32)> = fused
.into_iter()
.map(|(id, s)| (id.to_string(), s))
.collect();
ranked.sort_by(|x, y| y.1.total_cmp(&x.1));
ranked
}
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return f32::NEG_INFINITY;
}
let mut dot = 0.0;
let mut norm_a = 0.0;
let mut norm_b = 0.0;
for (x, y) in a.iter().zip(b.iter()) {
dot += x * y;
norm_a += x * x;
norm_b += y * y;
}
if norm_a == 0.0 || norm_b == 0.0 {
return f32::NEG_INFINITY;
}
dot / (norm_a.sqrt() * norm_b.sqrt())
}
fn term_overlap_score(text: &str, query: &str) -> f32 {
let haystack = text.to_lowercase();
let terms: Vec<String> = query
.to_lowercase()
.split_whitespace()
.map(str::to_string)
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect();
if terms.is_empty() {
return f32::NEG_INFINITY;
}
let hits = terms.iter().filter(|t| haystack.contains(*t)).count();
if hits == 0 {
f32::NEG_INFINITY
} else {
hits as f32 / terms.len() as f32
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn poisoned_store_returns_errors_without_panicking() {
let store = InMemoryVectorStore::new();
let poisoned = std::panic::catch_unwind(|| {
let _guard = store.namespaces.lock().unwrap();
panic!("poison reference store");
});
assert!(poisoned.is_err());
let errors = [
store.upsert("org_1__index", Vec::new()).await.unwrap_err(),
store
.query(
"org_1__index",
VectorQuery {
top_k: 1,
..Default::default()
},
)
.await
.unwrap_err(),
store
.delete_by_document("org_1__index", "document")
.await
.unwrap_err(),
store.delete_namespace("org_1__index").await.unwrap_err(),
];
for error in errors {
assert_eq!(error.to_string(), "vector store poisoned");
}
}
fn record(id: &str, doc: &str, vector: Vec<f32>, text: &str) -> VectorRecord {
VectorRecord {
id: id.to_string(),
vector,
text: text.to_string(),
document_id: doc.to_string(),
}
}
#[test]
fn namespace_is_org_prefixed() {
assert_eq!(
index_namespace(1, "kidx_00000000000000000000000000000001"),
"org_1__kidx_00000000000000000000000000000001"
);
}
#[tokio::test]
async fn vector_query_ranks_by_cosine() {
let store = InMemoryVectorStore::new();
let ns = index_namespace(1, "kidx_00000000000000000000000000000001");
store
.upsert(
&ns,
vec![
record("kchk_a", "kidoc_1", vec![1.0, 0.0], "alpha"),
record("kchk_b", "kidoc_1", vec![0.0, 1.0], "beta"),
record("kchk_c", "kidoc_2", vec![0.9, 0.1], "gamma"),
],
)
.await
.unwrap();
let matches = store
.query(
&ns,
VectorQuery {
vector: Some(vec![1.0, 0.0]),
text: None,
top_k: 2,
},
)
.await
.unwrap();
let ids: Vec<_> = matches.iter().map(|m| m.id.as_str()).collect();
assert_eq!(ids, vec!["kchk_a", "kchk_c"]);
assert_eq!(matches[0].document_id, "kidoc_1");
}
#[tokio::test]
async fn upsert_replaces_existing_id() {
let store = InMemoryVectorStore::new();
let ns = "org_1__kidx_x";
store
.upsert(ns, vec![record("kchk_a", "kidoc_1", vec![1.0, 0.0], "old")])
.await
.unwrap();
store
.upsert(ns, vec![record("kchk_a", "kidoc_1", vec![0.0, 1.0], "new")])
.await
.unwrap();
let matches = store
.query(
ns,
VectorQuery {
vector: Some(vec![0.0, 1.0]),
text: None,
top_k: 5,
},
)
.await
.unwrap();
assert_eq!(matches.len(), 1);
assert!(matches[0].score > 0.99);
}
#[tokio::test]
async fn text_query_ranks_by_term_overlap() {
let store = InMemoryVectorStore::new();
let ns = "org_1__kidx_x";
store
.upsert(
ns,
vec![
record("kchk_a", "kidoc_1", vec![0.0], "the quick brown fox"),
record("kchk_b", "kidoc_1", vec![0.0], "a slow green turtle"),
],
)
.await
.unwrap();
let matches = store
.query(
ns,
VectorQuery {
vector: None,
text: Some("quick fox".to_string()),
top_k: 5,
},
)
.await
.unwrap();
assert_eq!(matches.len(), 1);
assert_eq!(matches[0].id, "kchk_a");
}
#[tokio::test]
async fn delete_by_document_and_namespace() {
let store = InMemoryVectorStore::new();
let ns = "org_1__kidx_x";
store
.upsert(
ns,
vec![
record("kchk_a", "kidoc_1", vec![1.0], "a"),
record("kchk_b", "kidoc_2", vec![1.0], "b"),
],
)
.await
.unwrap();
store.delete_by_document(ns, "kidoc_1").await.unwrap();
let after = store
.query(
ns,
VectorQuery {
vector: Some(vec![1.0]),
text: None,
top_k: 5,
},
)
.await
.unwrap();
assert_eq!(after.len(), 1);
assert_eq!(after[0].id, "kchk_b");
store.delete_namespace(ns).await.unwrap();
let empty = store
.query(
ns,
VectorQuery {
vector: Some(vec![1.0]),
text: None,
top_k: 5,
},
)
.await
.unwrap();
assert!(empty.is_empty());
}
#[tokio::test]
async fn hybrid_query_fuses_both_signals() {
let store = InMemoryVectorStore::new();
let ns = "org_1__kidx_x";
store
.upsert(
ns,
vec![
record(
"kchk_a",
"kidoc_1",
vec![1.0, 0.0],
"database indexing guide",
),
record("kchk_b", "kidoc_1", vec![0.0, 1.0], "cooking recipes"),
],
)
.await
.unwrap();
let matches = store
.query(
ns,
VectorQuery {
vector: Some(vec![1.0, 0.0]),
text: Some("database".to_string()),
top_k: 2,
},
)
.await
.unwrap();
assert_eq!(matches[0].id, "kchk_a");
}
}