use std::sync::Arc;
use tokio::sync::RwLock;
use crate::RerankConfig;
use crate::config::model::ModelConfig;
use crate::domain::{
BatchEmbedRequest, BatchEmbedResponse, EmbedRequest, EmbedResponse, RerankRequest,
RerankResponse,
};
use crate::engine::{AnyEngine, EngineFactory};
use crate::error::VecboostError;
use crate::registry::{EmbeddingModule, RerankModule};
use crate::service::embedding::EmbeddingService;
use crate::service::rerank::RerankService;
#[derive(Debug, Clone, Default)]
pub struct LibraryConfig {
pub model_config: ModelConfig,
pub cache_size: usize,
pub rerank_config: Option<RerankConfig>,
}
impl LibraryConfig {
pub fn from_model_config(model_config: ModelConfig) -> Self {
Self {
model_config,
..Default::default()
}
}
}
pub struct VecBoostLibrary {
kit: Arc<trait_kit::AsyncKit<trait_kit::AsyncReady>>,
}
impl VecBoostLibrary {
pub async fn new(config: LibraryConfig) -> Result<Self, VecboostError> {
let engine = EngineFactory::create(
config.model_config.engine_type.clone(),
&config.model_config,
)?;
let engine: Arc<RwLock<AnyEngine>> = Arc::new(RwLock::new(engine));
let embedding_service = if config.cache_size > 0 {
Arc::new(RwLock::new(EmbeddingService::with_cache(
engine.clone(),
Some(config.model_config.clone()),
config.cache_size,
)))
} else {
Arc::new(RwLock::new(EmbeddingService::new(
engine.clone(),
Some(config.model_config.clone()),
)))
};
let rerank_service = Arc::new(RwLock::new(RerankService::new(
engine,
Some(config.model_config),
)));
let mut kit = trait_kit::AsyncKit::new();
kit.set_config(embedding_service);
kit.set_config(rerank_service);
kit.set_config(config.rerank_config.unwrap_or_default());
kit.register::<EmbeddingModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to register EmbeddingModule: {}", e))
})?;
kit.register::<RerankModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to register RerankModule: {}", e))
})?;
kit.register_lifecycle::<EmbeddingModule>();
kit.register_lifecycle::<RerankModule>();
let kit = kit.build().await.map_err(|e| {
VecboostError::InternalError(format!("Failed to build AsyncKit: {}", e))
})?;
Ok(Self { kit: Arc::new(kit) })
}
pub async fn embed(&self, text: &str) -> Result<EmbedResponse, VecboostError> {
let service = self.kit.require::<EmbeddingModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to require EmbeddingModule: {}", e))
})?;
let svc = service.read().await;
svc.process_text(
EmbedRequest {
text: text.to_string(),
normalize: None,
},
None,
)
.await
}
pub async fn embed_batch(&self, texts: &[String]) -> Result<BatchEmbedResponse, VecboostError> {
let service = self.kit.require::<EmbeddingModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to require EmbeddingModule: {}", e))
})?;
let svc = service.read().await;
svc.process_batch(
BatchEmbedRequest {
texts: texts.to_vec(),
mode: None,
normalize: None,
},
None,
)
.await
}
pub async fn rerank(
&self,
query: &str,
documents: &[String],
top_k: Option<usize>,
) -> Result<RerankResponse, VecboostError> {
let service = self.kit.require::<RerankModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to require RerankModule: {}", e))
})?;
let rerank_config = self.kit.config::<RerankConfig>().ok().unwrap_or_default();
let svc = service.read().await;
svc.process_rerank(
RerankRequest {
query: query.to_string(),
documents: documents.to_vec(),
top_k,
return_documents: None,
},
rerank_config.max_documents_per_query,
rerank_config.max_query_length,
)
.await
}
pub fn embed_sync(&self, text: &str) -> Result<EmbedResponse, VecboostError> {
let future = self.embed(text);
Self::block_on_future(future)
}
pub fn embed_batch_sync(&self, texts: &[String]) -> Result<BatchEmbedResponse, VecboostError> {
let future = self.embed_batch(texts);
Self::block_on_future(future)
}
pub fn rerank_sync(
&self,
query: &str,
documents: &[String],
top_k: Option<usize>,
) -> Result<RerankResponse, VecboostError> {
let future = self.rerank(query, documents, top_k);
Self::block_on_future(future)
}
fn block_on_future<F: std::future::Future>(future: F) -> F::Output
where
F::Output: IntoSyncResult,
{
if tokio::runtime::Handle::try_current().is_ok() {
return F::Output::into_sync_result(
"sync API called from within a tokio runtime; use the async variants (embed / embed_batch / rerank) instead",
);
}
static SHARED_RUNTIME: std::sync::OnceLock<std::io::Result<tokio::runtime::Runtime>> =
std::sync::OnceLock::new();
let rt = SHARED_RUNTIME.get_or_init(|| {
tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_all()
.build()
});
match rt {
Ok(rt) => rt.block_on(future),
Err(e) => F::Output::into_sync_result(format!(
"Failed to create tokio runtime for sync API: {e}"
)),
}
}
}
trait IntoSyncResult {
fn into_sync_result(message: impl Into<String>) -> Self;
}
impl<T> IntoSyncResult for Result<T, VecboostError> {
fn into_sync_result(message: impl Into<String>) -> Self {
Err(VecboostError::InternalError(message.into()))
}
}
pub struct VecBoostModuleBuilder {
model_config: ModelConfig,
cache_size: usize,
with_embedding: bool,
with_rerank: bool,
rerank_config: Option<RerankConfig>,
}
impl VecBoostModuleBuilder {
pub fn new(model_config: ModelConfig) -> Self {
Self {
model_config,
cache_size: 0,
with_embedding: false,
with_rerank: false,
rerank_config: None,
}
}
pub fn embedding(mut self) -> Self {
self.with_embedding = true;
self
}
pub fn rerank(mut self) -> Self {
self.with_rerank = true;
self
}
pub fn cache_size(mut self, size: usize) -> Self {
self.cache_size = size;
self
}
pub fn rerank_config(mut self, config: RerankConfig) -> Self {
self.rerank_config = Some(config);
self
}
pub async fn build(self, kit: &mut trait_kit::AsyncKit) -> Result<(), VecboostError> {
let engine =
EngineFactory::create(self.model_config.engine_type.clone(), &self.model_config)?;
let engine: Arc<RwLock<AnyEngine>> = Arc::new(RwLock::new(engine));
if self.with_embedding {
let embedding_service = if self.cache_size > 0 {
Arc::new(RwLock::new(EmbeddingService::with_cache(
engine.clone(),
Some(self.model_config.clone()),
self.cache_size,
)))
} else {
Arc::new(RwLock::new(EmbeddingService::new(
engine.clone(),
Some(self.model_config.clone()),
)))
};
kit.set_config(embedding_service);
kit.register::<EmbeddingModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to register EmbeddingModule: {}", e))
})?;
kit.register_lifecycle::<EmbeddingModule>();
}
if self.with_rerank {
let rerank_service = Arc::new(RwLock::new(RerankService::new(
engine,
Some(self.model_config),
)));
kit.set_config(rerank_service);
kit.set_config(self.rerank_config.unwrap_or_default());
kit.register::<RerankModule>().map_err(|e| {
VecboostError::InternalError(format!("Failed to register RerankModule: {}", e))
})?;
kit.register_lifecycle::<RerankModule>();
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::model::Precision;
use crate::engine::InferenceEngine;
use async_trait::async_trait;
struct MockEngine {
dimension: usize,
}
impl MockEngine {
fn new(dimension: usize) -> Self {
Self { dimension }
}
}
#[async_trait]
impl InferenceEngine for MockEngine {
fn embed(&self, _text: &str) -> Result<Vec<f32>, VecboostError> {
Ok(vec![1.0; self.dimension])
}
fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, VecboostError> {
Ok(texts.iter().map(|_| vec![1.0; self.dimension]).collect())
}
fn precision(&self) -> &Precision {
&Precision::Fp32
}
fn supports_mixed_precision(&self) -> bool {
false
}
fn rerank(&self, _query: &str, document: &str) -> Result<f32, VecboostError> {
Ok((document.len() as f32) / 100.0)
}
fn rerank_batch(
&self,
query: &str,
documents: &[String],
) -> Result<Vec<f32>, VecboostError> {
documents
.iter()
.map(|doc| self.rerank(query, doc))
.collect()
}
fn supports_rerank(&self) -> bool {
true
}
async fn try_fallback_to_cpu(
&mut self,
_config: &ModelConfig,
) -> Result<(), VecboostError> {
Ok(())
}
}
async fn make_test_library() -> VecBoostLibrary {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine::new(128)));
let embedding_service = Arc::new(RwLock::new(EmbeddingService::new(engine.clone(), None)));
let rerank_service = Arc::new(RwLock::new(RerankService::new(engine, None)));
let mut kit = trait_kit::AsyncKit::new();
kit.set_config(embedding_service);
kit.set_config(rerank_service);
kit.set_config(RerankConfig::default());
kit.register::<EmbeddingModule>().unwrap();
kit.register::<RerankModule>().unwrap();
let kit = kit.build().await.unwrap();
VecBoostLibrary { kit: Arc::new(kit) }
}
#[tokio::test]
async fn test_library_construction_succeeds() {
let lib = make_test_library().await;
assert!(lib.kit.contains::<EmbeddingModule>());
assert!(lib.kit.contains::<RerankModule>());
}
#[tokio::test]
async fn test_embed_returns_correct_dimension() {
let lib = make_test_library().await;
let response = lib.embed("hello world").await.unwrap();
assert_eq!(response.dimension, 128);
assert_eq!(response.embedding.len(), 128);
}
#[tokio::test]
async fn test_embed_batch_returns_correct_count() {
let lib = make_test_library().await;
let texts = vec!["hello".to_string(), "world".to_string()];
let response = lib.embed_batch(&texts).await.unwrap();
assert_eq!(response.embeddings.len(), 2);
}
#[tokio::test]
async fn test_rerank_returns_sorted_results() {
let lib = make_test_library().await;
let response = lib
.rerank(
"what is rust?",
&[
"short".to_string(),
"a much longer document about programming".to_string(),
"medium length doc".to_string(),
],
None,
)
.await
.unwrap();
assert_eq!(response.results.len(), 3);
for i in 1..response.results.len() {
assert!(
response.results[i - 1].score >= response.results[i].score,
"Results not sorted by score descending"
);
}
}
#[tokio::test]
async fn test_rerank_with_top_k() {
let lib = make_test_library().await;
let response = lib
.rerank(
"test",
&["a".to_string(), "bb".to_string(), "ccc".to_string()],
Some(2),
)
.await
.unwrap();
assert_eq!(response.results.len(), 2);
}
#[tokio::test]
async fn sync_api_inside_runtime_returns_error_not_panic() {
let lib = make_test_library().await;
let err = lib
.embed_sync("hello")
.expect_err("must error inside tokio context");
let msg = err.error_detail().to_string();
assert!(
msg.contains("async variants") || msg.contains("tokio runtime"),
"error should point to async alternatives, got: {msg}"
);
}
#[test]
fn test_sync_embed_outside_runtime() {
let lib = {
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(make_test_library())
};
let response = lib.embed_sync("hello").unwrap();
assert_eq!(response.dimension, 128);
}
#[test]
fn test_sync_rerank_outside_runtime() {
let lib = {
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(make_test_library())
};
let response = lib
.rerank_sync(
"query",
&["doc a".to_string(), "doc longer b".to_string()],
None,
)
.unwrap();
assert_eq!(response.results.len(), 2);
}
#[test]
fn test_library_config_default() {
let config = LibraryConfig::default();
assert_eq!(config.cache_size, 0);
assert!(config.rerank_config.is_none());
}
#[test]
fn test_library_config_from_model_config() {
let model_config = ModelConfig {
name: "test-model".to_string(),
expected_dimension: Some(768),
..Default::default()
};
let config = LibraryConfig::from_model_config(model_config);
assert_eq!(config.model_config.name, "test-model");
assert_eq!(config.model_config.expected_dimension, Some(768));
assert_eq!(config.cache_size, 0);
}
async fn make_test_kit_with_builder(
with_embedding: bool,
with_rerank: bool,
) -> trait_kit::AsyncKit<trait_kit::AsyncReady> {
let engine: Arc<RwLock<dyn InferenceEngine + Send + Sync>> =
Arc::new(RwLock::new(MockEngine::new(128)));
let mut kit = trait_kit::AsyncKit::new();
if with_embedding {
let embedding_service =
Arc::new(RwLock::new(EmbeddingService::new(engine.clone(), None)));
kit.set_config(embedding_service);
kit.register::<EmbeddingModule>().unwrap();
kit.register_lifecycle::<EmbeddingModule>();
}
if with_rerank {
let rerank_service = Arc::new(RwLock::new(RerankService::new(engine, None)));
kit.set_config(rerank_service);
kit.set_config(RerankConfig::default());
kit.register::<RerankModule>().unwrap();
kit.register_lifecycle::<RerankModule>();
}
kit.build().await.unwrap()
}
#[tokio::test]
async fn test_module_builder_embedding_only() {
let kit = make_test_kit_with_builder(true, false).await;
assert!(kit.contains::<EmbeddingModule>());
assert!(!kit.contains::<RerankModule>());
let svc = kit.require::<EmbeddingModule>().unwrap();
let guard = svc.read().await;
let resp = guard
.process_text(
EmbedRequest {
text: "test".to_string(),
normalize: None,
},
None,
)
.await
.unwrap();
assert_eq!(resp.dimension, 128);
}
#[tokio::test]
async fn test_module_builder_rerank_only() {
let kit = make_test_kit_with_builder(false, true).await;
assert!(!kit.contains::<EmbeddingModule>());
assert!(kit.contains::<RerankModule>());
let svc = kit.require::<RerankModule>().unwrap();
let guard = svc.read().await;
let resp = guard
.process_rerank(
RerankRequest {
query: "query".to_string(),
documents: vec!["doc a".to_string(), "doc b".to_string()],
top_k: None,
return_documents: None,
},
100,
1000,
)
.await
.unwrap();
assert_eq!(resp.results.len(), 2);
}
#[tokio::test]
async fn test_module_builder_both_modules() {
let kit = make_test_kit_with_builder(true, true).await;
assert!(kit.contains::<EmbeddingModule>());
assert!(kit.contains::<RerankModule>());
let _embed_svc = kit.require::<EmbeddingModule>().unwrap();
let _rerank_svc = kit.require::<RerankModule>().unwrap();
}
#[tokio::test]
async fn test_module_builder_neither_module() {
let kit = make_test_kit_with_builder(false, false).await;
assert!(!kit.contains::<EmbeddingModule>());
assert!(!kit.contains::<RerankModule>());
}
#[test]
fn test_module_builder_chaining_api() {
let model_config = ModelConfig::default();
let _builder = VecBoostModuleBuilder::new(model_config)
.embedding()
.rerank()
.cache_size(500)
.rerank_config(RerankConfig::default());
}
#[tokio::test]
async fn test_embed_empty_string_returns_error() {
let lib = make_test_library().await;
let result = lib.embed("").await;
assert!(
result.is_err(),
"empty string should return validation error"
);
}
#[tokio::test]
async fn test_rerank_empty_documents_returns_error() {
let lib = make_test_library().await;
let result = lib.rerank("query", &[], None).await;
assert!(result.is_err(), "empty documents should return error");
}
#[tokio::test]
async fn test_rerank_empty_query_returns_error() {
let lib = make_test_library().await;
let result = lib.rerank("", &["doc".to_string()], None).await;
assert!(result.is_err(), "empty query should return error");
}
#[test]
fn test_library_config_with_cache_size() {
let config = LibraryConfig {
cache_size: 1000,
..Default::default()
};
assert_eq!(config.cache_size, 1000);
}
#[test]
fn test_library_config_with_rerank_config() {
let config = LibraryConfig {
rerank_config: Some(RerankConfig::default()),
..Default::default()
};
assert!(config.rerank_config.is_some());
}
#[test]
fn test_module_builder_new_defaults() {
let builder = VecBoostModuleBuilder::new(ModelConfig::default());
assert!(!builder.with_embedding);
assert!(!builder.with_rerank);
assert_eq!(builder.cache_size, 0);
assert!(builder.rerank_config.is_none());
}
#[test]
fn test_module_builder_embedding_only_chain() {
let builder = VecBoostModuleBuilder::new(ModelConfig::default())
.embedding()
.cache_size(100);
assert!(builder.with_embedding);
assert!(!builder.with_rerank);
assert_eq!(builder.cache_size, 100);
}
#[test]
fn test_module_builder_rerank_only_chain() {
let builder = VecBoostModuleBuilder::new(ModelConfig::default())
.rerank()
.rerank_config(RerankConfig::default());
assert!(!builder.with_embedding);
assert!(builder.with_rerank);
assert!(builder.rerank_config.is_some());
}
#[test]
fn test_mock_engine_trait_method_coverage() {
let engine = MockEngine::new(256);
assert_eq!(engine.embed("test").unwrap().len(), 256);
assert_eq!(
engine.embed_batch(&["a".into(), "b".into()]).unwrap().len(),
2
);
assert_eq!(*engine.precision(), Precision::Fp32);
assert!(!engine.supports_mixed_precision());
assert!(engine.supports_rerank());
assert!(engine.rerank("q", "doc").unwrap() > 0.0);
assert_eq!(
engine
.rerank_batch("q", &["a".into(), "b".into()])
.unwrap()
.len(),
2
);
}
#[tokio::test]
async fn test_mock_engine_try_fallback_coverage() {
let mut engine = MockEngine::new(64);
let cfg = ModelConfig::default();
let _ = engine.try_fallback_to_cpu(&cfg).await;
}
}