mod config;
mod diarizer;
mod features;
pub use config::{
ADAPTER_TYPE, FRAME_DURATION_SECS, MAX_SPEAKERS, MODEL_ID, PostProcessConfig, SAMPLE_RATE,
SortformerConfig, SortformerError,
};
pub use diarizer::SortformerDiarizer;
#[cfg(feature = "download")]
use crate::models::{AdapterError, AdapterFactory, AdapterRegistry, AdapterStage, BuiltinAdapter};
#[cfg(feature = "download")]
use std::sync::Arc;
#[cfg(feature = "download")]
pub fn register_with(registry: &mut AdapterRegistry) -> Result<(), AdapterError> {
let factory: AdapterFactory = Arc::new(|| {
Box::new(BuiltinAdapter {
stage: AdapterStage::Diarizer,
id: ADAPTER_TYPE.to_owned(),
})
});
registry.register(AdapterStage::Diarizer, ADAPTER_TYPE, factory)?;
let _ = registry.register_alias(AdapterStage::Diarizer, "latest", ADAPTER_TYPE);
let _ = registry.register_alias(AdapterStage::Diarizer, "v2", ADAPTER_TYPE);
Ok(())
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn adapter_type_constant_matches_manifest_convention() {
assert_eq!(ADAPTER_TYPE, "sortformer-v2");
assert_eq!(MODEL_ID, "sortformer_v2");
assert_eq!(MAX_SPEAKERS, 4);
}
#[cfg(feature = "download")]
#[test]
fn register_with_empty_registry() {
let mut reg = AdapterRegistry::new();
register_with(&mut reg).unwrap();
assert!(reg.contains(AdapterStage::Diarizer, ADAPTER_TYPE));
let err = register_with(&mut reg).expect_err("duplicate");
assert!(matches!(err, AdapterError::AlreadyRegistered { .. }));
}
#[cfg(feature = "download")]
#[test]
fn register_aliases_resolve_to_adapter_type() {
let mut reg = AdapterRegistry::new();
register_with(&mut reg).unwrap();
for alias in ["latest", "v2"] {
let resolved = reg.resolve(AdapterStage::Diarizer, alias).unwrap();
assert_eq!(resolved, ADAPTER_TYPE, "alias {alias}");
}
let resolved = reg.resolve(AdapterStage::Diarizer, ADAPTER_TYPE).unwrap();
assert_eq!(resolved, ADAPTER_TYPE);
let handle = reg.create(AdapterStage::Diarizer, "v2").unwrap();
let adapter = handle.downcast::<BuiltinAdapter>().unwrap();
assert_eq!(adapter.id, ADAPTER_TYPE);
}
}