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 std::sync::atomic::{AtomicU64, Ordering};
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())
}
}
fn unique_name(suffix: &str) -> String {
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::SeqCst);
format!("mock-reranker-{suffix}-{id}")
}
#[test]
fn register_list_unregister_roundtrip() {
let name = unique_name("roundtrip");
register_reranker_backend(Arc::new(MockRerankerBackend { name: name.clone() })).unwrap();
assert!(list_reranker_backends().unwrap().contains(&name));
unregister_reranker_backend(&name).unwrap();
assert!(!list_reranker_backends().unwrap().contains(&name));
}
#[test]
fn empty_name_rejected_via_global_api() {
let result = register_reranker_backend(Arc::new(MockRerankerBackend { name: String::new() }));
assert!(matches!(result, Err(XbergError::Validation { .. })));
}
#[test]
fn duplicate_name_rejected_via_global_api() {
let name = unique_name("dup");
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 { .. })));
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 name = unique_name("clear");
register_reranker_backend(Arc::new(MockRerankerBackend { name: name.clone() })).unwrap();
assert!(list_reranker_backends().unwrap().contains(&name));
clear_reranker_backends().unwrap();
assert!(!list_reranker_backends().unwrap().contains(&name));
}
}