use super::*;
#[test]
fn test_diarization_config_default() {
let config = DiarizationConfig::default();
assert_eq!(config.min_speakers, 1);
assert!(config.max_speakers.is_none());
assert!((config.min_segment_duration - 0.5).abs() < f32::EPSILON);
}
#[test]
fn test_diarization_config_for_realtime() {
let config = DiarizationConfig::for_realtime();
assert_eq!(config.max_speakers, Some(4));
assert!((config.min_segment_duration - 0.3).abs() < f32::EPSILON);
}
#[test]
fn test_diarization_config_for_accuracy() {
let config = DiarizationConfig::for_accuracy();
assert!(config.max_speakers.is_none());
assert!((config.min_segment_duration - 0.5).abs() < f32::EPSILON);
}
#[test]
fn test_diarization_config_with_max_speakers() {
let config = DiarizationConfig::default().with_max_speakers(3);
assert_eq!(config.max_speakers, Some(3));
}
#[test]
fn test_diarization_config_with_min_segment_duration() {
let config = DiarizationConfig::default().with_min_segment_duration(1.0);
assert!((config.min_segment_duration - 1.0).abs() < f32::EPSILON);
}
#[test]
fn test_diarization_result_new() {
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(1, 2.0, 4.0, 0.85),
];
let embeddings = vec![
SpeakerEmbedding::new(vec![0.1; 256], 0),
SpeakerEmbedding::new(vec![0.2; 256], 1),
];
let result = DiarizationResult::new(segments, 2, embeddings, 4.0);
assert_eq!(result.num_speakers(), 2);
assert_eq!(result.segments().len(), 2);
assert!((result.duration() - 4.0).abs() < f32::EPSILON);
}
#[test]
fn test_diarization_result_segments_for_speaker() {
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(1, 2.0, 4.0, 0.85),
SpeakerSegment::new(0, 4.0, 6.0, 0.88),
];
let result = DiarizationResult::new(segments, 2, Vec::new(), 6.0);
let speaker0_segments = result.segments_for_speaker(0);
assert_eq!(speaker0_segments.len(), 2);
let speaker1_segments = result.segments_for_speaker(1);
assert_eq!(speaker1_segments.len(), 1);
}
#[test]
fn test_diarization_result_speaking_time() {
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(1, 2.0, 4.0, 0.85),
SpeakerSegment::new(0, 4.0, 6.0, 0.88),
];
let result = DiarizationResult::new(segments, 2, Vec::new(), 6.0);
assert!((result.speaking_time(0) - 4.0).abs() < f32::EPSILON);
assert!((result.speaking_time(1) - 2.0).abs() < f32::EPSILON);
}
#[test]
fn test_diarization_result_speaker_turns() {
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(1, 2.0, 4.0, 0.85),
SpeakerSegment::new(0, 4.0, 6.0, 0.88),
];
let result = DiarizationResult::new(segments, 2, Vec::new(), 6.0);
let turns = result.speaker_turns();
assert_eq!(turns.len(), 2);
assert_eq!(turns[0].from_speaker(), 0);
assert_eq!(turns[0].to_speaker(), 1);
assert_eq!(turns[1].from_speaker(), 1);
assert_eq!(turns[1].to_speaker(), 0);
}
#[test]
fn test_diarization_result_no_turns_single_speaker() {
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(0, 2.0, 4.0, 0.85),
];
let result = DiarizationResult::new(segments, 1, Vec::new(), 4.0);
let turns = result.speaker_turns();
assert!(turns.is_empty());
}
#[test]
fn test_diarizer_new() {
let diarizer = Diarizer::new(DiarizationConfig::default());
assert_eq!(diarizer.config().min_speakers, 1);
}
#[test]
fn test_diarizer_default_config() {
let diarizer = Diarizer::default_config();
assert!(diarizer.config().max_speakers.is_none());
}
#[test]
fn test_diarizer_process_empty_audio() {
let diarizer = Diarizer::default_config();
let audio: Vec<f32> = vec![];
let result = diarizer.process(&audio, 16000);
assert!(result.is_ok());
let result = result.expect("should succeed");
assert_eq!(result.num_speakers(), 0);
assert!(result.segments().is_empty());
}
#[test]
fn test_diarizer_process_silence() {
let diarizer = Diarizer::default_config();
let audio: Vec<f32> = vec![0.0; 16000]; let result = diarizer.process(&audio, 16000);
assert!(result.is_ok());
let result = result.expect("should succeed");
assert!(result.segments().is_empty());
}
#[test]
fn test_diarizer_merge_segments_same_speaker() {
let diarizer = Diarizer::default_config();
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(0, 2.05, 4.0, 0.85), ];
let merged = diarizer.merge_segments(segments);
assert_eq!(merged.len(), 1);
assert!((merged[0].start() - 0.0).abs() < f32::EPSILON);
assert!((merged[0].end() - 4.0).abs() < f32::EPSILON);
}
#[test]
fn test_diarizer_merge_segments_different_speakers() {
let diarizer = Diarizer::default_config();
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(1, 2.0, 4.0, 0.85),
];
let merged = diarizer.merge_segments(segments);
assert_eq!(merged.len(), 2);
}
#[test]
fn test_diarizer_merge_filters_short_segments() {
let config = DiarizationConfig::default().with_min_segment_duration(1.0);
let diarizer = Diarizer::new(config);
let segments = vec![
SpeakerSegment::new(0, 0.0, 0.3, 0.9), SpeakerSegment::new(1, 0.5, 2.0, 0.85),
];
let merged = diarizer.merge_segments(segments);
assert_eq!(merged.len(), 1);
assert_eq!(merged[0].speaker_id(), 1);
}
#[test]
fn test_assign_speaker_labels_basic() {
let diarizer = Diarizer::default_config();
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(0, 2.0, 4.0, 0.85),
SpeakerSegment::new(0, 4.0, 6.0, 0.88),
];
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 0),
SpeakerEmbedding::new(vec![1.0; 256], 0),
];
let clustering_config = ClusteringConfig::default();
let clustering = SpectralClustering::new(clustering_config);
let cluster_result = clustering.cluster(&embeddings, None, 1).expect("cluster");
let labeled = diarizer
.assign_speaker_labels(&segments, &cluster_result)
.expect("should assign labels");
assert_eq!(labeled.len(), 3);
for seg in &labeled {
assert!(seg.speaker_id() < 10); }
}
#[test]
fn test_assign_speaker_labels_mismatch_error() {
let diarizer = Diarizer::default_config();
let segments = vec![
SpeakerSegment::new(0, 0.0, 2.0, 0.9),
SpeakerSegment::new(0, 2.0, 4.0, 0.85),
];
let embeddings = vec![SpeakerEmbedding::new(vec![1.0; 256], 0)];
let clustering_config = ClusteringConfig::default();
let clustering = SpectralClustering::new(clustering_config);
let cluster_result = clustering.cluster(&embeddings, None, 1).expect("cluster");
let result = diarizer.assign_speaker_labels(&segments, &cluster_result);
assert!(result.is_err());
}
#[test]
fn test_extract_segment_embeddings_basic() {
let diarizer = Diarizer::default_config();
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 2)
.map(|i| (i as f32 * 0.01).sin())
.collect();
let segments = vec![
SpeakerSegment::new(0, 0.0, 1.0, 0.9),
SpeakerSegment::new(0, 1.0, 2.0, 0.85),
];
let embeddings = diarizer
.extract_segment_embeddings(&audio, sample_rate, &segments)
.expect("should extract embeddings");
assert_eq!(embeddings.len(), 2);
}
#[test]
fn test_extract_segment_embeddings_skip_invalid() {
let diarizer = Diarizer::default_config();
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize)
.map(|i| (i as f32 * 0.01).sin())
.collect();
let segments = vec![
SpeakerSegment::new(0, 0.0, 0.5, 0.9),
SpeakerSegment::new(0, 2.0, 3.0, 0.85), ];
let embeddings = diarizer
.extract_segment_embeddings(&audio, sample_rate, &segments)
.expect("should succeed");
assert!(embeddings.len() <= 2);
}
#[test]
fn test_extract_segment_embeddings_empty() {
let diarizer = Diarizer::default_config();
let audio = vec![0.0f32; 16000];
let embeddings = diarizer
.extract_segment_embeddings(&audio, 16000, &[])
.expect("should succeed");
assert!(embeddings.is_empty());
}
#[test]
fn test_diarizer_process_with_synthetic_speech() {
let diarizer = Diarizer::default_config();
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 3)
.map(|i| {
let t = i as f32 / sample_rate as f32;
if t < 1.5 {
(t * 200.0 * std::f32::consts::TAU).sin() * 0.5
} else {
(t * 350.0 * std::f32::consts::TAU).sin() * 0.5
}
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 3.0).abs() < 0.1);
}
#[test]
fn test_diarizer_cluster_speakers_kmeans_config() {
let mut config = DiarizationConfig::default();
config.clustering.algorithm = ClusteringAlgorithm::KMeans;
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 2)
.map(|i| {
let t = i as f32 / sample_rate as f32;
(t * 300.0 * std::f32::consts::TAU).sin() * 0.5
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 2.0).abs() < 0.1);
}
#[test]
fn test_diarizer_cluster_speakers_agglomerative_config() {
let mut config = DiarizationConfig::default();
config.clustering.algorithm = ClusteringAlgorithm::Agglomerative;
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 2)
.map(|i| {
let t = i as f32 / sample_rate as f32;
(t * 250.0 * std::f32::consts::TAU).sin() * 0.5
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 2.0).abs() < 0.1);
}
#[test]
fn test_diarizer_process_with_max_speakers() {
let config = DiarizationConfig::default().with_max_speakers(2);
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 3)
.map(|i| {
let t = i as f32 / sample_rate as f32;
if t < 1.0 {
(t * 200.0 * std::f32::consts::TAU).sin() * 0.5
} else if t < 2.0 {
(t * 400.0 * std::f32::consts::TAU).sin() * 0.5
} else {
(t * 200.0 * std::f32::consts::TAU).sin() * 0.5
}
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!(result.num_speakers() <= 2);
}
#[test]
fn test_diarizer_process_realtime_config() {
let config = DiarizationConfig::for_realtime();
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 2)
.map(|i| {
let t = i as f32 / sample_rate as f32;
(t * 300.0 * std::f32::consts::TAU).sin() * 0.5
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 2.0).abs() < 0.1);
}
#[test]
fn test_diarizer_process_loud_two_speaker_audio() {
let config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 4)
.map(|i| {
let t = i as f32 / sample_rate as f32;
if t < 1.5 {
(t * 150.0 * std::f32::consts::TAU).sin() * 0.8
} else if t < 2.0 {
0.0 } else {
(t * 500.0 * std::f32::consts::TAU).sin() * 0.7
}
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 4.0).abs() < 0.1);
assert!(result.num_speakers() <= 3);
}
#[test]
fn test_diarizer_cluster_speakers_direct() {
let config = DiarizationConfig::default().with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 3)
.map(|i| {
let t = i as f32 / sample_rate as f32;
(t * 300.0 * std::f32::consts::TAU).sin() * 0.9
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 3.0).abs() < 0.1);
}
#[test]
fn test_diarizer_process_accuracy_config() {
let config = DiarizationConfig::for_accuracy();
let diarizer = Diarizer::new(config);
let sample_rate = 16000u32;
let audio: Vec<f32> = (0..sample_rate as usize * 2)
.map(|i| {
let t = i as f32 / sample_rate as f32;
(t * 300.0 * std::f32::consts::TAU).sin() * 0.5
})
.collect();
let result = diarizer
.process(&audio, sample_rate)
.expect("should succeed");
assert!((result.duration() - 2.0).abs() < 0.1);
}
fn generate_speech_with_silence(
sample_rate: u32,
segments: &[(f32, f32, f32)], total_duration: f32,
) -> Vec<f32> {
let total_samples = (total_duration * sample_rate as f32) as usize;
let mut audio = vec![0.0f32; total_samples];
for &(start, end, freq) in segments {
let s = (start * sample_rate as f32) as usize;
let e = ((end * sample_rate as f32) as usize).min(total_samples);
for i in s..e {
let t = i as f32 / sample_rate as f32;
audio[i] = (t * freq * std::f32::consts::TAU).sin() * 0.8;
}
}
audio
}
#[test]
fn test_process_full_pipeline_single_speaker() {
let config = DiarizationConfig::default().with_min_segment_duration(0.2);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 3.0, 300.0)], 4.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!(
result.num_speakers() >= 1,
"expected >=1 speaker, got {}; VAD should detect speech region",
result.num_speakers()
);
assert!(
!result.segments().is_empty(),
"expected non-empty segments from full pipeline"
);
}
#[test]
fn test_process_full_pipeline_two_speakers() {
let config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.2);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 2.5, 200.0), (3.0, 4.5, 500.0)], 5.5);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!(
result.num_speakers() >= 1,
"expected >=1 speaker from two speech regions, got {}",
result.num_speakers()
);
assert!((result.duration() - 5.5).abs() < 0.1);
}
#[test]
fn test_cluster_speakers_kmeans_with_vad_triggering_audio() {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.2);
config.clustering.algorithm = ClusteringAlgorithm::KMeans;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 3.0, 300.0)], 4.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!(
result.num_speakers() >= 1,
"KMeans path: expected >=1 speaker, got {}",
result.num_speakers()
);
}
#[test]
fn test_cluster_speakers_agglomerative_with_vad_triggering_audio() {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.2);
config.clustering.algorithm = ClusteringAlgorithm::Agglomerative;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 3.0, 300.0)], 4.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!(
result.num_speakers() >= 1,
"Agglomerative path: expected >=1 speaker, got {}",
result.num_speakers()
);
}
#[test]
fn test_process_speaker_embeddings_populated() {
let config = DiarizationConfig::default().with_min_segment_duration(0.2);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 3.0, 300.0)], 4.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
if result.num_speakers() > 0 {
assert!(
!result.speaker_embeddings().is_empty(),
"step 6: speaker_embeddings should be populated when speakers detected"
);
}
}
#[test]
fn test_process_merge_adjacent_same_speaker_segments() {
let config = DiarizationConfig::default().with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 4.0, 250.0)], 5.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
if !result.segments().is_empty() {
assert!(
result.segments().len() <= 5,
"continuous speech should merge, got {} segments",
result.segments().len()
);
}
}
#[test]
#[allow(clippy::expect_used)]
fn test_cluster_speakers_spectral_with_forced_segments() {
let config = DiarizationConfig::default().with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(2.0, 5.0, 440.0)], 7.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!((result.duration() - 7.0).abs() < 0.1);
if result.num_speakers() > 0 {
assert!(!result.segments().is_empty());
assert!(!result.speaker_embeddings().is_empty());
}
}
#[test]
#[allow(clippy::expect_used)]
fn test_process_exercises_all_steps_with_two_speech_bursts() {
let config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.0, 3.0, 200.0), (4.0, 6.0, 600.0)], 7.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!((result.duration() - 7.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(
!result.segments().is_empty(),
"with speakers detected, segments should not be empty"
);
}
}
#[test]
#[allow(clippy::expect_used)]
fn test_process_spectral_algorithm_exercises_cluster_speakers() {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.1);
config.clustering.algorithm = ClusteringAlgorithm::Spectral;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.5, 3.5, 350.0)], 5.0);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!((result.duration() - 5.0).abs() < 0.1);
}
#[test]
#[allow(clippy::expect_used)]
fn test_process_long_audio_multiple_speakers() {
let config = DiarizationConfig::default()
.with_max_speakers(4)
.with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(
sr,
&[(1.0, 3.0, 150.0), (4.0, 6.0, 400.0), (7.0, 9.0, 250.0)],
10.0,
);
let result = diarizer.process(&audio, sr).expect("should succeed");
assert!((result.duration() - 10.0).abs() < 0.1);
let turns = result.speaker_turns();
let _ = turns.len();
}
#[test]
fn test_process_impulse_audio_exercises_full_pipeline() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let mut audio = vec![0.0f32; sr as usize * 6];
for i in (sr as usize * 2)..(sr as usize * 4) {
let t = i as f32 / sr as f32;
audio[i] = (t * 220.0 * std::f32::consts::TAU).sin() * 0.7
+ (t * 440.0 * std::f32::consts::TAU).sin() * 0.3;
}
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 6.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(!result.segments().is_empty());
assert!(!result.speaker_embeddings().is_empty());
for emb in result.speaker_embeddings() {
assert_eq!(emb.dim(), 256);
}
}
Ok(())
}
#[test]
fn test_process_two_bursts_forces_multi_segment_clustering() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(2)
.with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(1.5, 3.5, 150.0), (5.5, 7.5, 600.0)], 9.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 9.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(!result.segments().is_empty());
for seg in result.segments() {
assert!(seg.start() >= 0.0);
assert!(seg.end() > seg.start());
assert!(seg.end() <= 9.5); }
}
Ok(())
}
#[test]
fn test_cluster_speakers_kmeans_via_process_with_reliable_vad() -> WhisperResult<()> {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.1);
config.clustering.algorithm = ClusteringAlgorithm::KMeans;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(2.0, 5.0, 300.0)], 7.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 7.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(!result.segments().is_empty());
}
Ok(())
}
#[test]
fn test_cluster_speakers_agglomerative_via_process_with_reliable_vad() -> WhisperResult<()> {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.1);
config.clustering.algorithm = ClusteringAlgorithm::Agglomerative;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(2.0, 5.0, 350.0)], 7.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 7.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(!result.segments().is_empty());
}
Ok(())
}
#[test]
fn test_process_step5_merge_produces_fewer_segments() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_min_segment_duration(0.05)
.with_max_speakers(2);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(2.0, 6.0, 280.0)], 8.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 8.0).abs() < 0.1);
if result.num_speakers() >= 1 {
let total_segments = result.segments().len();
assert!(
total_segments <= 10,
"continuous 4s speech should merge into few segments, got {}",
total_segments
);
}
Ok(())
}
#[test]
fn test_process_step6_centroids_match_num_speakers() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_min_segment_duration(0.1)
.with_max_speakers(4);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(
sr,
&[(1.0, 2.5, 180.0), (3.5, 5.0, 450.0), (6.0, 7.5, 320.0)],
9.0,
);
let result = diarizer.process(&audio, sr)?;
if result.num_speakers() > 0 {
assert_eq!(
result.speaker_embeddings().len(),
result.num_speakers(),
"centroids count must equal num_speakers"
);
}
Ok(())
}
#[test]
fn test_process_short_audio_with_high_silence_ratio() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(0.5, 1.5, 400.0)], 2.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 2.0).abs() < 0.1);
Ok(())
}
#[test]
fn test_process_min_duration_filtering_in_merge() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(1.0);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_speech_with_silence(sr, &[(2.0, 2.5, 300.0)], 5.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 5.0).abs() < 0.1);
Ok(())
}
#[test]
fn test_diarization_result_speaker_turns_empty_segments() {
let result = DiarizationResult::new(Vec::new(), 0, Vec::new(), 0.0);
let turns = result.speaker_turns();
assert!(turns.is_empty());
}
#[test]
fn test_diarization_result_speaker_turns_single_segment() {
let segments = vec![SpeakerSegment::new(0, 0.0, 2.0, 0.9)];
let result = DiarizationResult::new(segments, 1, Vec::new(), 2.0);
let turns = result.speaker_turns();
assert!(turns.is_empty());
}
#[test]
fn test_diarization_result_speaking_time_nonexistent_speaker() {
let segments = vec![SpeakerSegment::new(0, 0.0, 2.0, 0.9)];
let result = DiarizationResult::new(segments, 1, Vec::new(), 2.0);
assert!((result.speaking_time(99) - 0.0).abs() < f32::EPSILON);
}
#[test]
fn test_process_with_non_standard_sample_rate() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 44100u32;
let total_samples = (6.0 * sr as f32) as usize;
let speech_start = (2.0 * sr as f32) as usize;
let speech_end = (4.0 * sr as f32) as usize;
let mut audio = vec![0.0f32; total_samples];
for i in speech_start..speech_end {
let t = i as f32 / sr as f32;
audio[i] = (t * 300.0 * std::f32::consts::TAU).sin() * 0.8;
}
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 6.0).abs() < 0.2);
Ok(())
}
fn generate_vad_triggering_audio(
sample_rate: u32,
speech_regions: &[(f32, f32, f32)], total_duration: f32,
) -> Vec<f32> {
let total_samples = (total_duration * sample_rate as f32) as usize;
let mut audio = vec![0.0f32; total_samples];
for &(start, end, freq) in speech_regions {
let s = (start * sample_rate as f32) as usize;
let e = ((end * sample_rate as f32) as usize).min(total_samples);
for i in s..e {
let t = i as f32 / sample_rate as f32;
let am = 0.3f32.mul_add((t * 5.0 * std::f32::consts::TAU).sin(), 0.7);
let signal = (t * freq * std::f32::consts::TAU).sin()
+ 0.5 * (t * freq * 2.0 * std::f32::consts::TAU).sin()
+ 0.25 * (t * freq * 3.0 * std::f32::consts::TAU).sin();
audio[i] = signal * am * 0.6;
}
}
audio
}
#[test]
fn test_process_guaranteed_vad_single_burst() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(sr, &[(3.0, 7.0, 300.0)], 10.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 10.0).abs() < 0.1);
assert!(
result.num_speakers() >= 1,
"VAD must detect speech: num_speakers={}, segments={}",
result.num_speakers(),
result.segments().len()
);
assert!(
!result.segments().is_empty(),
"pipeline steps 2-6 must produce segments"
);
assert!(
!result.speaker_embeddings().is_empty(),
"step 6 must produce speaker embeddings"
);
Ok(())
}
#[test]
fn test_process_guaranteed_vad_two_bursts_cluster_speakers() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(sr, &[(2.0, 5.0, 200.0), (7.0, 10.0, 700.0)], 12.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 12.0).abs() < 0.1);
assert!(
result.num_speakers() >= 1,
"two speech bursts must yield speakers, got num_speakers={}",
result.num_speakers()
);
assert!(
!result.segments().is_empty(),
"must have segments from two speech regions"
);
assert_eq!(
result.speaker_embeddings().len(),
result.num_speakers(),
"centroids count must equal num_speakers"
);
Ok(())
}
#[test]
fn test_process_cluster_speakers_kmeans_branch() -> WhisperResult<()> {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.05);
config.clustering.algorithm = ClusteringAlgorithm::KMeans;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(sr, &[(2.0, 6.0, 350.0)], 8.0);
let result = diarizer.process(&audio, sr)?;
assert!(
result.num_speakers() >= 1,
"KMeans branch: VAD must detect speech"
);
assert!(!result.segments().is_empty());
Ok(())
}
#[test]
fn test_process_cluster_speakers_agglomerative_branch() -> WhisperResult<()> {
let mut config = DiarizationConfig::default().with_min_segment_duration(0.05);
config.clustering.algorithm = ClusteringAlgorithm::Agglomerative;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(sr, &[(2.0, 6.0, 450.0)], 8.0);
let result = diarizer.process(&audio, sr)?;
assert!(
result.num_speakers() >= 1,
"Agglomerative branch: VAD must detect speech"
);
assert!(!result.segments().is_empty());
Ok(())
}
#[test]
fn test_process_merge_step_with_guaranteed_vad() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(sr, &[(3.0, 9.0, 250.0)], 12.0);
let result = diarizer.process(&audio, sr)?;
assert!(
result.num_speakers() >= 1,
"continuous speech must be detected"
);
assert!(
result.segments().len() <= 8,
"6s continuous speech should merge, got {} segments",
result.segments().len()
);
Ok(())
}
#[test]
fn test_process_assign_labels_valid_speaker_ids() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(4)
.with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(
sr,
&[(1.0, 3.0, 200.0), (5.0, 7.0, 500.0), (9.0, 11.0, 350.0)],
13.0,
);
let result = diarizer.process(&audio, sr)?;
for seg in result.segments() {
assert!(
seg.speaker_id() < result.num_speakers(),
"speaker_id {} must be < num_speakers {}",
seg.speaker_id(),
result.num_speakers()
);
}
Ok(())
}
#[test]
fn test_process_three_bursts_speaker_turns() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(4)
.with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(
sr,
&[(1.0, 3.0, 150.0), (4.0, 6.0, 600.0), (7.0, 9.0, 150.0)],
10.0,
);
let result = diarizer.process(&audio, sr)?;
assert!(
result.num_speakers() >= 1,
"three bursts must detect speakers"
);
let turns = result.speaker_turns();
if result.num_speakers() >= 2 {
assert!(
!turns.is_empty(),
"with >=2 speakers, there should be speaker turns"
);
}
Ok(())
}
#[test]
fn test_process_embedding_extraction_dimensions() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(sr, &[(2.0, 6.0, 400.0)], 8.0);
let result = diarizer.process(&audio, sr)?;
for emb in result.speaker_embeddings() {
assert_eq!(emb.dim(), 256, "speaker embedding must be 256-dimensional");
}
Ok(())
}
#[test]
fn test_process_merge_filters_short_segments_via_pipeline() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(2.0); let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(
sr,
&[(1.0, 1.5, 300.0), (3.0, 3.5, 400.0), (5.0, 8.0, 300.0)],
10.0,
);
let result = diarizer.process(&audio, sr)?;
for seg in result.segments() {
assert!(
seg.duration() >= 1.9, "segment duration {:.2}s should be >= 2.0s after merge filtering",
seg.duration()
);
}
Ok(())
}
#[test]
fn test_process_cluster_speakers_spectral_multi_segment() -> WhisperResult<()> {
let mut config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.05);
config.clustering.algorithm = ClusteringAlgorithm::Spectral;
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_vad_triggering_audio(
sr,
&[(1.0, 3.0, 180.0), (4.0, 6.0, 500.0), (7.0, 9.0, 320.0)],
10.0,
);
let result = diarizer.process(&audio, sr)?;
assert!(
result.num_speakers() >= 1,
"spectral clustering must produce speakers"
);
assert!(!result.speaker_embeddings().is_empty());
Ok(())
}
#[test]
fn test_cluster_speakers_direct_spectral() -> WhisperResult<()> {
let config = DiarizationConfig::default();
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![0.95; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 1),
SpeakerEmbedding::new(vec![-0.95; 256], 1),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 4);
Ok(())
}
#[test]
fn test_cluster_speakers_direct_kmeans() -> WhisperResult<()> {
let mut config = DiarizationConfig::default();
config.clustering.algorithm = ClusteringAlgorithm::KMeans;
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![0.9; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 1),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 3);
Ok(())
}
#[test]
fn test_cluster_speakers_direct_agglomerative() -> WhisperResult<()> {
let mut config = DiarizationConfig::default();
config.clustering.algorithm = ClusteringAlgorithm::Agglomerative;
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 1),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 2);
Ok(())
}
#[test]
fn test_cluster_speakers_single_embedding() -> WhisperResult<()> {
let diarizer = Diarizer::default_config();
let embeddings = vec![SpeakerEmbedding::new(vec![0.5; 256], 0)];
let result = diarizer.cluster_speakers(&embeddings)?;
assert_eq!(result.num_clusters(), 1);
assert_eq!(result.labels(), &[0]);
Ok(())
}
#[test]
fn test_cluster_speakers_empty_embeddings() -> WhisperResult<()> {
let diarizer = Diarizer::default_config();
let embeddings: Vec<SpeakerEmbedding> = Vec::new();
let result = diarizer.cluster_speakers(&embeddings)?;
assert_eq!(result.num_clusters(), 0);
assert!(result.labels().is_empty());
Ok(())
}
#[test]
fn test_cluster_speakers_with_max_speakers_constraint() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_max_speakers(2);
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(
vec![1.0, 0.0, 0.0]
.into_iter()
.chain(vec![0.0; 253])
.collect(),
0,
),
SpeakerEmbedding::new(
vec![0.0, 1.0, 0.0]
.into_iter()
.chain(vec![0.0; 253])
.collect(),
1,
),
SpeakerEmbedding::new(
vec![0.0, 0.0, 1.0]
.into_iter()
.chain(vec![0.0; 253])
.collect(),
2,
),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(
result.num_clusters() <= 2,
"should respect max_speakers=2, got {}",
result.num_clusters()
);
Ok(())
}
fn generate_broadband_speech(
sample_rate: u32,
speech_regions: &[(f32, f32)],
total_duration: f32,
) -> Vec<f32> {
let total_samples = (total_duration * sample_rate as f32) as usize;
let mut audio = vec![0.0f32; total_samples];
let freqs = [
100.0, 200.0, 300.0, 440.0, 600.0, 800.0, 1000.0, 1500.0, 2000.0,
];
for &(start, end) in speech_regions {
let s = (start * sample_rate as f32) as usize;
let e = ((end * sample_rate as f32) as usize).min(total_samples);
for i in s..e {
let t = i as f32 / sample_rate as f32;
let mut val = 0.0f32;
for (fi, &freq) in freqs.iter().enumerate() {
let phase = t * freq * std::f32::consts::TAU;
val += phase.sin() * (1.0 / (fi as f32 + 1.0));
}
let am = 0.4f32.mul_add((t * 3.0 * std::f32::consts::TAU).sin(), 0.6);
audio[i] = val * am * 0.3;
}
}
audio
}
#[test]
fn test_process_full_pipeline_direct_broadband() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_broadband_speech(sr, &[(2.0, 5.0)], 7.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 7.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(
!result.segments().is_empty(),
"step 4-5: segments must be populated"
);
assert!(
!result.speaker_embeddings().is_empty(),
"step 6: centroids must be populated"
);
assert_eq!(
result.speaker_embeddings().len(),
result.num_speakers(),
"centroids must match num_speakers"
);
}
Ok(())
}
#[test]
fn test_process_two_speakers_broadband() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_broadband_speech(sr, &[(1.0, 3.0), (4.5, 6.5)], 8.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 8.0).abs() < 0.1);
if result.num_speakers() >= 1 {
for seg in result.segments() {
assert!(seg.speaker_id() < result.num_speakers());
assert!(seg.end() > seg.start());
}
}
Ok(())
}
#[test]
fn test_process_deep_pipeline_broadband_comprehensive() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(4)
.with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_broadband_speech(sr, &[(3.0, 7.0)], 10.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 10.0).abs() < 0.1);
if result.num_speakers() >= 1 {
for seg in result.segments() {
assert!(
seg.speaker_id() < result.num_speakers(),
"speaker_id {} >= num_speakers {}",
seg.speaker_id(),
result.num_speakers()
);
assert!(seg.end() > seg.start());
}
assert!(
result.segments().len() <= 15,
"4s speech should merge to <15 segments, got {}",
result.segments().len()
);
assert_eq!(result.speaker_embeddings().len(), result.num_speakers());
for emb in result.speaker_embeddings() {
assert_eq!(emb.dim(), 256);
}
}
Ok(())
}
#[test]
fn test_process_two_broadband_bursts_exercises_clustering() -> WhisperResult<()> {
let config = DiarizationConfig::default()
.with_max_speakers(3)
.with_min_segment_duration(0.05);
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = generate_broadband_speech(sr, &[(1.5, 3.5), (5.0, 7.0)], 9.0);
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 9.0).abs() < 0.1);
if result.num_speakers() >= 1 {
assert!(!result.segments().is_empty());
assert!(!result.speaker_embeddings().is_empty());
for speaker_id in 0..result.num_speakers() {
assert!(result.speaking_time(speaker_id) >= 0.0);
}
}
Ok(())
}
#[test]
fn test_process_pure_silence_hits_early_return() -> WhisperResult<()> {
let config = DiarizationConfig::default();
let diarizer = Diarizer::new(config);
let sr = 16000u32;
let audio = vec![0.0f32; sr as usize * 5];
let result = diarizer.process(&audio, sr)?;
assert_eq!(result.num_speakers(), 0);
assert!(result.segments().is_empty());
assert!(result.speaker_embeddings().is_empty());
assert!((result.duration() - 5.0).abs() < 0.1);
Ok(())
}
#[test]
fn test_process_at_8khz_sample_rate() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_min_segment_duration(0.1);
let diarizer = Diarizer::new(config);
let sr = 8000u32;
let total_samples = (6.0 * sr as f32) as usize;
let speech_start = (2.0 * sr as f32) as usize;
let speech_end = (4.0 * sr as f32) as usize;
let mut audio = vec![0.0f32; total_samples];
let freqs = [100.0, 200.0, 300.0, 440.0, 600.0, 800.0];
for i in speech_start..speech_end {
let t = i as f32 / sr as f32;
let mut val = 0.0f32;
for (fi, &freq) in freqs.iter().enumerate() {
val += (t * freq * std::f32::consts::TAU).sin() * (1.0 / (fi as f32 + 1.0));
}
let am = 0.4f32.mul_add((t * 3.0 * std::f32::consts::TAU).sin(), 0.6);
audio[i] = val * am * 0.3;
}
let result = diarizer.process(&audio, sr)?;
assert!((result.duration() - 6.0).abs() < 0.5);
Ok(())
}
#[test]
fn test_cluster_speakers_five_embeddings_two_clusters() -> WhisperResult<()> {
let config = DiarizationConfig::default().with_max_speakers(3);
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![0.9; 256], 0),
SpeakerEmbedding::new(vec![0.95; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 1),
SpeakerEmbedding::new(vec![-0.9; 256], 1),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 5);
for &label in result.labels() {
assert!(label < result.num_clusters());
}
let centroids = result.cluster_centroids();
assert_eq!(centroids.len(), result.num_clusters());
Ok(())
}
#[test]
fn test_cluster_speakers_spectral_explicit() -> WhisperResult<()> {
let mut config = DiarizationConfig::default();
config.clustering.algorithm = ClusteringAlgorithm::Spectral;
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 1),
SpeakerEmbedding::new(vec![0.5; 256], 0),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 3);
Ok(())
}
#[test]
fn test_cluster_speakers_kmeans_min_speakers() -> WhisperResult<()> {
let mut config = DiarizationConfig::default().with_max_speakers(4);
config.clustering.algorithm = ClusteringAlgorithm::KMeans;
config.min_speakers = 2;
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(vec![1.0; 256], 0),
SpeakerEmbedding::new(vec![0.8; 256], 0),
SpeakerEmbedding::new(vec![-1.0; 256], 1),
SpeakerEmbedding::new(vec![-0.8; 256], 1),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 4);
Ok(())
}
#[test]
fn test_cluster_speakers_agglomerative_three_groups() -> WhisperResult<()> {
let mut config = DiarizationConfig::default().with_max_speakers(4);
config.clustering.algorithm = ClusteringAlgorithm::Agglomerative;
let diarizer = Diarizer::new(config);
let embeddings = vec![
SpeakerEmbedding::new(
vec![1.0, 0.0, 0.0]
.into_iter()
.chain(vec![0.0; 253])
.collect(),
0,
),
SpeakerEmbedding::new(
vec![0.0, 1.0, 0.0]
.into_iter()
.chain(vec![0.0; 253])
.collect(),
1,
),
SpeakerEmbedding::new(
vec![0.0, 0.0, 1.0]
.into_iter()
.chain(vec![0.0; 253])
.collect(),
2,
),
];
let result = diarizer.cluster_speakers(&embeddings)?;
assert!(result.num_clusters() >= 1);
assert_eq!(result.labels().len(), 3);
Ok(())
}