use crate::Result;
use crate::plugins::Plugin;
use async_trait::async_trait;
use std::sync::Arc;
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
pub trait EmbeddingBackend: Plugin {
fn dimensions(&self) -> usize;
async fn embed(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>>;
}
#[cfg_attr(alef, alef(skip))]
pub fn register_embedding_backend(backend: Arc<dyn EmbeddingBackend>) -> Result<()> {
use crate::plugins::registry::get_embedding_backend_registry;
let registry = get_embedding_backend_registry();
let mut registry = registry.write();
registry.register(backend)
}
#[cfg_attr(alef, alef(skip))]
pub fn unregister_embedding_backend(name: &str) -> Result<()> {
use crate::plugins::registry::get_embedding_backend_registry;
let registry = get_embedding_backend_registry();
let mut registry = registry.write();
registry.remove(name)
}
pub fn clear_embedding_backends() -> Result<()> {
use crate::plugins::registry::get_embedding_backend_registry;
let registry = get_embedding_backend_registry();
let mut registry = registry.write();
registry.shutdown_all()
}
pub fn list_embedding_backends() -> Result<Vec<String>> {
use crate::plugins::registry::get_embedding_backend_registry;
let registry = get_embedding_backend_registry();
let registry = registry.read();
Ok(registry.list())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::XbergError;
use crate::plugins::Plugin;
use crate::plugins::registry::test_support::EmbeddingRegistryGuard;
struct MockEmbeddingBackend {
name: String,
dimensions: usize,
}
impl Plugin for MockEmbeddingBackend {
fn name(&self) -> &str {
&self.name
}
fn version(&self) -> String {
"1.0.0".to_string()
}
fn initialize(&self) -> Result<()> {
Ok(())
}
fn shutdown(&self) -> Result<()> {
Ok(())
}
}
#[async_trait]
impl EmbeddingBackend for MockEmbeddingBackend {
fn dimensions(&self) -> usize {
self.dimensions
}
async fn embed(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>> {
Ok(texts.iter().map(|_| vec![0.5; self.dimensions]).collect())
}
}
#[test]
fn register_list_unregister_roundtrip() {
let _guard = EmbeddingRegistryGuard::acquire();
let name = "mock-roundtrip".to_string();
register_embedding_backend(Arc::new(MockEmbeddingBackend {
name: name.clone(),
dimensions: 384,
}))
.unwrap();
assert_eq!(list_embedding_backends().unwrap(), vec![name.clone()]);
unregister_embedding_backend(&name).unwrap();
assert!(list_embedding_backends().unwrap().is_empty());
}
#[test]
fn empty_name_rejected_via_global_api() {
let _guard = EmbeddingRegistryGuard::acquire();
let result = register_embedding_backend(Arc::new(MockEmbeddingBackend {
name: String::new(),
dimensions: 384,
}));
assert!(matches!(result, Err(XbergError::Validation { .. })));
assert!(
list_embedding_backends().unwrap().is_empty(),
"a rejected registration must not leave anything behind"
);
}
#[test]
fn zero_dimensions_rejected_via_global_api() {
let _guard = EmbeddingRegistryGuard::acquire();
let result = register_embedding_backend(Arc::new(MockEmbeddingBackend {
name: "mock-zero".to_string(),
dimensions: 0,
}));
assert!(matches!(result, Err(XbergError::Validation { .. })));
assert!(
list_embedding_backends().unwrap().is_empty(),
"a rejected registration must not leave anything behind"
);
}
#[tokio::test]
async fn mock_backend_returns_expected_shape() {
let backend = MockEmbeddingBackend {
name: "local".to_string(),
dimensions: 5,
};
let vectors = backend
.embed(vec!["one".into(), "two".into(), "three".into()])
.await
.unwrap();
assert_eq!(vectors.len(), 3);
assert!(vectors.iter().all(|v| v.len() == 5));
}
#[test]
fn register_list_clear_list_roundtrip() {
let _guard = EmbeddingRegistryGuard::acquire();
let name = "mock-clear".to_string();
register_embedding_backend(Arc::new(MockEmbeddingBackend {
name: name.clone(),
dimensions: 128,
}))
.unwrap();
assert_eq!(list_embedding_backends().unwrap(), vec![name]);
clear_embedding_backends().unwrap();
assert!(list_embedding_backends().unwrap().is_empty());
}
}