use std::sync::Arc;
use async_trait::async_trait;
use khive_score::DeterministicScore;
use khive_storage::types::{TextQueryMode, TextSearchRequest, VectorSearchRequest};
use khive_storage::{TextSearch, VectorStore};
use uuid::Uuid;
use crate::error::{Result, RetrievalError};
use crate::hybrid::{KeywordSearch, VectorSearch};
fn storage_err_to_retrieval(
err: khive_storage::StorageError,
context: &'static str,
) -> RetrievalError {
use khive_storage::StorageError;
match &err {
StorageError::Timeout { .. } => {
RetrievalError::Hnsw(format!("{context}: {err}"))
}
StorageError::InvalidInput { message, .. } => {
RetrievalError::InvalidQuery(format!("{context}: {message}"))
}
_ => {
RetrievalError::Hnsw(format!("{context}: {err}"))
}
}
}
pub struct StorageVectorSearch {
store: Arc<dyn VectorStore>,
}
impl StorageVectorSearch {
pub fn new(store: Arc<dyn VectorStore>) -> Self {
Self { store }
}
}
#[async_trait]
impl VectorSearch for StorageVectorSearch {
type Id = Uuid;
async fn vector_search(
&self,
embedding: &[f32],
top_k: usize,
) -> Result<Vec<(Uuid, DeterministicScore)>> {
let request = VectorSearchRequest {
query_vectors: vec![embedding.to_vec()],
top_k: top_k as u32,
namespace: None,
kind: None,
embedding_model: None,
filter: None,
backend_hints: None,
};
let hits = self
.store
.search(request)
.await
.map_err(|e| storage_err_to_retrieval(e, "vector search"))?;
Ok(hits
.into_iter()
.map(|hit| (hit.subject_id, hit.score))
.collect())
}
}
pub struct StorageKeywordSearch {
search: Arc<dyn TextSearch>,
}
impl StorageKeywordSearch {
pub fn new(search: Arc<dyn TextSearch>) -> Self {
Self { search }
}
}
#[async_trait]
impl KeywordSearch for StorageKeywordSearch {
type Id = Uuid;
async fn keyword_search(
&self,
text: &str,
top_k: usize,
) -> Result<Vec<(Uuid, DeterministicScore)>> {
let request = TextSearchRequest {
query: text.to_string(),
mode: TextQueryMode::Plain,
filter: None,
top_k: top_k as u32,
snippet_chars: 0, };
let hits = self
.search
.search(request)
.await
.map_err(|e| storage_err_to_retrieval(e, "keyword search"))?;
Ok(hits
.into_iter()
.map(|hit| (hit.subject_id, hit.score))
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
use khive_db::StorageBackend;
use khive_storage::types::TextDocument;
use khive_types::SubstrateKind;
fn test_backend() -> StorageBackend {
StorageBackend::memory().expect("memory backend")
}
#[tokio::test]
async fn vector_search_basic_roundtrip() {
let backend = test_backend();
let store = backend.vectors("test_vs", 3).unwrap();
let id1 = Uuid::new_v4();
let id2 = Uuid::new_v4();
store
.insert(
id1,
SubstrateKind::Entity,
"local",
"content",
vec![vec![1.0, 0.0, 0.0]],
)
.await
.unwrap();
store
.insert(
id2,
SubstrateKind::Entity,
"local",
"content",
vec![vec![0.0, 1.0, 0.0]],
)
.await
.unwrap();
let adapter = StorageVectorSearch::new(store);
let hits = adapter.vector_search(&[1.0, 0.0, 0.0], 2).await.unwrap();
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].0, id1);
assert!(hits[0].1.to_f64() > 0.9);
}
#[tokio::test]
async fn vector_search_respects_top_k() {
let backend = test_backend();
let store = backend.vectors("test_topk", 3).unwrap();
for _ in 0..5 {
store
.insert(
Uuid::new_v4(),
SubstrateKind::Entity,
"local",
"content",
vec![vec![1.0, 0.0, 0.0]],
)
.await
.unwrap();
}
let adapter = StorageVectorSearch::new(store);
let hits = adapter.vector_search(&[1.0, 0.0, 0.0], 3).await.unwrap();
assert_eq!(hits.len(), 3);
}
#[tokio::test]
async fn vector_search_empty_store() {
let backend = test_backend();
let store = backend.vectors("test_empty", 3).unwrap();
let adapter = StorageVectorSearch::new(store);
let hits = adapter.vector_search(&[1.0, 0.0, 0.0], 5).await.unwrap();
assert!(hits.is_empty());
}
#[tokio::test]
async fn vector_search_returns_deterministic_scores() {
let backend = test_backend();
let store = backend.vectors("test_det", 3).unwrap();
let id = Uuid::new_v4();
store
.insert(
id,
SubstrateKind::Entity,
"local",
"content",
vec![vec![1.0, 0.0, 0.0]],
)
.await
.unwrap();
let adapter = StorageVectorSearch::new(store);
let hits1 = adapter.vector_search(&[1.0, 0.0, 0.0], 1).await.unwrap();
let hits2 = adapter.vector_search(&[1.0, 0.0, 0.0], 1).await.unwrap();
assert_eq!(hits1[0].1, hits2[0].1);
}
#[tokio::test]
async fn keyword_search_basic_roundtrip() {
let backend = test_backend();
let store = backend.text("test_ks").unwrap();
let id1 = Uuid::new_v4();
let id2 = Uuid::new_v4();
store
.upsert_document(TextDocument {
subject_id: id1,
kind: SubstrateKind::Entity,
namespace: "test".to_string(),
title: Some("Rust Programming".to_string()),
body: "Rust is a systems programming language.".to_string(),
tags: vec![],
metadata: None,
updated_at: chrono::Utc::now(),
})
.await
.unwrap();
store
.upsert_document(TextDocument {
subject_id: id2,
kind: SubstrateKind::Entity,
namespace: "test".to_string(),
title: Some("Python Guide".to_string()),
body: "Python is a high-level programming language.".to_string(),
tags: vec![],
metadata: None,
updated_at: chrono::Utc::now(),
})
.await
.unwrap();
let adapter = StorageKeywordSearch::new(store);
let hits = adapter.keyword_search("Rust", 10).await.unwrap();
assert!(!hits.is_empty());
assert_eq!(hits[0].0, id1);
assert!(hits[0].1.to_f64() > 0.0);
}
#[tokio::test]
async fn keyword_search_respects_top_k() {
let backend = test_backend();
let store = backend.text("test_ks_topk").unwrap();
for i in 0..5 {
store
.upsert_document(TextDocument {
subject_id: Uuid::new_v4(),
kind: SubstrateKind::Note,
namespace: "test".to_string(),
title: Some(format!("Doc {}", i)),
body: format!("Programming topic number {}.", i),
tags: vec![],
metadata: None,
updated_at: chrono::Utc::now(),
})
.await
.unwrap();
}
let adapter = StorageKeywordSearch::new(store);
let hits = adapter.keyword_search("programming", 3).await.unwrap();
assert!(hits.len() <= 3);
}
#[tokio::test]
async fn keyword_search_empty_store() {
let backend = test_backend();
let store = backend.text("test_ks_empty").unwrap();
let adapter = StorageKeywordSearch::new(store);
let hits = adapter.keyword_search("anything", 5).await.unwrap();
assert!(hits.is_empty());
}
#[tokio::test]
async fn keyword_search_no_match() {
let backend = test_backend();
let store = backend.text("test_ks_nomatch").unwrap();
store
.upsert_document(TextDocument {
subject_id: Uuid::new_v4(),
kind: SubstrateKind::Entity,
namespace: "test".to_string(),
title: Some("Alpha".to_string()),
body: "Alpha article content.".to_string(),
tags: vec![],
metadata: None,
updated_at: chrono::Utc::now(),
})
.await
.unwrap();
let adapter = StorageKeywordSearch::new(store);
let hits = adapter
.keyword_search("nonexistent_xyz_term", 5)
.await
.unwrap();
assert!(hits.is_empty());
}
#[tokio::test]
async fn adapters_produce_fusible_results() {
use crate::hybrid::{fuse_search_results, HybridConfig};
let backend = test_backend();
let vec_store = backend.vectors("test_fuse", 3).unwrap();
let text_store = backend.text("test_fuse").unwrap();
let id = Uuid::new_v4();
vec_store
.insert(
id,
SubstrateKind::Note,
"local",
"content",
vec![vec![1.0, 0.0, 0.0]],
)
.await
.unwrap();
text_store
.upsert_document(TextDocument {
subject_id: id,
kind: SubstrateKind::Note,
namespace: "test".to_string(),
title: Some("Test".to_string()),
body: "Test document for fusion.".to_string(),
tags: vec![],
metadata: None,
updated_at: chrono::Utc::now(),
})
.await
.unwrap();
let vec_adapter = StorageVectorSearch::new(vec_store);
let kw_adapter = StorageKeywordSearch::new(text_store);
let vec_hits = vec_adapter
.vector_search(&[1.0, 0.0, 0.0], 5)
.await
.unwrap();
let kw_hits = kw_adapter.keyword_search("Test", 5).await.unwrap();
assert!(!vec_hits.is_empty());
assert!(!kw_hits.is_empty());
assert_eq!(vec_hits[0].0, id);
assert_eq!(kw_hits[0].0, id);
let config = HybridConfig::new(10);
let fused = fuse_search_results(vec![vec_hits, kw_hits], &config);
assert!(!fused.is_empty());
assert_eq!(fused[0].0, id);
}
}