tokio_prompt_orchestrator/
adaptive_timeout.rs1use dashmap::DashMap;
9use std::collections::VecDeque;
10use std::time::Duration;
11
12#[derive(Debug, Clone)]
14pub struct TimeoutSample {
15 pub model_name: String,
17 pub duration_ms: u64,
19 pub success: bool,
21}
22
23#[derive(Debug)]
25pub struct ModelTimeoutStats {
26 samples: VecDeque<TimeoutSample>,
28 pub ema_ms: f64,
30}
31
32impl ModelTimeoutStats {
33 fn new() -> Self {
34 Self {
35 samples: VecDeque::with_capacity(100),
36 ema_ms: 5_000.0,
37 }
38 }
39
40 fn push(&mut self, sample: TimeoutSample) {
42 const ALPHA: f64 = 0.1;
43 self.ema_ms = ALPHA * sample.duration_ms as f64 + (1.0 - ALPHA) * self.ema_ms;
44 if self.samples.len() >= 100 {
45 self.samples.pop_front();
46 }
47 self.samples.push_back(sample);
48 }
49
50 pub fn success_rate(&self) -> f64 {
52 if self.samples.is_empty() {
53 return 0.0;
54 }
55 let successes = self.samples.iter().filter(|s| s.success).count();
56 successes as f64 / self.samples.len() as f64
57 }
58
59 fn percentile(&self, p: f64) -> Option<f64> {
63 if self.samples.is_empty() {
64 return None;
65 }
66 let mut durations: Vec<u64> = self.samples.iter().map(|s| s.duration_ms).collect();
67 durations.sort_unstable();
68 let idx = ((p * (durations.len() as f64 - 1.0)).round() as usize)
69 .min(durations.len() - 1);
70 Some(durations[idx] as f64)
71 }
72
73 pub fn p50(&self) -> Option<f64> {
75 self.percentile(0.50)
76 }
77
78 pub fn p95(&self) -> Option<f64> {
80 self.percentile(0.95)
81 }
82
83 pub fn p99(&self) -> Option<f64> {
85 self.percentile(0.99)
86 }
87
88 pub fn sample_count(&self) -> usize {
90 self.samples.len()
91 }
92}
93
94#[derive(Debug, Clone)]
96pub struct TimeoutSummary {
97 pub model: String,
99 pub p50_ms: f64,
101 pub p95_ms: f64,
103 pub p99_ms: f64,
105 pub success_rate: f64,
107 pub ema_ms: f64,
109 pub sample_count: usize,
111}
112
113pub struct AdaptiveTimeoutManager {
118 stats: DashMap<String, ModelTimeoutStats>,
119}
120
121impl AdaptiveTimeoutManager {
122 pub fn new() -> Self {
124 Self {
125 stats: DashMap::new(),
126 }
127 }
128
129 pub fn record_outcome(&self, model: &str, duration_ms: u64, success: bool) {
133 let sample = TimeoutSample {
134 model_name: model.to_string(),
135 duration_ms,
136 success,
137 };
138 self.stats
139 .entry(model.to_string())
140 .or_insert_with(ModelTimeoutStats::new)
141 .push(sample);
142 }
143
144 pub fn get_timeout(&self, model: &str) -> Duration {
149 const MIN_MS: f64 = 5_000.0;
150 const MAX_MS: f64 = 120_000.0;
151 const DEFAULT_MS: f64 = 30_000.0;
152
153 let timeout_ms = self
154 .stats
155 .get(model)
156 .and_then(|s| s.p95())
157 .map(|p95| p95 * 1.5)
158 .unwrap_or(DEFAULT_MS);
159
160 let clamped = timeout_ms.clamp(MIN_MS, MAX_MS);
161 Duration::from_millis(clamped as u64)
162 }
163
164 pub fn adjust_for_load(&self, model: &str, queue_depth: usize) -> Duration {
169 let base = self.get_timeout(model);
170 let scale = ((queue_depth as f64 / 10.0).sqrt()).max(1.0);
171 let scaled_ms = base.as_millis() as f64 * scale;
172 let clamped = scaled_ms.min(120_000.0);
174 Duration::from_millis(clamped as u64)
175 }
176
177 pub fn model_summary(&self, model: &str) -> Option<TimeoutSummary> {
180 let stats = self.stats.get(model)?;
181 Some(TimeoutSummary {
182 model: model.to_string(),
183 p50_ms: stats.p50().unwrap_or(0.0),
184 p95_ms: stats.p95().unwrap_or(0.0),
185 p99_ms: stats.p99().unwrap_or(0.0),
186 success_rate: stats.success_rate(),
187 ema_ms: stats.ema_ms,
188 sample_count: stats.sample_count(),
189 })
190 }
191}
192
193impl Default for AdaptiveTimeoutManager {
194 fn default() -> Self {
195 Self::new()
196 }
197}
198
199#[cfg(test)]
200mod tests {
201 use super::*;
202
203 #[test]
204 fn test_record_and_get_timeout() {
205 let mgr = AdaptiveTimeoutManager::new();
206 for ms in [100u64, 200, 300, 400, 500, 600, 700, 800, 900, 1000] {
207 mgr.record_outcome("gpt-4o", ms, true);
208 }
209 let t = mgr.get_timeout("gpt-4o");
210 assert!(t >= Duration::from_secs(5));
212 assert!(t <= Duration::from_secs(120));
213 }
214
215 #[test]
216 fn test_default_timeout_for_unknown_model() {
217 let mgr = AdaptiveTimeoutManager::new();
218 assert_eq!(mgr.get_timeout("unknown"), Duration::from_secs(30));
219 }
220
221 #[test]
222 fn test_adjust_for_load_scales_up() {
223 let mgr = AdaptiveTimeoutManager::new();
224 let base = mgr.adjust_for_load("m", 0);
225 let loaded = mgr.adjust_for_load("m", 100);
226 assert!(loaded >= base);
227 }
228
229 #[test]
230 fn test_success_rate() {
231 let mgr = AdaptiveTimeoutManager::new();
232 mgr.record_outcome("m", 100, true);
233 mgr.record_outcome("m", 100, false);
234 let summary = mgr.model_summary("m").expect("summary present");
235 assert!((summary.success_rate - 0.5).abs() < 1e-9);
236 }
237}