use crate::types::{SpeakerId, SpeakerTurn, TimeRange, Word, WordAlignment};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum UncoveredPolicy {
#[default]
None,
LastTurn,
}
pub fn speaker_at(turns: &[SpeakerTurn], t: f64) -> Option<SpeakerId> {
turns
.iter()
.find(|turn| turn.time.contains_instant(t))
.map(|turn| turn.speaker)
}
pub fn speaker_at_stable(turns: &[SpeakerTurn], t: f64) -> Option<SpeakerId> {
turns
.iter()
.find(|turn| turn.stable && turn.time.contains_instant(t))
.map(|turn| turn.speaker)
}
#[inline]
pub fn midpoint(time: &TimeRange) -> f64 {
time.midpoint()
}
fn resolve_at(
turns: &[SpeakerTurn],
t: f64,
stable_only: bool,
policy: UncoveredPolicy,
) -> Option<SpeakerId> {
let hit = if stable_only {
speaker_at_stable(turns, t)
} else {
speaker_at(turns, t)
};
hit.or_else(|| match policy {
UncoveredPolicy::None => None,
UncoveredPolicy::LastTurn => turns.last().map(|turn| turn.speaker),
})
}
pub fn assign_speakers_by_midpoint(
word_times: impl IntoIterator<Item = TimeRange>,
turns: &[SpeakerTurn],
policy: UncoveredPolicy,
stable_only: bool,
) -> Vec<Option<SpeakerId>> {
word_times
.into_iter()
.map(|time| resolve_at(turns, time.midpoint(), stable_only, policy))
.collect()
}
pub fn label_words(
words: &[Word],
turns: &[SpeakerTurn],
policy: UncoveredPolicy,
stable_only: bool,
) -> Vec<WordAlignment> {
words
.iter()
.map(|w| {
let speaker = resolve_at(turns, w.time.midpoint(), stable_only, policy);
WordAlignment {
word: w.word.clone(),
time: w.time,
speaker,
confidence: if speaker.is_some() { 1.0 } else { 0.0 },
interpolated: false,
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn turn(id: u32, start: f64, end: f64) -> SpeakerTurn {
SpeakerTurn::new(SpeakerId(id), TimeRange { start, end })
}
fn turn_unstable(id: u32, start: f64, end: f64) -> SpeakerTurn {
SpeakerTurn::with_stability(SpeakerId(id), TimeRange { start, end }, false)
}
#[test]
fn speaker_at_first_covering_turn_wins() {
let turns = vec![turn(0, 0.0, 2.0), turn(1, 1.5, 3.0)];
assert_eq!(speaker_at(&turns, 1.7), Some(SpeakerId(0)));
assert_eq!(speaker_at(&turns, 2.5), Some(SpeakerId(1)));
assert_eq!(speaker_at(&turns, 4.0), None);
}
#[test]
fn streaming_last_turn_fallback() {
let turns = vec![turn(0, 0.0, 1.0), turn(1, 1.0, 2.0)];
let times = [TimeRange {
start: 2.5,
end: 2.7,
}];
let labels = assign_speakers_by_midpoint(times, &turns, UncoveredPolicy::LastTurn, false);
assert_eq!(labels, vec![Some(SpeakerId(1))]);
}
#[test]
fn last_turn_fallback_ignores_stable_only_for_tail() {
let turns = vec![turn_unstable(7, 0.0, 1.0)];
let times = [TimeRange {
start: 5.0,
end: 5.1,
}];
let labels = assign_speakers_by_midpoint(times, &turns, UncoveredPolicy::LastTurn, true);
assert_eq!(labels, vec![Some(SpeakerId(7))]);
}
#[test]
fn offline_none_leaves_uncovered_unset() {
let turns = vec![turn(0, 0.0, 1.0)];
let times = [TimeRange {
start: 5.0,
end: 5.2,
}];
let labels = assign_speakers_by_midpoint(times, &turns, UncoveredPolicy::None, false);
assert_eq!(labels, vec![None]);
}
#[test]
fn empty_turns_clear_all_with_last_turn_policy() {
let times = [TimeRange {
start: 0.0,
end: 0.1,
}];
let labels = assign_speakers_by_midpoint(times, &[], UncoveredPolicy::LastTurn, false);
assert_eq!(labels, vec![None]);
}
#[test]
fn stable_only_skips_provisional() {
let turns = vec![turn_unstable(0, 0.0, 2.0), turn(1, 1.0, 3.0)];
assert_eq!(speaker_at_stable(&turns, 0.5), None);
assert_eq!(speaker_at_stable(&turns, 1.5), Some(SpeakerId(1)));
}
#[test]
fn label_words_midpoint() {
let turns = vec![turn(0, 0.0, 1.0), turn(1, 1.0, 2.0)];
let words = vec![
Word {
word: "hello".into(),
time: TimeRange {
start: 0.1,
end: 0.3,
},
confidence: 0.9,
},
Word {
word: "there".into(),
time: TimeRange {
start: 1.2,
end: 1.4,
},
confidence: 0.8,
},
];
let out = label_words(&words, &turns, UncoveredPolicy::None, false);
assert_eq!(out[0].speaker, Some(SpeakerId(0)));
assert_eq!(out[1].speaker, Some(SpeakerId(1)));
assert!((out[0].confidence - 1.0).abs() < f32::EPSILON);
}
#[test]
fn midpoint_matches_time_range() {
let tr = TimeRange {
start: 1.0,
end: 3.0,
};
assert!((midpoint(&tr) - 2.0).abs() < f64::EPSILON);
assert!((tr.midpoint() - 2.0).abs() < f64::EPSILON);
}
}