1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
50pub enum ScaleDecision {
51 Stable,
53 ScaleUp {
55 by: usize,
57 },
58 ScaleDown {
60 by: usize,
62 },
63}
64
65#[derive(Debug, Clone, Serialize, Deserialize)]
67pub struct AdaptivePoolConfig {
68 pub min_workers: usize,
70 pub max_workers: usize,
72 pub scale_up_threshold: f64,
74 pub scale_down_threshold: f64,
76 pub latency_threshold_ms: f64,
78 pub cooldown: Duration,
80 pub kalman_process_noise: f64,
82 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#[derive(Debug, Clone)]
121pub struct KalmanFilter {
122 pub x: f64,
124 pub p: f64,
126 pub q: f64,
128 pub r: f64,
130 initialised: bool,
132}
133
134impl KalmanFilter {
135 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 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 let x_pred = self.x;
158 let p_pred = self.p + self.q;
159
160 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 pub fn estimate(&self) -> f64 {
170 self.x
171 }
172}
173
174#[derive(Debug, Clone)]
176pub struct LatencyEma {
177 value: f64,
178 alpha: f64,
179 sample_count: u64,
180}
181
182impl LatencyEma {
183 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 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 pub fn current(&self) -> f64 {
205 self.value
206 }
207
208 pub fn count(&self) -> u64 {
210 self.sample_count
211 }
212}
213
214struct PoolState {
216 kalman: KalmanFilter,
217 latency_ema: LatencyEma,
218 current_workers: usize,
219 last_scale_at: Option<Instant>,
220}
221
222#[derive(Debug, Clone, Serialize)]
224pub struct PoolStats {
225 pub current_workers: usize,
227 pub estimated_queue_depth: f64,
229 pub latency_ema_ms: f64,
231 pub total_scale_ups: u64,
233 pub total_scale_downs: u64,
235}
236
237pub 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 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 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 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 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 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); 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 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 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 pub fn observation_count(&self) -> usize {
382 self.observations.load(Ordering::Relaxed)
383 }
384}
385
386pub 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 for _ in 0..20 {
455 kf.update(100.0);
456 }
457 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 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 for _ in 0..5 {
522 pool.evaluate(200, 1000.0).await;
523 }
524 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), 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 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, 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}