use crate::Result;
use crate::plugins::Plugin;
use std::sync::Arc;
pub trait TokenizerBackend: Plugin {
fn count_tokens(&self, text: &str) -> usize;
}
#[cfg_attr(alef, alef(skip))]
pub fn register_tokenizer_backend(backend: Arc<dyn TokenizerBackend>) -> Result<()> {
use crate::plugins::registry::get_tokenizer_backend_registry;
let registry = get_tokenizer_backend_registry();
let mut registry = registry.write();
registry.register(backend)
}
#[cfg_attr(alef, alef(skip))]
pub fn unregister_tokenizer_backend(name: &str) -> Result<()> {
use crate::plugins::registry::get_tokenizer_backend_registry;
let registry = get_tokenizer_backend_registry();
let mut registry = registry.write();
registry.remove(name)
}
pub fn clear_tokenizer_backends() -> Result<()> {
use crate::plugins::registry::get_tokenizer_backend_registry;
let registry = get_tokenizer_backend_registry();
let mut registry = registry.write();
registry.shutdown_all()
}
pub fn list_tokenizer_backends() -> Result<Vec<String>> {
use crate::plugins::registry::get_tokenizer_backend_registry;
let registry = get_tokenizer_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 MockTokenizerBackend {
name: String,
}
impl Plugin for MockTokenizerBackend {
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(())
}
}
impl TokenizerBackend for MockTokenizerBackend {
fn count_tokens(&self, text: &str) -> usize {
text.split_whitespace().count().max(usize::from(!text.is_empty()))
}
}
fn unique_name(suffix: &str) -> String {
static COUNTER: AtomicU64 = AtomicU64::new(0);
let id = COUNTER.fetch_add(1, Ordering::SeqCst);
format!("mock-tokenizer-{suffix}-{id}")
}
#[test]
fn register_list_unregister_roundtrip() {
let name = unique_name("roundtrip");
register_tokenizer_backend(Arc::new(MockTokenizerBackend { name: name.clone() })).unwrap();
assert!(list_tokenizer_backends().unwrap().contains(&name));
unregister_tokenizer_backend(&name).unwrap();
assert!(!list_tokenizer_backends().unwrap().contains(&name));
}
#[test]
fn empty_name_rejected_via_global_api() {
let result = register_tokenizer_backend(Arc::new(MockTokenizerBackend { name: String::new() }));
assert!(matches!(result, Err(XbergError::Validation { .. })));
}
#[test]
fn register_clear_roundtrip() {
let name = unique_name("clear");
register_tokenizer_backend(Arc::new(MockTokenizerBackend { name: name.clone() })).unwrap();
assert!(list_tokenizer_backends().unwrap().contains(&name));
clear_tokenizer_backends().unwrap();
assert!(!list_tokenizer_backends().unwrap().contains(&name));
}
#[test]
fn mock_backend_counts_words() {
let backend = MockTokenizerBackend {
name: "counter".to_string(),
};
assert_eq!(backend.count_tokens("one two three"), 3);
assert_eq!(backend.count_tokens(""), 0);
}
}