use super::*;
fn base_offline_input<'a>(
raw: &'a [f32],
segs: &'a [f64],
count: &'a [u8],
plda: &'a diaric::plda::PldaTransform,
) -> diaric::offline::OfflineInput<'a> {
diaric::offline::OfflineInput::new(
raw,
1,
3,
segs,
1,
count,
1,
diaric::reconstruct::SlidingWindow::new(0.0, 10.0, 1.0),
diaric::reconstruct::SlidingWindow::new(0.0, 0.0619375, 0.016875),
plda,
)
}
#[test]
fn as_str_round_trips_from_str_for_every_spelling() {
for &sp in CLUSTER_BACKEND_SPELLINGS {
let parsed: ClusterBackend = sp.parse().expect("table spelling parses");
assert_eq!(
parsed.as_str(),
sp,
"FromStr → as_str must round-trip the discriminant"
);
}
}
#[test]
fn display_matches_as_str_for_every_spelling() {
for &sp in CLUSTER_BACKEND_SPELLINGS {
let parsed: ClusterBackend = sp.parse().expect("table spelling parses");
assert_eq!(
parsed.to_string(),
sp,
"Display must equal as_str (derive_more delegation)"
);
}
}
#[test]
fn spellings_roster_is_the_known_set() {
assert_eq!(CLUSTER_BACKEND_SPELLINGS, &["offline", "online"]);
}
#[test]
fn offline_spelling_maps_to_default_payload() {
assert_eq!(
"offline".parse::<ClusterBackend>().unwrap(),
ClusterBackend::Offline(OfflineOptions::new())
);
}
#[test]
fn online_spelling_maps_to_default_payload() {
assert_eq!(
"online".parse::<ClusterBackend>().unwrap(),
ClusterBackend::Online(OnlineOptions::new())
);
}
#[test]
fn default_is_still_offline_not_online() {
assert_eq!(
ClusterBackend::default(),
ClusterBackend::Offline(OfflineOptions::new())
);
assert_ne!(
ClusterBackend::default(),
ClusterBackend::Online(OnlineOptions::new())
);
}
#[test]
fn unknown_name_is_opaque_error() {
assert!("".parse::<ClusterBackend>().is_err());
assert!("Offline".parse::<ClusterBackend>().is_err()); assert!("Online".parse::<ClusterBackend>().is_err()); assert_eq!(
"nope".parse::<ClusterBackend>().unwrap_err(),
"also-nope".parse::<ClusterBackend>().unwrap_err()
);
}
#[test]
fn default_is_offline_with_default_options() {
assert_eq!(
ClusterBackend::default(),
ClusterBackend::Offline(OfflineOptions::new())
);
}
#[test]
fn cluster_backend_is_copy() {
let a = ClusterBackend::default();
let b = a;
assert_eq!(a, b);
}
#[test]
fn offline_options_new_matches_default() {
assert_eq!(OfflineOptions::new(), OfflineOptions::default());
}
#[test]
fn offline_options_defaults_match_const_literals() {
let o = OfflineOptions::new();
assert_eq!(o.threshold(), DEFAULT_THRESHOLD);
assert_eq!(o.fa(), DEFAULT_FA);
assert_eq!(o.fb(), DEFAULT_FB);
assert_eq!(o.max_iters(), DEFAULT_MAX_ITERS);
assert_eq!(o.min_duration_off(), DEFAULT_MIN_DURATION_OFF);
assert_eq!(o.threshold(), 0.6);
assert_eq!(o.fa(), 0.07);
assert_eq!(o.fb(), 0.8);
assert_eq!(o.max_iters(), 20);
assert_eq!(o.min_duration_off(), 0.0);
}
#[test]
fn defaults_equal_diaric() {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let raw = vec![0.0f32; crate::audio::speaker::embed::EMBEDDING_DIM * 3];
let segs = vec![0.0f64; 3];
let count = vec![0u8; 1];
let input = base_offline_input(&raw, &segs, &count, &plda);
let o = OfflineOptions::default();
assert_eq!(
o.threshold(),
input.threshold(),
"threshold drifted from diaric"
);
assert_eq!(o.fa(), input.fa(), "fa drifted from diaric");
assert_eq!(o.fb(), input.fb(), "fb drifted from diaric");
assert_eq!(
o.max_iters(),
input.max_iters(),
"max_iters drifted from diaric"
);
assert_eq!(
o.min_duration_off(),
input.min_duration_off(),
"min_duration_off drifted from diaric"
);
}
#[test]
fn offline_options_builders_and_setters() {
let o = OfflineOptions::new()
.with_threshold(0.33)
.with_fa(0.11)
.with_fb(0.77)
.with_max_iters(9)
.with_min_duration_off(0.5);
assert_eq!(o.threshold(), 0.33);
assert_eq!(o.fa(), 0.11);
assert_eq!(o.fb(), 0.77);
assert_eq!(o.max_iters(), 9);
assert_eq!(o.min_duration_off(), 0.5);
let mut m = OfflineOptions::new();
m.set_threshold(0.1);
m.set_fa(0.2);
m.set_fb(0.3);
m.set_max_iters(4);
m.set_min_duration_off(0.05);
assert_eq!(m.threshold(), 0.1);
assert_eq!(m.fa(), 0.2);
assert_eq!(m.fb(), 0.3);
assert_eq!(m.max_iters(), 4);
assert_eq!(m.min_duration_off(), 0.05);
}
#[test]
fn threshold_fa_fb_builders_accept_non_finite_like_dia() {
let o = OfflineOptions::new()
.with_threshold(f64::NAN)
.with_fa(f64::INFINITY)
.with_fb(f64::NEG_INFINITY);
assert!(o.threshold().is_nan());
assert_eq!(o.fa(), f64::INFINITY);
assert_eq!(o.fb(), f64::NEG_INFINITY);
}
#[test]
#[should_panic(expected = "min_duration_off must be finite and >= 0")]
fn with_min_duration_off_panics_on_negative() {
let _ = OfflineOptions::new().with_min_duration_off(-1.0);
}
#[test]
#[should_panic(expected = "min_duration_off must be finite and >= 0")]
fn with_min_duration_off_panics_on_nan() {
let _ = OfflineOptions::new().with_min_duration_off(f64::NAN);
}
#[test]
#[should_panic(expected = "min_duration_off must be finite and >= 0")]
fn with_min_duration_off_panics_on_positive_infinity() {
let _ = OfflineOptions::new().with_min_duration_off(f64::INFINITY);
}
#[test]
fn with_min_duration_off_accepts_zero_and_positive() {
assert_eq!(
OfflineOptions::new()
.with_min_duration_off(0.0)
.min_duration_off(),
0.0
);
assert_eq!(
OfflineOptions::new()
.with_min_duration_off(2.5)
.min_duration_off(),
2.5
);
}
#[test]
fn apply_to_maps_each_knob_to_its_dia_field() {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let raw = vec![0.0f32; crate::audio::speaker::embed::EMBEDDING_DIM * 3];
let segs = vec![0.0f64; 3];
let count = vec![0u8; 1];
let base = base_offline_input(&raw, &segs, &count, &plda);
let opts = OfflineOptions::new()
.with_threshold(0.31)
.with_fa(0.12)
.with_fb(0.73)
.with_max_iters(7)
.with_min_duration_off(0.4);
let input = opts.apply_to(base);
assert_eq!(input.threshold(), 0.31);
assert_eq!(input.fa(), 0.12);
assert_eq!(input.fb(), 0.73);
assert_eq!(input.max_iters(), 7);
assert_eq!(input.min_duration_off(), 0.4);
assert_eq!(input.raw_embeddings(), raw.as_slice());
assert_eq!(input.count(), count.as_slice());
}
#[test]
fn apply_to_default_is_a_no_op_over_dia_defaults() {
let plda = diaric::plda::PldaTransform::new().expect("hermetic PLDA weights load");
let raw = vec![0.0f32; crate::audio::speaker::embed::EMBEDDING_DIM * 3];
let segs = vec![0.0f64; 3];
let count = vec![0u8; 1];
let base = base_offline_input(&raw, &segs, &count, &plda);
let (t, fa, fb, mi, md) = (
base.threshold(),
base.fa(),
base.fb(),
base.max_iters(),
base.min_duration_off(),
);
let out = OfflineOptions::default().apply_to(base);
assert_eq!(out.threshold(), t);
assert_eq!(out.fa(), fa);
assert_eq!(out.fb(), fb);
assert_eq!(out.max_iters(), mi);
assert_eq!(out.min_duration_off(), md);
}
#[cfg(feature = "serde")]
#[test]
fn serde_discriminant_tag_equals_as_str_for_every_spelling() {
for &sp in CLUSTER_BACKEND_SPELLINGS {
let backend: ClusterBackend = sp.parse().unwrap();
let value = serde_json::to_value(backend).unwrap();
let obj = value
.as_object()
.expect("externally-tagged enum is an object");
assert_eq!(obj.len(), 1, "exactly one discriminant key");
assert_eq!(
obj.keys().next().unwrap(),
sp,
"serde tag must equal as_str"
);
}
}
#[cfg(feature = "serde")]
#[test]
fn serde_offline_empty_payload_is_full_defaults() {
let b: ClusterBackend = serde_json::from_str(r#"{"offline":{}}"#).unwrap();
assert_eq!(b, ClusterBackend::Offline(OfflineOptions::new()));
}
#[cfg(feature = "serde")]
#[test]
fn serde_offline_partial_payload_defaults_other_knobs() {
let b: ClusterBackend = serde_json::from_str(r#"{"offline":{"threshold":0.42}}"#).unwrap();
let ClusterBackend::Offline(o) = b else {
panic!("expected Offline")
};
assert_eq!(o.threshold(), 0.42);
assert_eq!(o.fa(), DEFAULT_FA);
assert_eq!(o.fb(), DEFAULT_FB);
assert_eq!(o.max_iters(), DEFAULT_MAX_ITERS);
assert_eq!(o.min_duration_off(), DEFAULT_MIN_DURATION_OFF);
}
#[cfg(feature = "serde")]
#[test]
fn serde_non_default_round_trips_without_silent_flip() {
let b = ClusterBackend::Offline(
OfflineOptions::new()
.with_threshold(0.55)
.with_fa(0.09)
.with_fb(0.71)
.with_max_iters(33)
.with_min_duration_off(1.25),
);
let json = serde_json::to_string(&b).unwrap();
let back: ClusterBackend = serde_json::from_str(&json).unwrap();
assert_eq!(back, b);
}
#[cfg(feature = "serde")]
#[test]
fn serde_options_empty_object_is_full_defaults() {
let o: OfflineOptions = serde_json::from_str("{}").unwrap();
assert_eq!(o, OfflineOptions::new());
}
#[cfg(feature = "serde")]
#[test]
fn serde_serialize_rejects_non_finite_threshold_fa_fb() {
assert!(serde_json::to_string(&OfflineOptions::new().with_threshold(f64::NAN)).is_err());
assert!(serde_json::to_string(&OfflineOptions::new().with_fa(f64::INFINITY)).is_err());
assert!(serde_json::to_string(&OfflineOptions::new().with_fb(f64::NEG_INFINITY)).is_err());
assert!(
serde_json::to_string(&ClusterBackend::Offline(
OfflineOptions::new().with_threshold(f64::NAN)
))
.is_err()
);
}
#[cfg(feature = "serde")]
#[test]
fn serde_deserialize_rejects_negative_min_duration_off() {
assert!(serde_json::from_str::<OfflineOptions>(r#"{"min_duration_off":-1.0}"#).is_err());
assert!(
serde_json::from_str::<ClusterBackend>(r#"{"offline":{"min_duration_off":-0.001}}"#).is_err()
);
}
#[cfg(feature = "serde")]
#[test]
fn serde_deserialize_accepts_valid_min_duration_off() {
let o: OfflineOptions = serde_json::from_str(r#"{"min_duration_off":0.75}"#).unwrap();
assert_eq!(o.min_duration_off(), 0.75);
let z: OfflineOptions = serde_json::from_str(r#"{"min_duration_off":0.0}"#).unwrap();
assert_eq!(z.min_duration_off(), 0.0);
}
#[test]
fn online_options_new_matches_default() {
assert_eq!(OnlineOptions::new(), OnlineOptions::default());
}
#[test]
fn online_options_defaults_match_const_literals() {
let o = OnlineOptions::new();
assert_eq!(o.speaker_threshold(), DEFAULT_SPEAKER_THRESHOLD);
assert_eq!(o.embedding_threshold(), DEFAULT_EMBEDDING_THRESHOLD);
assert_eq!(o.min_speech_duration(), DEFAULT_MIN_SPEECH_DURATION);
assert_eq!(o.speaker_threshold(), 0.65);
assert_eq!(o.embedding_threshold(), 0.45);
assert_eq!(o.min_speech_duration(), 1.0);
}
#[test]
fn online_defaults_equal_diaric() {
let diaric_default = diaric::cluster::online::OnlineClusterOptions::default();
let o = OnlineOptions::default();
assert_eq!(
o.speaker_threshold(),
diaric_default.speaker_threshold(),
"speaker_threshold drifted from diaric"
);
assert_eq!(
o.embedding_threshold(),
diaric_default.embedding_threshold(),
"embedding_threshold drifted from diaric"
);
assert_eq!(
o.min_speech_duration(),
diaric_default.min_speech_duration(),
"min_speech_duration drifted from diaric"
);
}
#[test]
fn online_from_clustering_threshold_matches_dia_ratios() {
let o = OnlineOptions::from_clustering_threshold(0.7);
let dia = diaric::cluster::online::OnlineClusterOptions::from_clustering_threshold(0.7);
assert_eq!(o.speaker_threshold(), dia.speaker_threshold());
assert_eq!(o.embedding_threshold(), dia.embedding_threshold());
assert_eq!(o.min_speech_duration(), dia.min_speech_duration());
assert!((o.speaker_threshold() - 0.84).abs() < 1e-6);
assert!((o.embedding_threshold() - 0.56).abs() < 1e-6);
assert_eq!(o.min_speech_duration(), 1.0);
}
#[test]
fn online_options_builders_and_setters() {
let o = OnlineOptions::new()
.with_speaker_threshold(0.9)
.with_embedding_threshold(0.3)
.with_min_speech_duration(2.5);
assert_eq!(o.speaker_threshold(), 0.9);
assert_eq!(o.embedding_threshold(), 0.3);
assert_eq!(o.min_speech_duration(), 2.5);
let mut m = OnlineOptions::new();
m.set_speaker_threshold(1.1);
m.set_embedding_threshold(0.2);
m.set_min_speech_duration(0.0);
assert_eq!(m.speaker_threshold(), 1.1);
assert_eq!(m.embedding_threshold(), 0.2);
assert_eq!(m.min_speech_duration(), 0.0);
}
#[test]
fn online_threshold_boundaries_accept_zero_and_two() {
assert_eq!(
OnlineOptions::new()
.with_speaker_threshold(0.0)
.speaker_threshold(),
0.0
);
assert_eq!(
OnlineOptions::new()
.with_embedding_threshold(2.0)
.embedding_threshold(),
2.0
);
}
#[test]
#[should_panic(expected = "speaker_threshold must be a finite cosine distance")]
fn with_speaker_threshold_panics_on_nan() {
let _ = OnlineOptions::new().with_speaker_threshold(f32::NAN);
}
#[test]
#[should_panic(expected = "speaker_threshold must be a finite cosine distance")]
fn with_speaker_threshold_panics_above_two() {
let _ = OnlineOptions::new().with_speaker_threshold(2.5);
}
#[test]
#[should_panic(expected = "speaker_threshold must be a finite cosine distance")]
fn with_speaker_threshold_panics_on_negative() {
let _ = OnlineOptions::new().with_speaker_threshold(-0.1);
}
#[test]
#[should_panic(expected = "embedding_threshold must be a finite cosine distance")]
fn with_embedding_threshold_panics_on_infinity() {
let _ = OnlineOptions::new().with_embedding_threshold(f32::INFINITY);
}
#[test]
#[should_panic(expected = "min_speech_duration must be finite and >= 0")]
fn with_min_speech_duration_panics_on_negative() {
let _ = OnlineOptions::new().with_min_speech_duration(-1.0);
}
#[test]
#[should_panic(expected = "min_speech_duration must be finite and >= 0")]
fn with_min_speech_duration_panics_on_infinity() {
let _ = OnlineOptions::new().with_min_speech_duration(f32::INFINITY);
}
#[test]
#[should_panic(expected = "speaker_threshold must be a finite cosine distance")]
fn from_clustering_threshold_overflow_panics() {
let _ = OnlineOptions::from_clustering_threshold(2.0);
}
#[test]
fn online_to_dia_options_maps_each_knob() {
let opts = OnlineOptions::new()
.with_speaker_threshold(0.71)
.with_embedding_threshold(0.33)
.with_min_speech_duration(1.75);
let dia = opts.to_dia_options();
assert_eq!(dia.speaker_threshold(), 0.71);
assert_eq!(dia.embedding_threshold(), 0.33);
assert_eq!(dia.min_speech_duration(), 1.75);
}
#[test]
fn online_to_dia_options_default_equals_dia_default() {
let dia = OnlineOptions::default().to_dia_options();
let dia_default = diaric::cluster::online::OnlineClusterOptions::default();
assert_eq!(dia.speaker_threshold(), dia_default.speaker_threshold());
assert_eq!(dia.embedding_threshold(), dia_default.embedding_threshold());
assert_eq!(dia.min_speech_duration(), dia_default.min_speech_duration());
}
#[test]
fn online_clusterer_is_deterministic_given_fixed_order() {
use diaric::{
cluster::online::{Assignment, OnlineClusterer},
embed::{EMBEDDING_DIM, Embedding},
};
let make = |block: usize| -> Embedding {
let mut raw = [0.0f32; EMBEDDING_DIM];
raw[(block * 64)..((block + 1) * 64)].fill(1.0);
Embedding::normalize_from(raw).expect("nonzero")
};
let seq = [
(make(0), 2.0f32),
(make(1), 2.0),
(make(0), 2.0), (make(2), 2.0),
(make(1), 2.0), ];
let run = || -> (Vec<Assignment>, Vec<[f32; EMBEDDING_DIM]>) {
let mut c = OnlineClusterer::try_new(OnlineOptions::default().to_dia_options())
.expect("default OnlineOptions are valid");
let mut assigns = Vec::new();
for (e, d) in &seq {
assigns.push(c.assign(e, *d));
}
let centroids: Vec<[f32; EMBEDDING_DIM]> =
c.speaker_ids().map(|id| *c.centroid(id).unwrap()).collect();
(assigns, centroids)
};
let (a1, cent1) = run();
let (a2, cent2) = run();
assert_eq!(a1, a2, "assignments must be identical across runs");
assert_eq!(cent1, cent2, "centroids must be bit-identical across runs");
assert_eq!(
a1,
vec![
Assignment::New(1),
Assignment::New(2),
Assignment::Existing(1),
Assignment::New(3),
Assignment::Existing(2),
]
);
}
#[cfg(feature = "serde")]
#[test]
fn serde_online_empty_payload_is_full_defaults() {
let b: ClusterBackend = serde_json::from_str(r#"{"online":{}}"#).unwrap();
assert_eq!(b, ClusterBackend::Online(OnlineOptions::new()));
let o: OnlineOptions = serde_json::from_str("{}").unwrap();
assert_eq!(o, OnlineOptions::new());
}
#[cfg(feature = "serde")]
#[test]
fn serde_online_partial_payload_defaults_other_knobs() {
let b: ClusterBackend = serde_json::from_str(r#"{"online":{"speaker_threshold":0.9}}"#).unwrap();
let ClusterBackend::Online(o) = b else {
panic!("expected Online")
};
assert_eq!(o.speaker_threshold(), 0.9);
assert_eq!(o.embedding_threshold(), DEFAULT_EMBEDDING_THRESHOLD);
assert_eq!(o.min_speech_duration(), DEFAULT_MIN_SPEECH_DURATION);
}
#[cfg(feature = "serde")]
#[test]
fn serde_online_non_default_round_trips_without_silent_flip() {
let b = ClusterBackend::Online(
OnlineOptions::new()
.with_speaker_threshold(0.71)
.with_embedding_threshold(0.33)
.with_min_speech_duration(1.75),
);
let json = serde_json::to_string(&b).unwrap();
let back: ClusterBackend = serde_json::from_str(&json).unwrap();
assert_eq!(back, b);
}
#[cfg(feature = "serde")]
#[test]
fn serde_online_rejects_non_finite_and_out_of_range_thresholds() {
assert!(serde_json::from_str::<OnlineOptions>(r#"{"speaker_threshold":2.5}"#).is_err());
assert!(serde_json::from_str::<OnlineOptions>(r#"{"speaker_threshold":-0.1}"#).is_err());
assert!(serde_json::from_str::<OnlineOptions>(r#"{"embedding_threshold":3.0}"#).is_err());
assert!(
serde_json::from_str::<ClusterBackend>(r#"{"online":{"speaker_threshold":2.5}}"#).is_err()
);
}
#[cfg(feature = "serde")]
#[test]
fn serde_online_serialize_helper_rejects_out_of_range() {
#[derive(serde::Serialize)]
struct Bare {
#[serde(with = "super::finite_threshold_f32")]
v: f32,
}
assert!(serde_json::to_string(&Bare { v: 2.5 }).is_err());
assert!(serde_json::to_string(&Bare { v: f32::NAN }).is_err());
assert!(serde_json::to_string(&Bare { v: 0.65 }).is_ok());
}
#[cfg(feature = "serde")]
#[test]
fn serde_online_rejects_negative_min_speech_duration() {
assert!(serde_json::from_str::<OnlineOptions>(r#"{"min_speech_duration":-1.0}"#).is_err());
assert!(
serde_json::from_str::<ClusterBackend>(r#"{"online":{"min_speech_duration":-0.001}}"#).is_err()
);
let o: OnlineOptions = serde_json::from_str(r#"{"min_speech_duration":0.0}"#).unwrap();
assert_eq!(o.min_speech_duration(), 0.0);
}