use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
use tracing::{debug, warn};
#[derive(Clone)]
pub struct BackpressureMonitor {
sent_count: Arc<AtomicU64>,
consumed_count: Arc<AtomicU64>,
active: Arc<AtomicBool>,
last_adjustment: Arc<parking_lot::Mutex<Instant>>,
}
impl BackpressureMonitor {
pub fn new() -> Self {
Self {
sent_count: Arc::new(AtomicU64::new(0)),
consumed_count: Arc::new(AtomicU64::new(0)),
active: Arc::new(AtomicBool::new(true)),
last_adjustment: Arc::new(parking_lot::Mutex::new(Instant::now())),
}
}
pub fn record_send(&self) {
self.sent_count.fetch_add(1, Ordering::Relaxed);
}
pub fn record_consume(&self) {
self.consumed_count.fetch_add(1, Ordering::Relaxed);
}
pub fn get_lag(&self) -> u64 {
let sent = self.sent_count.load(Ordering::Relaxed);
let consumed = self.consumed_count.load(Ordering::Relaxed);
sent.saturating_sub(consumed)
}
pub fn get_consumption_rate(&self, duration: Duration) -> f64 {
let consumed = self.consumed_count.load(Ordering::Relaxed);
consumed as f64 / duration.as_secs_f64()
}
pub fn stop(&self) {
self.active.store(false, Ordering::Relaxed);
}
pub fn should_apply_backpressure(&self, threshold: u64) -> bool {
self.get_lag() > threshold
}
pub fn recommend_buffer_size(
&self,
config: &crate::runtime::stream_config::StreamConfig,
) -> usize {
if !config.adaptive_buffering {
return config.channel_buffer_size;
}
let lag = self.get_lag();
let current_size = config.channel_buffer_size;
let mut last_adj = self.last_adjustment.lock();
let now = Instant::now();
if now.duration_since(*last_adj) < Duration::from_secs(1) {
return current_size;
}
let new_size = if lag > current_size as u64 * 2 {
debug!(
"High backpressure detected (lag: {}), increasing buffer size",
lag
);
(current_size * 2).min(config.max_buffer_size)
} else if lag < current_size as u64 / 4 {
debug!(
"Low buffer utilization (lag: {}), decreasing buffer size",
lag
);
(current_size / 2).max(config.min_buffer_size)
} else {
current_size
};
if new_size != current_size {
*last_adj = now;
warn!(
"Adjusting buffer size from {} to {} (lag: {})",
current_size, new_size, lag
);
}
new_size
}
}
pub struct BackpressureSender<T> {
inner: mpsc::Sender<T>,
monitor: BackpressureMonitor,
threshold: u64,
}
impl<T> BackpressureSender<T> {
pub fn new(inner: mpsc::Sender<T>, monitor: BackpressureMonitor, threshold: u64) -> Self {
Self {
inner,
monitor,
threshold,
}
}
pub async fn send_with_backpressure(&self, value: T) -> Result<(), mpsc::error::SendError<T>> {
if self.monitor.should_apply_backpressure(self.threshold) {
let lag = self.monitor.get_lag();
debug!("Applying backpressure (lag: {}), waiting before send", lag);
let delay = Duration::from_millis((lag as f64 * 0.1).min(100.0) as u64);
tokio::time::sleep(delay).await;
}
self.monitor.record_send();
self.inner.send(value).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_backpressure_monitor() {
let monitor = BackpressureMonitor::new();
for _ in 0..10 {
monitor.record_send();
}
assert_eq!(monitor.get_lag(), 10);
for _ in 0..5 {
monitor.record_consume();
}
assert_eq!(monitor.get_lag(), 5);
assert!(monitor.should_apply_backpressure(3));
assert!(!monitor.should_apply_backpressure(10));
}
#[tokio::test]
async fn test_consumption_rate() {
let monitor = BackpressureMonitor::new();
for _ in 0..100 {
monitor.record_consume();
}
let rate = monitor.get_consumption_rate(Duration::from_secs(10));
assert_eq!(rate, 10.0); }
}