use super::stability::prefer_current_speaker;
use crate::types::SpeakerId;
use crate::utils::{cosine_similarity, l2_normalize};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct AssignResult {
pub speaker: SpeakerId,
pub confidence: f32,
pub stable: bool,
pub overflow_merged: bool,
}
struct CacheEntry {
id: SpeakerId,
centroid: Vec<f32>,
count: usize,
confidence_sum: f32,
last_step: u64,
hits: usize,
stable: bool,
}
impl CacheEntry {
fn keep_score(&self, now: u64) -> f32 {
let avg = if self.count == 0 {
0.0
} else {
self.confidence_sum / self.count as f32
};
let age = now.saturating_sub(self.last_step) as f32;
let recency = 1.0 / (1.0 + age);
let stable_bonus = if self.stable { 0.05 } else { 0.0 };
avg + 0.25 * recency + stable_bonus
}
}
pub struct ArrivalOrderSpeakerCache {
entries: Vec<CacheEntry>,
cap: usize,
match_threshold: f32,
min_hits_to_stable: usize,
prefer_current_margin: f32,
step: u64,
current_speaker: Option<SpeakerId>,
next_id: u32,
}
impl ArrivalOrderSpeakerCache {
#[allow(clippy::panic)] pub fn new(
cap: usize,
match_threshold: f32,
min_hits_to_stable: usize,
prefer_current_margin: f32,
) -> Self {
if cap == 0 {
panic!("ArrivalOrderSpeakerCache::new: cap must be > 0");
}
Self {
entries: Vec::with_capacity(cap.min(16)),
cap,
match_threshold,
min_hits_to_stable: min_hits_to_stable.max(1),
prefer_current_margin,
step: 0,
current_speaker: None,
next_id: 0,
}
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn cap(&self) -> usize {
self.cap
}
pub fn speaker_ids(&self) -> Vec<SpeakerId> {
let mut ids: Vec<_> = self.entries.iter().map(|e| e.id).collect();
ids.sort_by_key(|s| s.0);
ids
}
pub fn assign(&mut self, embedding: &[f32]) -> AssignResult {
self.step = self.step.saturating_add(1);
let now = self.step;
let mut candidates: Vec<(SpeakerId, f32)> = self
.entries
.iter()
.map(|e| (e.id, cosine_similarity(embedding, &e.centroid)))
.collect();
let preferred = prefer_current_speaker(
self.current_speaker,
&candidates,
self.prefer_current_margin,
);
if let Some(pref) = preferred
&& let Some((idx, _)) = candidates
.iter()
.enumerate()
.find(|(_, (id, _))| *id == pref)
{
candidates.swap(0, idx);
}
let best = candidates
.iter()
.copied()
.max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
let chosen = match (preferred, best) {
(Some(pref), Some((best_id, best_sim))) if pref != best_id => {
let pref_sim = candidates
.iter()
.find(|(id, _)| *id == pref)
.map(|(_, s)| *s)
.unwrap_or(f32::NEG_INFINITY);
if best_sim - pref_sim <= self.prefer_current_margin {
Some((pref, pref_sim))
} else {
Some((best_id, best_sim))
}
}
(_, b) => b,
};
if let Some((id, sim)) = chosen
&& (sim >= self.match_threshold || self.entries.len() >= self.cap)
{
let overflow = sim < self.match_threshold && self.entries.len() >= self.cap;
let result = self.update_entry(id, embedding, sim, now);
self.current_speaker = Some(result.speaker);
return AssignResult {
overflow_merged: overflow,
..result
};
}
if self.entries.len() < self.cap {
let result = self.create_entry(embedding, now);
self.current_speaker = Some(result.speaker);
return result;
}
if let Some((id, sim)) = best {
let result = self.update_entry(id, embedding, sim, now);
self.current_speaker = Some(result.speaker);
return AssignResult {
overflow_merged: true,
..result
};
}
let result = self.create_entry(embedding, now);
self.current_speaker = Some(result.speaker);
result
}
pub fn evict_weakest(&mut self) -> Option<SpeakerId> {
if self.entries.len() < self.cap {
return None;
}
let now = self.step;
let idx = self
.entries
.iter()
.enumerate()
.min_by(|(_, a), (_, b)| {
a.keep_score(now)
.partial_cmp(&b.keep_score(now))
.unwrap_or(std::cmp::Ordering::Equal)
})
.map(|(i, _)| i)?;
let evicted = self.entries.swap_remove(idx).id;
if self.current_speaker == Some(evicted) {
self.current_speaker = None;
}
Some(evicted)
}
fn create_entry(&mut self, embedding: &[f32], now: u64) -> AssignResult {
let id = SpeakerId(self.next_id);
self.next_id = self.next_id.saturating_add(1);
let mut centroid = embedding.to_vec();
l2_normalize(&mut centroid);
let hits = 1;
let stable = hits >= self.min_hits_to_stable;
self.entries.push(CacheEntry {
id,
centroid,
count: 1,
confidence_sum: 1.0,
last_step: now,
hits,
stable,
});
debug_assert!(self.entries.len() <= self.cap);
AssignResult {
speaker: id,
confidence: 1.0,
stable,
overflow_merged: false,
}
}
fn update_entry(
&mut self,
id: SpeakerId,
embedding: &[f32],
sim: f32,
now: u64,
) -> AssignResult {
let min_hits = self.min_hits_to_stable;
let Some(idx) = self.entries.iter().position(|e| e.id == id) else {
if let Some((fallback_id, sim2)) = self
.entries
.iter()
.map(|e| (e.id, cosine_similarity(embedding, &e.centroid)))
.max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
{
return self.update_entry(fallback_id, embedding, sim2, now);
}
if self.entries.len() < self.cap {
return self.create_entry(embedding, now);
}
return AssignResult {
speaker: id,
confidence: sim,
stable: false,
overflow_merged: true,
};
};
let entry = &mut self.entries[idx];
let n = entry.count as f32;
for (v, &e) in entry.centroid.iter_mut().zip(embedding.iter()) {
*v = (*v * n + e) / (n + 1.0);
}
l2_normalize(&mut entry.centroid);
entry.count += 1;
entry.confidence_sum += sim;
entry.last_step = now;
entry.hits += 1;
if entry.hits >= min_hits {
entry.stable = true;
}
AssignResult {
speaker: entry.id,
confidence: sim,
stable: entry.stable,
overflow_merged: false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn unit(dim: usize, axis: usize) -> Vec<f32> {
let mut v = vec![0.0f32; dim];
v[axis] = 1.0;
v
}
fn near(dim: usize, axis: usize, eps: f32) -> Vec<f32> {
let mut v = unit(dim, axis);
v[(axis + 1) % dim] = eps;
l2_normalize(&mut v);
v
}
#[test]
fn arrival_order_ids_are_stable_across_chunks() {
let mut cache = ArrivalOrderSpeakerCache::new(8, 0.5, 2, 0.05);
let a = unit(8, 0);
let b = unit(8, 1);
let r0 = cache.assign(&a);
let r1 = cache.assign(&b);
assert_eq!(r0.speaker, SpeakerId(0));
assert_eq!(r1.speaker, SpeakerId(1));
let r2 = cache.assign(&b);
let r3 = cache.assign(&a);
assert_eq!(r2.speaker, SpeakerId(1), "B must stay speaker 1");
assert_eq!(r3.speaker, SpeakerId(0), "A must stay speaker 0");
}
#[test]
fn cache_size_never_exceeds_cap() {
let cap = 3;
let mut cache = ArrivalOrderSpeakerCache::new(cap, 0.99, 2, 0.0);
for axis in 0..10 {
let emb = unit(16, axis % 16);
cache.assign(&emb);
assert!(
cache.len() <= cap,
"cache len {} exceeded cap {}",
cache.len(),
cap
);
}
assert_eq!(cache.len(), cap);
}
#[test]
fn overflow_merges_into_closest() {
let mut cache = ArrivalOrderSpeakerCache::new(2, 0.99, 3, 0.0);
let a = unit(4, 0);
let b = unit(4, 1);
cache.assign(&a);
cache.assign(&b);
let near_a = near(4, 0, 0.01);
let r = cache.assign(&near_a);
assert!(r.overflow_merged || r.speaker.0 < 2);
assert_eq!(cache.len(), 2);
assert!(r.speaker.0 < 2);
}
#[test]
fn provisional_then_stable_is_deterministic() {
let mut cache = ArrivalOrderSpeakerCache::new(4, 0.5, 3, 0.0);
let a = unit(8, 0);
let r1 = cache.assign(&a);
assert!(!r1.stable, "first hit is provisional");
let r2 = cache.assign(&near(8, 0, 0.01));
assert!(!r2.stable, "second hit still provisional at min_hits=3");
let r3 = cache.assign(&near(8, 0, 0.01));
assert!(r3.stable, "third hit reaches stability");
let r4 = cache.assign(&near(8, 0, 0.01));
assert!(r4.stable, "stability latches");
assert_eq!(r1.speaker, r4.speaker);
}
#[test]
fn hysteresis_suppresses_flicker_inside_cache() {
let mut cache = ArrivalOrderSpeakerCache::new(4, 0.4, 5, 0.15);
let a = unit(8, 0);
let b = unit(8, 1);
cache.assign(&a); cache.assign(&b); let r = cache.assign(&a);
assert_eq!(r.speaker, SpeakerId(0));
let mut e = vec![0.0f32; 8];
e[0] = 0.70;
e[1] = 0.78;
l2_normalize(&mut e);
let r2 = cache.assign(&e);
let sim_a = e[0]; let sim_b = e[1];
let gap = cosine_similarity(&e, &b) - cosine_similarity(&e, &a);
if gap <= 0.15 {
assert_eq!(r2.speaker, SpeakerId(0), "hysteresis should keep current A");
} else {
assert_eq!(r2.speaker, SpeakerId(1));
}
let _ = (sim_a, sim_b);
}
#[test]
#[should_panic(expected = "cap must be > 0")]
fn rejects_zero_cap() {
let _ = ArrivalOrderSpeakerCache::new(0, 0.5, 2, 0.0);
}
}