use crate::Result;
use crate::plugins::Plugin;
use async_trait::async_trait;
use std::sync::Arc;
#[doc(alias = "rerank")]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
pub trait RerankerBackend: Plugin {
async fn rerank(&self, query: String, documents: Vec<String>) -> Result<Vec<f32>>;
}
#[cfg_attr(alef, alef(skip))]
pub fn register_reranker_backend(backend: Arc<dyn RerankerBackend>) -> Result<()> {
use crate::plugins::registry::get_reranker_backend_registry;
let registry = get_reranker_backend_registry();
let mut registry = registry.write();
registry.register(backend)
}
#[cfg_attr(alef, alef(skip))]
pub fn unregister_reranker_backend(name: &str) -> Result<()> {
use crate::plugins::registry::get_reranker_backend_registry;
let registry = get_reranker_backend_registry();
let mut registry = registry.write();
registry.remove(name)
}
pub fn clear_reranker_backends() -> Result<()> {
use crate::plugins::registry::get_reranker_backend_registry;
let registry = get_reranker_backend_registry();
let mut registry = registry.write();
registry.shutdown_all()
}
pub fn list_reranker_backends() -> Result<Vec<String>> {
use crate::plugins::registry::get_reranker_backend_registry;
let registry = get_reranker_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::RerankerRegistryGuard;
struct MockRerankerBackend {
name: String,
}
impl Plugin for MockRerankerBackend {
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 RerankerBackend for MockRerankerBackend {
async fn rerank(&self, _query: String, documents: Vec<String>) -> Result<Vec<f32>> {
Ok(documents.iter().map(|_| 1.0_f32).collect())
}
}
#[test]
fn register_list_unregister_roundtrip() {
let _guard = RerankerRegistryGuard::acquire();
let name = "mock-reranker-roundtrip".to_string();
register_reranker_backend(Arc::new(MockRerankerBackend { name: name.clone() })).unwrap();
assert_eq!(list_reranker_backends().unwrap(), vec![name.clone()]);
unregister_reranker_backend(&name).unwrap();
assert!(list_reranker_backends().unwrap().is_empty());
}
#[test]
fn empty_name_rejected_via_global_api() {
let _guard = RerankerRegistryGuard::acquire();
let result = register_reranker_backend(Arc::new(MockRerankerBackend { name: String::new() }));
assert!(matches!(result, Err(XbergError::Validation { .. })));
assert!(
list_reranker_backends().unwrap().is_empty(),
"a rejected registration must not leave anything behind"
);
}
#[test]
fn duplicate_name_rejected_via_global_api() {
let _guard = RerankerRegistryGuard::acquire();
let name = "mock-reranker-dup".to_string();
register_reranker_backend(Arc::new(MockRerankerBackend { name: name.clone() })).unwrap();
let result = register_reranker_backend(Arc::new(MockRerankerBackend { name: name.clone() }));
assert!(matches!(result, Err(XbergError::Plugin { .. })));
assert_eq!(
list_reranker_backends().unwrap(),
vec![name.clone()],
"the rejected duplicate must not have been added alongside the original"
);
unregister_reranker_backend(&name).unwrap();
}
#[tokio::test]
async fn mock_backend_returns_expected_shape() {
let backend = MockRerankerBackend {
name: "local".to_string(),
};
let scores = backend
.rerank("query".to_string(), vec!["doc1".into(), "doc2".into(), "doc3".into()])
.await
.unwrap();
assert_eq!(scores.len(), 3);
}
#[test]
fn register_list_clear_list_roundtrip() {
let _guard = RerankerRegistryGuard::acquire();
let name = "mock-reranker-clear".to_string();
register_reranker_backend(Arc::new(MockRerankerBackend { name: name.clone() })).unwrap();
assert_eq!(list_reranker_backends().unwrap(), vec![name]);
clear_reranker_backends().unwrap();
assert!(list_reranker_backends().unwrap().is_empty());
}
}