Skip to main content

tokio_prompt_orchestrator/
adaptive_timeout.rs

1//! Adaptive timeout management for LLM requests.
2//!
3//! Tracks per-model latency samples and computes dynamic timeouts based on
4//! observed percentile latencies and an exponential moving average. Timeouts
5//! are bounded between 5 s and 120 s and can be scaled up for high queue
6//! depths via [`AdaptiveTimeoutManager::adjust_for_load`].
7
8use dashmap::DashMap;
9use std::collections::VecDeque;
10use std::time::Duration;
11
12/// A single latency/outcome sample for one LLM request.
13#[derive(Debug, Clone)]
14pub struct TimeoutSample {
15    /// Model name (e.g. `"gpt-4o"`).
16    pub model_name: String,
17    /// Observed round-trip duration in milliseconds.
18    pub duration_ms: u64,
19    /// Whether the request succeeded (`false` = timeout / error).
20    pub success: bool,
21}
22
23/// Per-model statistics maintained by [`AdaptiveTimeoutManager`].
24#[derive(Debug)]
25pub struct ModelTimeoutStats {
26    /// Rolling window of the most recent 100 samples.
27    samples: VecDeque<TimeoutSample>,
28    /// Exponential moving average of duration_ms (α = 0.1).
29    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    /// Record a new sample, evicting the oldest when the window is full.
41    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    /// Fraction of samples that succeeded (0.0 if no samples).
51    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    /// Compute a percentile from the sorted duration values.
60    ///
61    /// `p` must be in `[0.0, 1.0]`. Returns `None` when there are no samples.
62    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    /// 50th-percentile latency in milliseconds.
74    pub fn p50(&self) -> Option<f64> {
75        self.percentile(0.50)
76    }
77
78    /// 95th-percentile latency in milliseconds.
79    pub fn p95(&self) -> Option<f64> {
80        self.percentile(0.95)
81    }
82
83    /// 99th-percentile latency in milliseconds.
84    pub fn p99(&self) -> Option<f64> {
85        self.percentile(0.99)
86    }
87
88    /// Number of samples in the window.
89    pub fn sample_count(&self) -> usize {
90        self.samples.len()
91    }
92}
93
94/// A snapshot of timeout statistics for a single model.
95#[derive(Debug, Clone)]
96pub struct TimeoutSummary {
97    /// Model identifier.
98    pub model: String,
99    /// 50th-percentile latency (ms).
100    pub p50_ms: f64,
101    /// 95th-percentile latency (ms).
102    pub p95_ms: f64,
103    /// 99th-percentile latency (ms).
104    pub p99_ms: f64,
105    /// Fraction of successful requests (0.0–1.0).
106    pub success_rate: f64,
107    /// Exponential moving average of latency (ms).
108    pub ema_ms: f64,
109    /// Number of samples in the rolling window.
110    pub sample_count: usize,
111}
112
113/// Thread-safe, per-model adaptive timeout manager.
114///
115/// Uses a [`DashMap`] so that multiple tokio tasks can record outcomes and
116/// query timeouts concurrently without a global lock.
117pub struct AdaptiveTimeoutManager {
118    stats: DashMap<String, ModelTimeoutStats>,
119}
120
121impl AdaptiveTimeoutManager {
122    /// Create a new manager with an empty model registry.
123    pub fn new() -> Self {
124        Self {
125            stats: DashMap::new(),
126        }
127    }
128
129    /// Record the outcome of one request for `model`.
130    ///
131    /// This updates the rolling sample window and the EMA.
132    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    /// Compute the recommended timeout for `model`.
145    ///
146    /// Returns `p95 * 1.5`, clamped to `[5 s, 120 s]`.  Falls back to 30 s
147    /// when fewer than two samples have been collected.
148    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    /// Adjust the base timeout by the current queue depth.
165    ///
166    /// Scales `get_timeout` by `sqrt(queue_depth / 10).max(1.0)` so that a
167    /// backlogged queue gets proportionally more time.
168    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        // Re-clamp after scaling.
173        let clamped = scaled_ms.min(120_000.0);
174        Duration::from_millis(clamped as u64)
175    }
176
177    /// Return a [`TimeoutSummary`] snapshot for `model`, or `None` if the
178    /// model has never been seen.
179    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        // p95 of 10 samples ≈ 950 ms; * 1.5 = 1425 ms — well above 5 s min.
211        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}