use std::sync::{Arc, Mutex};
use super::convergence::{ConvergenceConfig, ConvergenceDetector, ConvergenceStatus};
use super::loop_detector::{LoopDetector, LoopDetectorConfig, Operation, ToolSignature};
pub use super::convergence::ConvergenceAction;
pub use super::convergence::ConvergenceConfigError;
pub use super::loop_detector::LoopStatus;
#[derive(Debug, Clone)]
pub enum DetectedPattern {
LoopDetected {
repetitions: usize,
pattern_description: String,
},
ConvergenceDetected {
similarity: f32,
consecutive_count: usize,
},
NoPattern,
}
#[derive(Debug, Clone)]
pub struct DetectionConfig {
pub loop_threshold: usize,
pub stop_threshold: usize,
pub enable_loop_detection: bool,
pub max_history: usize,
pub convergence_threshold: f32,
pub convergence_count: usize,
pub enable_convergence_detection: bool,
pub on_converge: ConvergenceAction,
}
impl Default for DetectionConfig {
fn default() -> Self {
Self {
loop_threshold: 3,
stop_threshold: 10,
enable_loop_detection: true,
max_history: 100,
convergence_threshold: 0.95,
convergence_count: 3,
enable_convergence_detection: true,
on_converge: ConvergenceAction::default(),
}
}
}
impl DetectionConfig {
#[must_use]
pub fn to_convergence_config(&self) -> ConvergenceConfig {
ConvergenceConfig {
enabled: self.enable_convergence_detection,
window_size: self.convergence_count,
similarity_threshold: self.convergence_threshold,
on_converge: self.on_converge,
}
}
#[must_use]
pub fn to_loop_detector_config(&self) -> LoopDetectorConfig {
LoopDetectorConfig {
window_size: self.max_history,
repetition_threshold: self.loop_threshold,
stop_threshold: self.stop_threshold,
..LoopDetectorConfig::default()
}
}
}
#[derive(Debug, Clone, Default)]
pub struct DetectionStats {
pub turns_analyzed: usize,
pub loops_detected: usize,
pub convergences_detected: usize,
pub current_streak: usize,
}
pub struct DetectionManager {
config: DetectionConfig,
loop_detector: Arc<LoopDetector>,
convergence_detector: Mutex<ConvergenceDetector>,
stats: Mutex<DetectionStats>,
}
impl DetectionManager {
pub fn new() -> Result<Self, ConvergenceConfigError> {
Self::new_with_config(DetectionConfig::default())
}
pub fn new_with_config(config: DetectionConfig) -> Result<Self, ConvergenceConfigError> {
let loop_detector = Arc::new(LoopDetector::new(
config.to_loop_detector_config(),
Arc::new(super::loop_detector::NoOpToolSignature),
));
let convergence_detector =
Mutex::new(ConvergenceDetector::new(config.to_convergence_config())?);
Ok(Self {
config,
loop_detector,
convergence_detector,
stats: Mutex::new(DetectionStats::default()),
})
}
pub fn new_with_loop_detector(
config: DetectionConfig,
loop_detector: LoopDetector,
) -> Result<Self, ConvergenceConfigError> {
let convergence_detector =
Mutex::new(ConvergenceDetector::new(config.to_convergence_config())?);
Ok(Self {
config,
loop_detector: Arc::new(loop_detector),
convergence_detector,
stats: Mutex::new(DetectionStats::default()),
})
}
pub fn new_with_signature(
signature: Arc<dyn ToolSignature>,
) -> Result<Self, ConvergenceConfigError> {
let config = DetectionConfig::default();
let mut ldc = config.to_loop_detector_config();
ldc.tool_thresholds = signature.tool_thresholds();
let loop_detector = Arc::new(LoopDetector::new(ldc, signature));
let convergence_detector =
Mutex::new(ConvergenceDetector::new(config.to_convergence_config())?);
let stats = Mutex::new(DetectionStats::default());
Ok(Self {
config,
loop_detector,
convergence_detector,
stats,
})
}
pub fn record_operation(&self, operation: Operation) -> DetectedPattern {
if !self.config.enable_loop_detection {
return DetectedPattern::NoPattern;
}
self.loop_detector.record(operation);
let status = self.loop_detector.check_loop();
if status.is_looping {
let mut guard = self.stats.lock().unwrap_or_else(|e| {
tracing::warn!("stats lock poisoned, recovering");
e.into_inner()
});
guard.turns_analyzed = guard.turns_analyzed.saturating_add(1);
guard.loops_detected = guard.loops_detected.saturating_add(1);
guard.current_streak = status.repetition_count;
DetectedPattern::LoopDetected {
repetitions: status.repetition_count,
pattern_description: status
.repeated_operations
.first()
.map(|op| format!("{}({})", op.tool, op.primary_param))
.unwrap_or_default(),
}
} else {
let mut guard = self.stats.lock().unwrap_or_else(|e| {
tracing::warn!("stats lock poisoned, recovering");
e.into_inner()
});
guard.turns_analyzed = guard.turns_analyzed.saturating_add(1);
DetectedPattern::NoPattern
}
}
pub fn signature(&self) -> &dyn ToolSignature {
self.loop_detector.signature()
}
pub fn record_tool_call(&self, tool: &str, input_hash: u64) -> DetectedPattern {
let operation = Operation::new(tool, format!("hash:{input_hash}"));
self.record_operation(operation)
}
pub fn record_tool_call_with_result(
&self,
tool: &str,
input_hash: u64,
result_hash: Option<u64>,
) -> DetectedPattern {
let operation =
Operation::new(tool, format!("hash:{input_hash}")).with_result_hash(result_hash);
self.record_operation(operation)
}
#[must_use]
pub fn check_loop(&self) -> LoopStatus {
self.loop_detector.check_loop()
}
#[must_use]
pub fn loop_detector(&self) -> &LoopDetector {
&self.loop_detector
}
#[must_use]
pub fn convergence_detector(&self) -> &Mutex<ConvergenceDetector> {
&self.convergence_detector
}
pub fn record_response(&self, response: &str) -> DetectedPattern {
if !self.config.enable_convergence_detection {
return DetectedPattern::NoPattern;
}
let status = self
.convergence_detector
.lock()
.unwrap_or_else(|e| {
tracing::warn!("convergence detector lock poisoned, recovering");
e.into_inner()
})
.add_response(response);
if status.detected {
let mut guard = self.stats.lock().unwrap_or_else(|e| {
tracing::warn!("stats lock poisoned, recovering");
e.into_inner()
});
guard.convergences_detected = guard.convergences_detected.saturating_add(1);
return DetectedPattern::ConvergenceDetected {
similarity: status.similarity_score,
consecutive_count: status.consecutive_count,
};
}
DetectedPattern::NoPattern
}
#[must_use]
pub fn check_convergence(&self) -> ConvergenceStatus {
self.convergence_detector
.lock()
.unwrap_or_else(|e| {
tracing::warn!("convergence detector lock poisoned, recovering");
e.into_inner()
})
.check_convergence()
}
#[must_use]
pub fn check_current_pattern(&self) -> DetectedPattern {
let loop_status = self.loop_detector.check_loop();
if loop_status.is_looping {
return DetectedPattern::LoopDetected {
repetitions: loop_status.repetition_count,
pattern_description: loop_status
.repeated_operations
.first()
.map(|op| format!("{}({})", op.tool, op.primary_param))
.unwrap_or_default(),
};
}
let convergence = self.check_convergence();
if convergence.detected {
return DetectedPattern::ConvergenceDetected {
similarity: convergence.similarity_score,
consecutive_count: convergence.consecutive_count,
};
}
DetectedPattern::NoPattern
}
#[must_use]
pub fn stats(&self) -> DetectionStats {
self.stats
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone()
}
#[must_use]
pub fn config(&self) -> &DetectionConfig {
&self.config
}
pub fn reset(&self) {
self.loop_detector.reset();
self.convergence_detector
.lock()
.unwrap_or_else(|e| {
tracing::warn!("convergence detector lock poisoned, recovering");
e.into_inner()
})
.clear();
*self
.stats
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = DetectionStats::default();
}
}
impl Default for DetectionManager {
fn default() -> Self {
let config = DetectionConfig::default();
let loop_config = config.to_loop_detector_config();
let loop_detector = Arc::new(LoopDetector::new(
loop_config,
Arc::new(super::loop_detector::NoOpToolSignature),
));
let convergence_detector = Mutex::new(ConvergenceDetector::default());
Self {
config,
loop_detector,
convergence_detector,
stats: Mutex::new(DetectionStats::default()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_no_pattern_initially() {
let dm = DetectionManager::new().unwrap();
assert!(matches!(
dm.check_current_pattern(),
DetectedPattern::NoPattern
));
}
#[test]
fn test_loop_detection() {
let dm = DetectionManager::new().unwrap();
for _ in 0..5 {
let result = dm.record_tool_call("read_file", 42);
if matches!(result, DetectedPattern::LoopDetected { .. }) {
return; }
}
panic!("Expected loop detection after 5 identical calls");
}
#[test]
fn test_convergence_detection() {
let dm = DetectionManager::new().unwrap();
let response = "I am working on the task and making progress";
for _ in 0..5 {
let result = dm.record_response(response);
if matches!(result, DetectedPattern::ConvergenceDetected { .. }) {
return; }
}
panic!("Expected convergence detection after 5 identical responses");
}
#[test]
fn test_reset() {
let dm = DetectionManager::new().unwrap();
dm.record_tool_call("read_file", 42);
dm.record_response("hello");
dm.reset();
assert!(matches!(
dm.check_current_pattern(),
DetectedPattern::NoPattern
));
let stats = dm.stats();
assert_eq!(stats.turns_analyzed, 0);
}
#[test]
fn test_no_detection_when_disabled() {
let config = DetectionConfig {
enable_loop_detection: false,
enable_convergence_detection: false,
..Default::default()
};
let dm = DetectionManager::new_with_config(config).unwrap();
for _ in 0..10 {
let result = dm.record_tool_call("read_file", 42);
assert!(matches!(result, DetectedPattern::NoPattern));
}
}
#[test]
fn test_loop_status_should_stop() {
let config = DetectionConfig {
loop_threshold: 3,
stop_threshold: 5,
..Default::default()
};
let dm = DetectionManager::new_with_config(config).unwrap();
for _ in 0..5 {
dm.record_tool_call("read_file", 42);
}
let status = dm.check_loop();
assert!(status.is_looping);
assert!(status.should_stop);
assert!(status.warning.is_some());
assert!(status.warning.unwrap().contains("STOPPING"));
}
#[test]
fn test_record_operation_with_operation_struct() {
let dm = DetectionManager::new().unwrap();
for _ in 0..5 {
dm.record_operation(Operation::new("Read", "/test/file.txt"));
}
let status = dm.check_loop();
assert!(status.is_looping);
assert!(status.repetition_count >= 3);
}
#[test]
fn test_record_operation_from_input() {
use super::super::loop_detector::NoOpToolSignature;
let dm = DetectionManager::new().unwrap();
let input = serde_json::json!({"file_path": "/test/file.txt"});
for _ in 0..5 {
dm.record_operation(Operation::from_input_with_signature(
"Read",
&input,
&NoOpToolSignature,
));
}
let status = dm.check_loop();
assert!(status.is_looping);
}
#[test]
fn test_result_aware_detection() {
let dm = DetectionManager::new().unwrap();
for i in 0..5 {
let hash = super::super::loop_detector::hash_result(&format!("output {i}"));
dm.record_tool_call_with_result("Bash", 42, hash);
}
let status = dm.check_loop();
assert!(!status.is_looping, "Different results should not be a loop");
}
#[test]
fn test_result_aware_same_result_is_loop() {
let dm = DetectionManager::new().unwrap();
let hash = super::super::loop_detector::hash_result("same output");
for _ in 0..5 {
dm.record_tool_call_with_result("Bash", 42, hash);
}
let status = dm.check_loop();
assert!(status.is_looping, "Same results should be detected as loop");
}
#[test]
fn test_access_loop_detector() {
let dm = DetectionManager::new().unwrap();
assert_eq!(dm.loop_detector().turn_count(), 0);
}
#[test]
fn test_config_to_loop_detector_config() {
let config = DetectionConfig {
loop_threshold: 5,
stop_threshold: 15,
max_history: 200,
..Default::default()
};
let ldc = config.to_loop_detector_config();
assert_eq!(ldc.repetition_threshold, 5);
assert_eq!(ldc.stop_threshold, 15);
assert_eq!(ldc.window_size, 200);
}
#[test]
fn test_detection_config_default_has_no_max_response_history() {
let config = DetectionConfig {
loop_threshold: 3,
stop_threshold: 5,
..Default::default()
};
assert_eq!(config.loop_threshold, 3);
assert_eq!(config.stop_threshold, 5);
let default = DetectionConfig::default();
assert_eq!(default.loop_threshold, 3);
assert_eq!(default.stop_threshold, 10);
assert!(default.enable_loop_detection);
assert_eq!(default.max_history, 100);
assert!((default.convergence_threshold - 0.95).abs() < f32::EPSILON);
assert_eq!(default.convergence_count, 3);
assert!(default.enable_convergence_detection);
}
}