Skip to main content

tokio_prompt_orchestrator/
adaptive_pool.rs

1//! Adaptive worker pool with Kalman-filter latency prediction.
2//!
3//! Monitors pipeline queue depths and inferred latency to decide when to
4//! spawn additional worker tasks or drain excess ones. Uses a one-dimensional
5//! Kalman filter to smooth noisy queue-depth observations and predict the
6//! short-term trend, avoiding oscillation from reactive on/off control.
7//!
8//! ## Design
9//!
10//! ```text
11//! ┌──────────────────────────────────────────────────────────┐
12//! │  AdaptivePool                                             │
13//! │                                                           │
14//! │  Queue depth ──► KalmanFilter ──► predicted depth        │
15//! │                                        │                 │
16//! │  Latency EMA ──────────────────────► ScaleDecision       │
17//! │                                        │                 │
18//! │                            spawn task / mark idle        │
19//! └──────────────────────────────────────────────────────────┘
20//! ```
21//!
22//! ## Kalman Filter
23//!
24//! A scalar 1-D Kalman filter tracks the queue depth signal:
25//!
26//! - **State** `x`: estimated true queue depth
27//! - **Process noise** `Q`: models how fast the queue can change (default 1.0)
28//! - **Measurement noise** `R`: models observation noise (default 5.0)
29//!
30//! The filter converges to the true depth in ~5–10 observations and provides a
31//! smooth signal that drives scaling without reacting to single-sample spikes.
32//!
33//! ## Scaling Policy
34//!
35//! | Condition | Action |
36//! |-----------|--------|
37//! | predicted_depth > `scale_up_threshold` AND latency > `latency_threshold_ms` | Recommend scale-up |
38//! | predicted_depth < `scale_down_threshold` AND pool_size > `min_workers` | Recommend scale-down |
39//! | otherwise | Stable |
40
41use serde::{Deserialize, Serialize};
42use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
43use std::sync::Arc;
44use std::time::{Duration, Instant};
45use tokio::sync::Mutex;
46use tracing::{debug, info};
47
48/// Scaling recommendation produced by [`AdaptivePool::evaluate`].
49#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
50pub enum ScaleDecision {
51    /// No change required — pool size is appropriate.
52    Stable,
53    /// Add `n` workers to handle increasing load.
54    ScaleUp {
55        /// Number of workers to add.
56        by: usize,
57    },
58    /// Remove `n` workers to reduce idle resource consumption.
59    ScaleDown {
60        /// Number of workers to remove.
61        by: usize,
62    },
63}
64
65/// Configuration for the adaptive pool controller.
66#[derive(Debug, Clone, Serialize, Deserialize)]
67pub struct AdaptivePoolConfig {
68    /// Minimum number of workers to maintain at all times.
69    pub min_workers: usize,
70    /// Maximum number of workers the pool may grow to.
71    pub max_workers: usize,
72    /// Queue depth (predicted) above which scale-up is triggered.
73    pub scale_up_threshold: f64,
74    /// Queue depth (predicted) below which scale-down is considered.
75    pub scale_down_threshold: f64,
76    /// P99 latency (ms) above which scale-up is triggered alongside depth check.
77    pub latency_threshold_ms: f64,
78    /// Minimum time between consecutive scale events (prevents thrashing).
79    pub cooldown: Duration,
80    /// Process noise for the Kalman filter (models queue volatility).
81    pub kalman_process_noise: f64,
82    /// Measurement noise for the Kalman filter (models observation noise).
83    pub kalman_measurement_noise: f64,
84}
85
86impl Default for AdaptivePoolConfig {
87    fn default() -> Self {
88        Self {
89            min_workers: 1,
90            max_workers: 32,
91            scale_up_threshold: 50.0,
92            scale_down_threshold: 5.0,
93            latency_threshold_ms: 500.0,
94            cooldown: Duration::from_secs(10),
95            kalman_process_noise: 1.0,
96            kalman_measurement_noise: 5.0,
97        }
98    }
99}
100
101/// Scalar 1-D Kalman filter.
102///
103/// Tracks a noisy scalar signal (e.g. queue depth) and produces a smoothed
104/// estimate. The filter runs in O(1) time and O(1) space.
105///
106/// ## Math
107///
108/// Predict:
109/// ```text
110/// x_pred = x
111/// p_pred = p + Q
112/// ```
113///
114/// Update:
115/// ```text
116/// K = p_pred / (p_pred + R)
117/// x = x_pred + K * (measurement - x_pred)
118/// p = (1 - K) * p_pred
119/// ```
120#[derive(Debug, Clone)]
121pub struct KalmanFilter {
122    /// Current state estimate.
123    pub x: f64,
124    /// Estimate covariance.
125    pub p: f64,
126    /// Process noise covariance.
127    pub q: f64,
128    /// Measurement noise covariance.
129    pub r: f64,
130    /// Whether the filter has been initialised with a first measurement.
131    initialised: bool,
132}
133
134impl KalmanFilter {
135    /// Create a new filter with the given noise parameters.
136    pub fn new(process_noise: f64, measurement_noise: f64) -> Self {
137        Self {
138            x: 0.0,
139            p: 1.0,
140            q: process_noise,
141            r: measurement_noise,
142            initialised: false,
143        }
144    }
145
146    /// Feed a new observation and return the filtered estimate.
147    ///
148    /// On the first call, the filter is initialised directly to the measurement.
149    pub fn update(&mut self, measurement: f64) -> f64 {
150        if !self.initialised {
151            self.x = measurement;
152            self.initialised = true;
153            return self.x;
154        }
155
156        // Predict step
157        let x_pred = self.x;
158        let p_pred = self.p + self.q;
159
160        // Update step
161        let k = p_pred / (p_pred + self.r);
162        self.x = x_pred + k * (measurement - x_pred);
163        self.p = (1.0 - k) * p_pred;
164
165        self.x
166    }
167
168    /// Return the current smoothed estimate without feeding a new observation.
169    pub fn estimate(&self) -> f64 {
170        self.x
171    }
172}
173
174/// EMA tracker for latency observations.
175#[derive(Debug, Clone)]
176pub struct LatencyEma {
177    value: f64,
178    alpha: f64,
179    sample_count: u64,
180}
181
182impl LatencyEma {
183    /// Create a new EMA tracker with the given smoothing factor (0 < alpha ≤ 1).
184    pub fn new(alpha: f64) -> Self {
185        Self {
186            value: 0.0,
187            alpha: alpha.clamp(0.001, 1.0),
188            sample_count: 0,
189        }
190    }
191
192    /// Feed a latency observation and return the new EMA.
193    pub fn update(&mut self, latency_ms: f64) -> f64 {
194        if self.sample_count == 0 {
195            self.value = latency_ms;
196        } else {
197            self.value = self.alpha * latency_ms + (1.0 - self.alpha) * self.value;
198        }
199        self.sample_count += 1;
200        self.value
201    }
202
203    /// Current EMA value.
204    pub fn current(&self) -> f64 {
205        self.value
206    }
207
208    /// Number of samples fed.
209    pub fn count(&self) -> u64 {
210        self.sample_count
211    }
212}
213
214/// Internal mutable state, behind a mutex to allow async update.
215struct PoolState {
216    kalman: KalmanFilter,
217    latency_ema: LatencyEma,
218    current_workers: usize,
219    last_scale_at: Option<Instant>,
220}
221
222/// Pool statistics snapshot.
223#[derive(Debug, Clone, Serialize)]
224pub struct PoolStats {
225    /// Current number of workers.
226    pub current_workers: usize,
227    /// Kalman-filtered queue depth estimate.
228    pub estimated_queue_depth: f64,
229    /// Latency EMA in milliseconds.
230    pub latency_ema_ms: f64,
231    /// Total scale-up decisions since creation.
232    pub total_scale_ups: u64,
233    /// Total scale-down decisions since creation.
234    pub total_scale_downs: u64,
235}
236
237/// Adaptive pool controller.
238///
239/// This is a **controller** — it recommends scaling decisions but does not
240/// itself spawn or kill tasks. The caller is responsible for acting on
241/// [`ScaleDecision`] values returned by [`AdaptivePool::evaluate`].
242///
243/// # Thread Safety
244///
245/// `AdaptivePool` is `Send + Sync` and may be shared across tasks.
246pub struct AdaptivePool {
247    config: AdaptivePoolConfig,
248    state: Mutex<PoolState>,
249    total_scale_ups: AtomicU64,
250    total_scale_downs: AtomicU64,
251    observations: AtomicUsize,
252}
253
254impl AdaptivePool {
255    /// Create a new pool controller with the given config and initial worker count.
256    pub fn new(config: AdaptivePoolConfig, initial_workers: usize) -> Arc<Self> {
257        let initial_workers = initial_workers.max(config.min_workers);
258        Arc::new(Self {
259            state: Mutex::new(PoolState {
260                kalman: KalmanFilter::new(
261                    config.kalman_process_noise,
262                    config.kalman_measurement_noise,
263                ),
264                latency_ema: LatencyEma::new(0.15),
265                current_workers: initial_workers,
266                last_scale_at: None,
267            }),
268            config,
269            total_scale_ups: AtomicU64::new(0),
270            total_scale_downs: AtomicU64::new(0),
271            observations: AtomicUsize::new(0),
272        })
273    }
274
275    /// Feed a new observation and return a scaling recommendation.
276    ///
277    /// Call this on each pipeline tick (e.g. every 500 ms) to get decisions.
278    ///
279    /// # Arguments
280    /// * `queue_depth` — raw queue depth observation (number of queued items).
281    /// * `latency_ms`  — recent P99 or average latency in milliseconds.
282    pub async fn evaluate(&self, queue_depth: usize, latency_ms: f64) -> ScaleDecision {
283        let mut state = self.state.lock().await;
284        self.observations.fetch_add(1, Ordering::Relaxed);
285
286        // Update filters
287        let smoothed_depth = state.kalman.update(queue_depth as f64);
288        let smoothed_latency = state.latency_ema.update(latency_ms);
289
290        debug!(
291            raw_depth = queue_depth,
292            smoothed_depth,
293            raw_latency_ms = latency_ms,
294            smoothed_latency_ms = smoothed_latency,
295            current_workers = state.current_workers,
296            "adaptive pool observation"
297        );
298
299        // Enforce cooldown
300        if let Some(last) = state.last_scale_at {
301            if last.elapsed() < self.config.cooldown {
302                return ScaleDecision::Stable;
303            }
304        }
305
306        let decision = if smoothed_depth > self.config.scale_up_threshold
307            && smoothed_latency > self.config.latency_threshold_ms
308            && state.current_workers < self.config.max_workers
309        {
310            // Scale up — add workers proportional to overload
311            let headroom = self.config.max_workers - state.current_workers;
312            let by = ((smoothed_depth / self.config.scale_up_threshold) as usize)
313                .max(1)
314                .min(headroom)
315                .min(4); // cap single-event scale-up at 4
316            ScaleDecision::ScaleUp { by }
317        } else if smoothed_depth < self.config.scale_down_threshold
318            && smoothed_latency < self.config.latency_threshold_ms * 0.5
319            && state.current_workers > self.config.min_workers
320        {
321            let excess = state.current_workers - self.config.min_workers;
322            let by = (excess / 2).max(1);
323            ScaleDecision::ScaleDown { by }
324        } else {
325            ScaleDecision::Stable
326        };
327
328        if decision != ScaleDecision::Stable {
329            state.last_scale_at = Some(Instant::now());
330            match decision {
331                ScaleDecision::ScaleUp { by } => {
332                    state.current_workers =
333                        (state.current_workers + by).min(self.config.max_workers);
334                    self.total_scale_ups.fetch_add(1, Ordering::Relaxed);
335                    info!(
336                        by,
337                        new_total = state.current_workers,
338                        depth = smoothed_depth,
339                        latency_ms = smoothed_latency,
340                        "adaptive pool: scale up"
341                    );
342                }
343                ScaleDecision::ScaleDown { by } => {
344                    state.current_workers =
345                        (state.current_workers - by).max(self.config.min_workers);
346                    self.total_scale_downs.fetch_add(1, Ordering::Relaxed);
347                    info!(
348                        by,
349                        new_total = state.current_workers,
350                        depth = smoothed_depth,
351                        latency_ms = smoothed_latency,
352                        "adaptive pool: scale down"
353                    );
354                }
355                ScaleDecision::Stable => {}
356            }
357        }
358
359        decision
360    }
361
362    /// Notify the controller that the worker count changed externally.
363    pub async fn set_workers(&self, count: usize) {
364        let mut state = self.state.lock().await;
365        state.current_workers = count.clamp(self.config.min_workers, self.config.max_workers);
366    }
367
368    /// Return a statistics snapshot.
369    pub async fn stats(&self) -> PoolStats {
370        let state = self.state.lock().await;
371        PoolStats {
372            current_workers: state.current_workers,
373            estimated_queue_depth: state.kalman.estimate(),
374            latency_ema_ms: state.latency_ema.current(),
375            total_scale_ups: self.total_scale_ups.load(Ordering::Relaxed),
376            total_scale_downs: self.total_scale_downs.load(Ordering::Relaxed),
377        }
378    }
379
380    /// Total number of observations fed.
381    pub fn observation_count(&self) -> usize {
382        self.observations.load(Ordering::Relaxed)
383    }
384}
385
386/// Runs the adaptive pool evaluation loop as a Tokio background task.
387///
388/// Polls `queue_depth_fn` and `latency_fn` at `interval`, feeds observations
389/// to the pool, and calls `on_decision` whenever a non-Stable decision is made.
390///
391/// Returns a [`tokio::task::JoinHandle`] for the background task.
392///
393/// # Example
394///
395/// ```no_run
396/// use std::sync::Arc;
397/// use std::time::Duration;
398/// use tokio_prompt_orchestrator::adaptive_pool::{AdaptivePool, AdaptivePoolConfig, ScaleDecision};
399///
400/// # async fn example() {
401/// let pool = AdaptivePool::new(AdaptivePoolConfig::default(), 2);
402/// let pool_clone = Arc::clone(&pool);
403///
404/// let handle = tokio_prompt_orchestrator::adaptive_pool::run_pool_controller(
405///     Arc::clone(&pool),
406///     Duration::from_millis(500),
407///     || 10usize,   // queue_depth_fn
408///     || 200.0f64,  // latency_ms_fn
409///     |decision| Box::pin(async move {
410///         match decision {
411///             ScaleDecision::ScaleUp { by } => println!("Spawning {by} workers"),
412///             ScaleDecision::ScaleDown { by } => println!("Draining {by} workers"),
413///             ScaleDecision::Stable => {}
414///         }
415///     }),
416/// );
417/// # }
418/// ```
419pub fn run_pool_controller<QFn, LFn, OFn, OFut>(
420    pool: Arc<AdaptivePool>,
421    interval: Duration,
422    queue_depth_fn: QFn,
423    latency_fn: LFn,
424    on_decision: OFn,
425) -> tokio::task::JoinHandle<()>
426where
427    QFn: Fn() -> usize + Send + 'static,
428    LFn: Fn() -> f64 + Send + 'static,
429    OFn: Fn(ScaleDecision) -> OFut + Send + 'static,
430    OFut: std::future::Future<Output = ()> + Send + 'static,
431{
432    tokio::spawn(async move {
433        let mut ticker = tokio::time::interval(interval);
434        loop {
435            ticker.tick().await;
436            let depth = queue_depth_fn();
437            let latency = latency_fn();
438            let decision = pool.evaluate(depth, latency).await;
439            if decision != ScaleDecision::Stable {
440                on_decision(decision).await;
441            }
442        }
443    })
444}
445
446#[cfg(test)]
447mod tests {
448    use super::*;
449
450    #[test]
451    fn test_kalman_converges() {
452        let mut kf = KalmanFilter::new(1.0, 5.0);
453        // Feed constant signal of 100
454        for _ in 0..20 {
455            kf.update(100.0);
456        }
457        // Should converge close to 100
458        assert!(
459            (kf.estimate() - 100.0).abs() < 5.0,
460            "estimate={}, expected ~100",
461            kf.estimate()
462        );
463    }
464
465    #[test]
466    fn test_kalman_tracks_step_change() {
467        let mut kf = KalmanFilter::new(1.0, 5.0);
468        for _ in 0..10 {
469            kf.update(50.0);
470        }
471        for _ in 0..10 {
472            kf.update(100.0);
473        }
474        // Should have moved toward 100
475        assert!(kf.estimate() > 70.0, "estimate={}", kf.estimate());
476    }
477
478    #[test]
479    fn test_latency_ema_initialises_to_first_sample() {
480        let mut ema = LatencyEma::new(0.1);
481        let v = ema.update(200.0);
482        assert_eq!(v, 200.0);
483    }
484
485    #[test]
486    fn test_latency_ema_blends_samples() {
487        let mut ema = LatencyEma::new(0.5);
488        ema.update(100.0);
489        let v = ema.update(200.0);
490        assert!((v - 150.0).abs() < 1.0, "v={v}");
491    }
492
493    #[tokio::test]
494    async fn test_stable_when_within_thresholds() {
495        let config = AdaptivePoolConfig {
496            scale_up_threshold: 100.0,
497            scale_down_threshold: 5.0,
498            latency_threshold_ms: 500.0,
499            min_workers: 1,
500            max_workers: 16,
501            cooldown: Duration::from_millis(0),
502            ..Default::default()
503        };
504        let pool = AdaptivePool::new(config, 2);
505        let decision = pool.evaluate(10, 100.0).await;
506        assert_eq!(decision, ScaleDecision::Stable);
507    }
508
509    #[tokio::test]
510    async fn test_scale_up_under_high_load() {
511        let config = AdaptivePoolConfig {
512            scale_up_threshold: 10.0,
513            latency_threshold_ms: 100.0,
514            min_workers: 1,
515            max_workers: 16,
516            cooldown: Duration::from_millis(0),
517            ..Default::default()
518        };
519        let pool = AdaptivePool::new(config, 2);
520        // Feed high load to warm up Kalman filter
521        for _ in 0..5 {
522            pool.evaluate(200, 1000.0).await;
523        }
524        // After warmup should see scale-up
525        let stats = pool.stats().await;
526        assert!(stats.current_workers > 2, "expected scale-up, got {} workers", stats.current_workers);
527    }
528
529    #[tokio::test]
530    async fn test_cooldown_prevents_rapid_scaling() {
531        let config = AdaptivePoolConfig {
532            scale_up_threshold: 10.0,
533            latency_threshold_ms: 50.0,
534            cooldown: Duration::from_secs(60), // very long cooldown
535            min_workers: 1,
536            max_workers: 16,
537            ..Default::default()
538        };
539        let pool = AdaptivePool::new(config, 2);
540        let d1 = pool.evaluate(200, 1000.0).await;
541        let d2 = pool.evaluate(200, 1000.0).await;
542        // Second decision should be Stable due to cooldown
543        if d1 != ScaleDecision::Stable {
544            assert_eq!(d2, ScaleDecision::Stable, "cooldown should suppress second scale event");
545        }
546    }
547
548    #[tokio::test]
549    async fn test_scale_down_under_low_load() {
550        let config = AdaptivePoolConfig {
551            scale_up_threshold: 100.0,
552            scale_down_threshold: 20.0,
553            latency_threshold_ms: 500.0,
554            min_workers: 1,
555            max_workers: 16,
556            cooldown: Duration::from_millis(0),
557            ..Default::default()
558        };
559        let pool = AdaptivePool::new(config, 8);
560
561        for _ in 0..10 {
562            pool.evaluate(1, 10.0).await;
563        }
564
565        let stats = pool.stats().await;
566        assert!(
567            stats.current_workers < 8,
568            "expected scale-down from 8, got {}",
569            stats.current_workers
570        );
571    }
572
573    #[tokio::test]
574    async fn test_min_workers_respected() {
575        let config = AdaptivePoolConfig {
576            scale_down_threshold: 100.0, // always trigger scale-down
577            latency_threshold_ms: 1000.0,
578            min_workers: 3,
579            max_workers: 16,
580            cooldown: Duration::from_millis(0),
581            ..Default::default()
582        };
583        let pool = AdaptivePool::new(config, 5);
584        for _ in 0..20 {
585            pool.evaluate(0, 0.0).await;
586        }
587        let stats = pool.stats().await;
588        assert!(
589            stats.current_workers >= 3,
590            "must not go below min_workers=3, got {}",
591            stats.current_workers
592        );
593    }
594}