use std::collections::HashMap;
use chrono::{DateTime, Utc};
use super::models::{CostEvent, EventType};
pub const TRANSIENT_ERRORS: &[&str] = &[
"rate_limit",
"timeout",
"5xx",
"server_error",
"connection_error",
];
pub fn error_likelihood(error_type: &str) -> f64 {
match error_type {
"rate_limit" => 1.0,
"timeout" => 0.9,
"5xx" => 0.85,
"server_error" => 0.85,
"connection_error" => 0.8,
_ => 0.8,
}
}
#[derive(Debug, Clone)]
pub struct HeuristicMatch {
pub is_retry: bool,
pub confidence: f64,
pub matched_event_id: Option<String>,
pub reason: String,
}
impl HeuristicMatch {
fn no_match() -> Self {
Self {
is_retry: false,
confidence: 0.0,
matched_event_id: None,
reason: String::new(),
}
}
}
pub struct RetryHeuristicEngine {
window_seconds: f64,
threshold: f64,
recent_events: HashMap<String, Vec<CostEvent>>, }
impl RetryHeuristicEngine {
pub fn new(window_seconds: f64, threshold: f64) -> Result<Self, crate::error::DexcostError> {
if window_seconds <= 0.0 {
return Err(crate::error::DexcostError::Config(format!(
"window_seconds must be positive, got {}",
window_seconds
)));
}
if threshold <= 0.0 || threshold > 1.0 {
return Err(crate::error::DexcostError::Config(format!(
"threshold must be between 0 and 1, got {}",
threshold
)));
}
Ok(Self {
window_seconds,
threshold,
recent_events: HashMap::new(),
})
}
pub fn window_seconds(&self) -> f64 {
self.window_seconds
}
pub fn threshold(&self) -> f64 {
self.threshold
}
pub fn record(&mut self, event: CostEvent) {
let window = self.window_seconds;
let cutoff: DateTime<Utc> = event.occurred_at;
let bucket = self.recent_events.entry(event.task_id.clone()).or_default();
bucket.retain(|e| {
let gap = (cutoff - e.occurred_at).num_milliseconds() as f64 / 1000.0;
gap >= 0.0 && gap <= window
});
bucket.push(event);
const PER_TASK_CAP: usize = 1000;
if bucket.len() > PER_TASK_CAP {
let drop_n = PER_TASK_CAP / 10;
bucket.drain(..drop_n);
}
}
pub fn check(&self, event: &CostEvent) -> HeuristicMatch {
let events = match self.recent_events.get(&event.task_id) {
Some(v) if !v.is_empty() => v,
_ => return HeuristicMatch::no_match(),
};
for candidate in events.iter().rev() {
if candidate.event_id == event.event_id {
continue;
}
if candidate.event_type != EventType::LlmCall {
continue;
}
if candidate.model != event.model {
continue;
}
let error_type = match candidate.details.get("error_type") {
Some(serde_json::Value::String(s)) => s.as_str(),
_ => return HeuristicMatch::no_match(),
};
if !TRANSIENT_ERRORS.contains(&error_type) {
return HeuristicMatch::no_match();
}
let gap =
(event.occurred_at - candidate.occurred_at).num_milliseconds() as f64 / 1000.0;
if gap < 0.0 || gap > self.window_seconds {
return HeuristicMatch::no_match();
}
let base = error_likelihood(error_type);
let time_decay = (1.0 - gap / self.window_seconds).max(0.0);
let confidence = base * time_decay;
if confidence >= self.threshold {
return HeuristicMatch {
is_retry: true,
confidence,
matched_event_id: Some(candidate.event_id.clone()),
reason: error_type.to_string(),
};
}
return HeuristicMatch::no_match();
}
HeuristicMatch::no_match()
}
}
pub const DEFAULT_WINDOW_SECONDS: f64 = 30.0;
pub const DEFAULT_THRESHOLD: f64 = 0.8;
#[cfg(test)]
mod tests {
use chrono::Duration;
use rust_decimal::Decimal;
use super::*;
use crate::core::models::{CostConfidence, CostEvent, EventType};
fn make_llm_event(task_id: &str, model: &str) -> CostEvent {
let mut e = CostEvent::new(task_id, EventType::LlmCall);
e.model = Some(model.to_string());
e.cost_usd = Decimal::ZERO;
e.cost_confidence = CostConfidence::Unknown;
e
}
fn with_error(mut event: CostEvent, error_type: &str) -> CostEvent {
event
.details
.insert("error_type".to_string(), serde_json::json!(error_type));
event
}
#[test]
fn test_detects_retry_after_transient_error() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.8).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-1", "gpt-4o");
first.occurred_at = base_time;
let first = with_error(first, "rate_limit"); engine.record(first.clone());
let mut second = make_llm_event("task-1", "gpt-4o");
second.occurred_at = base_time + Duration::seconds(1); engine.record(second.clone());
let result = engine.check(&second);
assert!(result.is_retry);
assert!(result.confidence >= 0.8);
assert_eq!(result.matched_event_id, Some(first.event_id.clone()));
assert_eq!(result.reason, "rate_limit");
}
#[test]
fn test_no_flag_different_model() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.8).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-1", "gpt-4o");
first.occurred_at = base_time;
let first = with_error(first, "rate_limit");
engine.record(first);
let mut second = make_llm_event("task-1", "claude-3.5-sonnet");
second.occurred_at = base_time + Duration::seconds(1);
engine.record(second.clone());
let result = engine.check(&second);
assert!(!result.is_retry);
}
#[test]
fn test_no_flag_different_task() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.8).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-A", "gpt-4o");
first.occurred_at = base_time;
let first = with_error(first, "rate_limit");
engine.record(first);
let mut second = make_llm_event("task-B", "gpt-4o");
second.occurred_at = base_time + Duration::seconds(1);
engine.record(second.clone());
let result = engine.check(&second);
assert!(!result.is_retry);
}
#[test]
fn test_no_flag_when_previous_succeeded() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.8).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-1", "gpt-4o");
first.occurred_at = base_time;
engine.record(first);
let mut second = make_llm_event("task-1", "gpt-4o");
second.occurred_at = base_time + Duration::seconds(1);
engine.record(second.clone());
let result = engine.check(&second);
assert!(!result.is_retry);
}
#[test]
fn test_no_flag_outside_window() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.8).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-1", "gpt-4o");
first.occurred_at = base_time;
let first = with_error(first, "rate_limit");
engine.record(first);
let mut second = make_llm_event("task-1", "gpt-4o");
second.occurred_at = base_time + Duration::seconds(31);
engine.record(second.clone());
let result = engine.check(&second);
assert!(!result.is_retry);
}
#[test]
fn test_confidence_decays_with_gap() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.001).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-1", "gpt-4o");
first.occurred_at = base_time;
let first = with_error(first, "rate_limit"); engine.record(first);
let mut second = make_llm_event("task-1", "gpt-4o");
second.occurred_at = base_time + Duration::seconds(15);
engine.record(second.clone());
let result = engine.check(&second);
assert!(result.is_retry); let expected = 1.0 * (1.0 - 15.0 / 30.0);
assert!((result.confidence - expected).abs() < 0.001);
}
#[test]
fn test_prunes_old_events() {
let mut engine = RetryHeuristicEngine::new(30.0, 0.8).unwrap();
let base_time = Utc::now();
let mut first = make_llm_event("task-1", "gpt-4o");
first.occurred_at = base_time;
let first = with_error(first, "rate_limit");
engine.record(first);
let mut second = make_llm_event("task-1", "gpt-4o");
second.occurred_at = base_time + Duration::seconds(31);
engine.record(second.clone());
let bucket = engine.recent_events.get("task-1").unwrap();
assert_eq!(bucket.len(), 1);
assert_eq!(bucket[0].event_id, second.event_id);
}
#[test]
fn test_defaults() {
assert_eq!(DEFAULT_WINDOW_SECONDS, 30.0);
assert_eq!(DEFAULT_THRESHOLD, 0.8);
let engine = RetryHeuristicEngine::new(DEFAULT_WINDOW_SECONDS, DEFAULT_THRESHOLD).unwrap();
assert_eq!(engine.window_seconds(), 30.0);
assert_eq!(engine.threshold(), 0.8);
}
}