use std::collections::HashSet;
use std::collections::VecDeque;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ConvergenceAction {
#[default]
Stop,
Warn,
SwitchPhase,
AskUser,
Compact,
}
#[derive(Debug, Clone, PartialEq, thiserror::Error)]
pub enum ConvergenceConfigError {
#[error("window_size must be at least 2, got {actual}")]
WindowTooSmall {
actual: usize,
},
#[error("similarity_threshold must be in [0.0, 1.0], got {actual}")]
ThresholdOutOfRange {
actual: f32,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConvergenceConfig {
pub enabled: bool,
pub window_size: usize,
pub similarity_threshold: f32,
#[serde(default)]
pub on_converge: ConvergenceAction,
}
impl Default for ConvergenceConfig {
fn default() -> Self {
Self {
enabled: true,
window_size: 3,
similarity_threshold: 0.95,
on_converge: ConvergenceAction::Stop,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct ConvergenceStatus {
pub detected: bool,
pub consecutive_count: usize,
pub similarity_score: f32,
pub similar_responses: Vec<String>,
pub action: ConvergenceAction,
}
impl ConvergenceStatus {
#[must_use]
pub fn no_convergence() -> Self {
Self {
detected: false,
consecutive_count: 0,
similarity_score: 0.0,
similar_responses: Vec::new(),
action: ConvergenceAction::Stop,
}
}
}
#[derive(Debug)]
pub struct ConvergenceDetector {
config: ConvergenceConfig,
pub(super) window: VecDeque<String>,
pub(super) consecutive_count: usize,
similar_responses: Vec<String>,
}
impl ConvergenceDetector {
pub fn new(config: ConvergenceConfig) -> Result<Self, ConvergenceConfigError> {
if config.window_size < 2 {
return Err(ConvergenceConfigError::WindowTooSmall {
actual: config.window_size,
});
}
if !(0.0..=1.0).contains(&config.similarity_threshold) {
return Err(ConvergenceConfigError::ThresholdOutOfRange {
actual: config.similarity_threshold,
});
}
let capacity = config.window_size;
Ok(Self {
config,
window: VecDeque::with_capacity(capacity),
consecutive_count: 0,
similar_responses: Vec::new(),
})
}
pub fn default_detector() -> Result<Self, ConvergenceConfigError> {
Self::new(ConvergenceConfig::default())
}
pub fn add_response(&mut self, response: &str) -> ConvergenceStatus {
if !self.config.enabled || response.is_empty() {
return ConvergenceStatus::no_convergence();
}
let mut max_similarity = 0.0;
for prev_response in &self.window {
let similarity = Self::compute_similarity(response, prev_response);
if similarity > max_similarity {
max_similarity = similarity;
}
}
let prev_is_similar = match self.window.back() {
Some(prev) => {
let sim = Self::compute_similarity(response, prev);
if sim >= self.config.similarity_threshold {
if !self.similar_responses.contains(&response.to_string()) {
self.similar_responses.push(response.to_string());
}
true
} else {
false
}
}
None => false,
};
if self.window.is_empty() {
self.consecutive_count = 1;
self.similar_responses.push(response.to_string());
} else if prev_is_similar {
self.consecutive_count = self.consecutive_count.saturating_add(1);
} else {
self.consecutive_count = 1;
self.similar_responses.clear();
self.similar_responses.push(response.to_string());
}
if self.window.len() >= self.config.window_size {
self.window.pop_front();
}
self.window.push_back(response.to_string());
let detected = self.consecutive_count >= self.config.window_size;
ConvergenceStatus {
detected,
consecutive_count: self.consecutive_count,
similarity_score: max_similarity,
similar_responses: self.similar_responses.clone(),
action: self.config.on_converge,
}
}
#[must_use]
pub fn check_convergence(&self) -> ConvergenceStatus {
if !self.config.enabled {
return ConvergenceStatus::no_convergence();
}
if self.window.len() < self.config.window_size {
return ConvergenceStatus::no_convergence();
}
ConvergenceStatus {
detected: self.consecutive_count >= self.config.window_size,
consecutive_count: self.consecutive_count,
similarity_score: 0.0,
similar_responses: self.similar_responses.clone(),
action: self.config.on_converge,
}
}
pub fn clear(&mut self) {
self.window.clear();
self.consecutive_count = 0;
self.similar_responses.clear();
}
#[must_use]
pub fn window(&self) -> &VecDeque<String> {
&self.window
}
#[must_use]
pub fn config(&self) -> &ConvergenceConfig {
&self.config
}
#[must_use]
pub fn consecutive_count(&self) -> usize {
self.consecutive_count
}
#[allow(clippy::cast_precision_loss)]
#[must_use]
pub fn compute_similarity(a: &str, b: &str) -> f32 {
if a.is_empty() || b.is_empty() {
return 0.0;
}
let a_norm = Self::normalize_text(a);
let b_norm = Self::normalize_text(b);
let a_words: HashSet<&str> = a_norm.split_whitespace().collect();
let b_words: HashSet<&str> = b_norm.split_whitespace().collect();
if a_words.is_empty() || b_words.is_empty() {
return 0.0;
}
let intersection = a_words.intersection(&b_words).count();
let union = a_words.union(&b_words).count();
if union == 0 {
return 0.0;
}
intersection as f32 / union as f32
}
fn normalize_text(text: &str) -> String {
text.to_lowercase()
.chars()
.map(|c| if c.is_alphanumeric() { c } else { ' ' })
.collect()
}
}
impl Default for ConvergenceDetector {
fn default() -> Self {
let config = ConvergenceConfig::default();
let window_capacity = config.window_size;
Self {
config,
window: VecDeque::with_capacity(window_capacity),
consecutive_count: 0,
similar_responses: Vec::new(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_convergence_detection() {
let config = ConvergenceConfig {
window_size: 3,
similarity_threshold: 0.5,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
let status1 = detector.add_response("alpha");
assert!(!status1.detected, "first response: no comparison possible");
assert_eq!(
status1.consecutive_count, 1,
"first response: starts a streak of 1"
);
let status2 = detector.add_response("beta");
assert!(!status2.detected, "second response: different from first");
assert_eq!(
status2.consecutive_count, 1,
"second response: new streak of 1"
);
let status3 = detector.add_response("beta");
assert!(!status3.detected, "third response: streak of 2, need 3");
assert_eq!(status3.consecutive_count, 2);
let status4 = detector.add_response("beta");
assert!(
status4.detected,
"fourth response: three consecutive 'beta'"
);
assert!(status4.consecutive_count >= 3);
}
#[test]
fn test_no_convergence_different_responses() {
let config = ConvergenceConfig {
window_size: 3,
similarity_threshold: 0.95,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
detector.add_response("Response one about apples");
detector.add_response("Response two about oranges");
detector.add_response("Response three about bananas");
let status = detector.check_convergence();
assert!(!status.detected);
}
#[test]
fn test_convergence_disabled() {
let config = ConvergenceConfig {
enabled: false,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
detector.add_response("Same");
detector.add_response("Same");
detector.add_response("Same");
let status = detector.check_convergence();
assert!(!status.detected);
}
#[test]
fn test_similarity_computation() {
let sim = ConvergenceDetector::compute_similarity("hello world", "hello world");
assert!(sim > 0.99);
let sim = ConvergenceDetector::compute_similarity("hello world", "hello there");
assert!(sim > 0.3 && sim < 0.7);
let sim = ConvergenceDetector::compute_similarity("hello world", "goodbye moon");
assert!(sim < 0.3);
}
#[test]
fn test_clear_detector() {
let mut detector = ConvergenceDetector::default_detector().unwrap();
detector.add_response("Test");
detector.add_response("Test");
assert!(!detector.window.is_empty());
detector.clear();
assert!(detector.window.is_empty());
assert_eq!(detector.consecutive_count, 0);
}
#[test]
fn test_config_window_too_small() {
let config = ConvergenceConfig {
window_size: 1,
..Default::default()
};
let err = ConvergenceDetector::new(config).unwrap_err();
assert!(matches!(
err,
ConvergenceConfigError::WindowTooSmall { actual: 1 }
));
assert!(err.to_string().contains("at least 2"));
}
#[test]
fn test_config_window_zero() {
let config = ConvergenceConfig {
window_size: 0,
..Default::default()
};
let err = ConvergenceDetector::new(config).unwrap_err();
assert!(matches!(
err,
ConvergenceConfigError::WindowTooSmall { actual: 0 }
));
}
#[test]
fn test_config_threshold_too_high() {
let config = ConvergenceConfig {
similarity_threshold: 1.5,
..Default::default()
};
let err = ConvergenceDetector::new(config).unwrap_err();
assert!(matches!(
err,
ConvergenceConfigError::ThresholdOutOfRange { .. }
));
assert!(err.to_string().contains("[0.0, 1.0]"));
}
#[test]
fn test_config_threshold_negative() {
let config = ConvergenceConfig {
similarity_threshold: -0.1,
..Default::default()
};
let err = ConvergenceDetector::new(config).unwrap_err();
assert!(matches!(
err,
ConvergenceConfigError::ThresholdOutOfRange { .. }
));
}
#[test]
fn test_streak_resets_on_dissimilar_response() {
let config = ConvergenceConfig {
window_size: 3,
similarity_threshold: 0.5,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
let s1 = detector.add_response("same text");
assert_eq!(s1.consecutive_count, 1);
let s2 = detector.add_response("same text");
assert_eq!(s2.consecutive_count, 2);
let s3 = detector.add_response("completely different content here");
assert_eq!(
s3.consecutive_count, 1,
"dissimilar response should reset streak"
);
let s4 = detector.add_response("completely different content here");
assert_eq!(s4.consecutive_count, 2);
}
#[test]
fn test_alternating_responses_never_converge() {
let config = ConvergenceConfig {
window_size: 3,
similarity_threshold: 0.5,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
detector.add_response("alpha alpha alpha");
detector.add_response("beta beta beta");
detector.add_response("alpha alpha alpha");
detector.add_response("beta beta beta");
detector.add_response("alpha alpha alpha");
let status = detector.check_convergence();
assert!(
!status.detected,
"alternating responses must not trigger convergence"
);
assert!(
status.consecutive_count <= 1,
"alternating responses should not build a streak"
);
}
#[test]
fn test_similar_after_gap_starts_fresh_streak() {
let config = ConvergenceConfig {
window_size: 3,
similarity_threshold: 0.5,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
detector.add_response("same text");
let s2 = detector.add_response("same text");
assert_eq!(s2.consecutive_count, 2);
detector.add_response("totally different stuff");
let s4 = detector.add_response("same text");
assert_eq!(
s4.consecutive_count, 1,
"similar response after a gap should start a fresh streak"
);
}
#[test]
fn test_converge_then_break_then_re_converge() {
let config = ConvergenceConfig {
window_size: 3,
similarity_threshold: 0.5,
..Default::default()
};
let mut detector = ConvergenceDetector::new(config).unwrap();
detector.add_response("loop loop loop");
detector.add_response("loop loop loop");
let s3 = detector.add_response("loop loop loop");
assert!(s3.detected);
let s4 = detector.add_response("something entirely new and different");
assert!(!s4.detected);
let s5 = detector.add_response("something entirely new and different");
assert!(!s5.detected, "only 2 consecutive so far");
let s6 = detector.add_response("something entirely new and different");
assert!(s6.detected, "3 consecutive of the new response");
}
#[test]
fn test_config_threshold_boundary_valid() {
let config_zero = ConvergenceConfig {
similarity_threshold: 0.0,
..Default::default()
};
assert!(ConvergenceDetector::new(config_zero).is_ok());
let config_one = ConvergenceConfig {
similarity_threshold: 1.0,
..Default::default()
};
assert!(ConvergenceDetector::new(config_one).is_ok());
}
}