use crate::core::{Error, Result};
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
use tokio::time::sleep;
use tracing::{info, warn};
pub struct StreamReconnectionManager {
attempts: AtomicU32,
max_attempts: u32,
base_delay: Duration,
last_success: RwLock<Option<Instant>>,
health_metrics: StreamHealthMetrics,
}
#[derive(Debug)]
struct StreamHealthMetrics {
successful_connections: AtomicU64,
failed_connections: AtomicU64,
avg_connection_duration: RwLock<Duration>,
last_error: RwLock<Option<Instant>>,
}
impl StreamReconnectionManager {
pub fn new(max_attempts: u32, base_delay: Duration) -> Self {
Self {
attempts: AtomicU32::new(0),
max_attempts,
base_delay,
last_success: RwLock::new(None),
health_metrics: StreamHealthMetrics {
successful_connections: AtomicU64::new(0),
failed_connections: AtomicU64::new(0),
avg_connection_duration: RwLock::new(Duration::from_secs(0)),
last_error: RwLock::new(None),
},
}
}
pub async fn reconnect<F, Fut, T>(&self, mut connect_fn: F) -> Result<T>
where
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Result<T>>,
{
let current_attempt = self.attempts.fetch_add(1, Ordering::SeqCst);
if current_attempt >= self.max_attempts {
self.attempts.store(0, Ordering::SeqCst);
return Err(Error::StreamClosed);
}
let delay = self.calculate_backoff(current_attempt);
info!(
"Attempting stream reconnection {} of {}, delay: {:?}",
current_attempt + 1,
self.max_attempts,
delay
);
sleep(delay).await;
match connect_fn().await {
Ok(result) => {
self.on_success().await;
Ok(result)
}
Err(e) => {
self.on_failure().await;
Err(e)
}
}
}
fn calculate_backoff(&self, attempt: u32) -> Duration {
let exponential = self.base_delay.as_millis() as u64 * 2u64.pow(attempt);
let max_delay = Duration::from_secs(60);
let delay = Duration::from_millis(exponential.min(max_delay.as_millis() as u64));
let jitter = (rand::random::<f64>() - 0.5) * 0.5;
let jittered_millis = (delay.as_millis() as f64 * (1.0 + jitter)) as u64;
Duration::from_millis(jittered_millis)
}
async fn on_success(&self) {
self.attempts.store(0, Ordering::SeqCst);
*self.last_success.write().await = Some(Instant::now());
self.health_metrics
.successful_connections
.fetch_add(1, Ordering::SeqCst);
}
async fn on_failure(&self) {
self.health_metrics
.failed_connections
.fetch_add(1, Ordering::SeqCst);
*self.health_metrics.last_error.write().await = Some(Instant::now());
}
pub async fn health_status(&self) -> StreamHealthStatus {
let success_count = self
.health_metrics
.successful_connections
.load(Ordering::SeqCst);
let failure_count = self
.health_metrics
.failed_connections
.load(Ordering::SeqCst);
let total = success_count + failure_count;
let success_rate = if total > 0 {
(success_count as f64 / total as f64) * 100.0
} else {
0.0
};
StreamHealthStatus {
success_rate,
total_connections: total,
current_attempts: self.attempts.load(Ordering::SeqCst),
last_success: *self.last_success.read().await,
last_error: *self.health_metrics.last_error.read().await,
}
}
}
#[derive(Debug)]
pub struct StreamHealthStatus {
pub success_rate: f64,
pub total_connections: u64,
pub current_attempts: u32,
pub last_success: Option<Instant>,
pub last_error: Option<Instant>,
}
pub struct CircuitBreaker {
state: Arc<RwLock<CircuitState>>,
failure_threshold: u32,
success_threshold: u32,
open_timeout: Duration,
failure_count: AtomicU32,
success_count: AtomicU32,
last_state_change: RwLock<Instant>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
impl CircuitBreaker {
pub fn new(failure_threshold: u32, success_threshold: u32, open_timeout: Duration) -> Self {
Self {
state: Arc::new(RwLock::new(CircuitState::Closed)),
failure_threshold,
success_threshold,
open_timeout,
failure_count: AtomicU32::new(0),
success_count: AtomicU32::new(0),
last_state_change: RwLock::new(Instant::now()),
}
}
pub async fn execute<F, Fut, T>(&self, operation: F) -> Result<T>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<T>>,
{
self.check_state_transition().await;
let current_state = self.state.read().await.clone();
match current_state {
CircuitState::Open => {
warn!("Circuit breaker is open, rejecting request");
Err(Error::ProcessError(
"Service unavailable - circuit breaker is open".to_string(),
))
}
CircuitState::Closed | CircuitState::HalfOpen => match operation().await {
Ok(result) => {
self.on_success().await;
Ok(result)
}
Err(e) => {
self.on_failure().await;
Err(e)
}
},
}
}
async fn check_state_transition(&self) {
let current_state = self.state.read().await.clone();
let last_change = *self.last_state_change.read().await;
match current_state {
CircuitState::Open => {
if last_change.elapsed() >= self.open_timeout {
self.transition_to(CircuitState::HalfOpen).await;
}
}
_ => {}
}
}
async fn on_success(&self) {
let current_state = self.state.read().await.clone();
match current_state {
CircuitState::HalfOpen => {
let count = self.success_count.fetch_add(1, Ordering::SeqCst) + 1;
if count >= self.success_threshold {
self.transition_to(CircuitState::Closed).await;
}
}
CircuitState::Closed => {
self.failure_count.store(0, Ordering::SeqCst);
}
_ => {}
}
}
async fn on_failure(&self) {
let current_state = self.state.read().await.clone();
match current_state {
CircuitState::Closed => {
let count = self.failure_count.fetch_add(1, Ordering::SeqCst) + 1;
if count >= self.failure_threshold {
self.transition_to(CircuitState::Open).await;
}
}
CircuitState::HalfOpen => {
self.transition_to(CircuitState::Open).await;
}
_ => {}
}
}
async fn transition_to(&self, new_state: CircuitState) {
let mut state = self.state.write().await;
if *state != new_state {
info!(
"Circuit breaker transitioning from {:?} to {:?}",
*state, new_state
);
*state = new_state;
*self.last_state_change.write().await = Instant::now();
self.failure_count.store(0, Ordering::SeqCst);
self.success_count.store(0, Ordering::SeqCst);
}
}
pub async fn current_state(&self) -> CircuitState {
self.state.read().await.clone()
}
}
pub struct TokenBucketRateLimiter {
capacity: u32,
tokens: Arc<RwLock<f64>>,
refill_rate: f64,
last_refill: Arc<RwLock<Instant>>,
}
impl TokenBucketRateLimiter {
pub fn new(capacity: u32, refill_rate: f64) -> Self {
Self {
capacity,
tokens: Arc::new(RwLock::new(capacity as f64)),
refill_rate,
last_refill: Arc::new(RwLock::new(Instant::now())),
}
}
pub async fn try_acquire(&self, tokens_needed: u32) -> Result<()> {
self.refill_tokens().await;
let mut tokens = self.tokens.write().await;
if *tokens >= tokens_needed as f64 {
*tokens -= tokens_needed as f64;
Ok(())
} else {
let _wait_time = self.calculate_wait_time(tokens_needed, *tokens);
Err(Error::RateLimitExceeded)
}
}
async fn refill_tokens(&self) {
let mut last_refill = self.last_refill.write().await;
let elapsed = last_refill.elapsed();
let tokens_to_add = elapsed.as_secs_f64() * self.refill_rate;
if tokens_to_add > 0.0 {
let mut tokens = self.tokens.write().await;
*tokens = (*tokens + tokens_to_add).min(self.capacity as f64);
*last_refill = Instant::now();
}
}
fn calculate_wait_time(&self, tokens_needed: u32, current_tokens: f64) -> Duration {
let tokens_deficit = tokens_needed as f64 - current_tokens;
let seconds_to_wait = tokens_deficit / self.refill_rate;
Duration::from_secs_f64(seconds_to_wait)
}
pub async fn available_tokens(&self) -> f64 {
self.refill_tokens().await;
*self.tokens.read().await
}
}
pub struct PartialResultRecovery {
buffer: Arc<RwLock<Vec<String>>>,
last_checkpoint: Arc<RwLock<Option<usize>>>,
max_buffer_size: usize,
}
impl PartialResultRecovery {
pub fn new(max_buffer_size: usize) -> Self {
Self {
buffer: Arc::new(RwLock::new(Vec::new())),
last_checkpoint: Arc::new(RwLock::new(None)),
max_buffer_size,
}
}
pub async fn save_partial(&self, data: String) -> Result<()> {
let mut buffer = self.buffer.write().await;
if buffer.len() >= self.max_buffer_size {
buffer.remove(0); }
buffer.push(data);
*self.last_checkpoint.write().await = Some(buffer.len());
Ok(())
}
pub async fn recover(&self) -> Vec<String> {
self.buffer.read().await.clone()
}
pub async fn last_checkpoint(&self) -> Option<usize> {
*self.last_checkpoint.read().await
}
pub async fn clear(&self) {
self.buffer.write().await.clear();
*self.last_checkpoint.write().await = None;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
#[ignore] async fn test_stream_reconnection_manager() {
let manager = StreamReconnectionManager::new(3, Duration::from_millis(100));
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
let attempt_count = Arc::new(AtomicUsize::new(0));
let result = manager
.reconnect(|| {
let count = Arc::clone(&attempt_count);
async move {
let current = count.fetch_add(1, Ordering::SeqCst) + 1;
if current < 3 {
Err(Error::StreamClosed)
} else {
Ok("success")
}
}
})
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), "success");
let health = manager.health_status().await;
assert_eq!(health.total_connections, 3);
assert!((health.success_rate - 33.33333333333333).abs() < 0.01);
}
#[tokio::test]
async fn test_circuit_breaker() {
let breaker = CircuitBreaker::new(2, 2, Duration::from_secs(1));
assert_eq!(breaker.current_state().await, CircuitState::Closed);
let _ = breaker
.execute(|| async { Err::<(), _>(Error::ProcessError("fail".to_string())) })
.await;
let _ = breaker
.execute(|| async { Err::<(), _>(Error::ProcessError("fail".to_string())) })
.await;
assert_eq!(breaker.current_state().await, CircuitState::Open);
let result = breaker.execute(|| async { Ok::<_, Error>("test") }).await;
assert!(result.is_err());
sleep(Duration::from_secs(2)).await;
let _ = breaker
.execute(|| async { Ok::<_, Error>("success") })
.await;
assert_eq!(breaker.current_state().await, CircuitState::HalfOpen);
}
#[tokio::test]
async fn test_token_bucket_rate_limiter() {
let limiter = TokenBucketRateLimiter::new(10, 5.0);
assert_eq!(limiter.available_tokens().await as u32, 10);
assert!(limiter.try_acquire(5).await.is_ok());
assert_eq!(limiter.available_tokens().await as u32, 5);
assert!(limiter.try_acquire(10).await.is_err());
sleep(Duration::from_secs(1)).await;
let tokens = limiter.available_tokens().await;
assert!(tokens > 5.0 && tokens <= 10.0);
}
#[tokio::test]
async fn test_partial_result_recovery() {
let recovery = PartialResultRecovery::new(3);
recovery.save_partial("chunk1".to_string()).await.unwrap();
recovery.save_partial("chunk2".to_string()).await.unwrap();
recovery.save_partial("chunk3".to_string()).await.unwrap();
let recovered = recovery.recover().await;
assert_eq!(recovered.len(), 3);
assert_eq!(recovered[0], "chunk1");
assert_eq!(recovered[2], "chunk3");
recovery.save_partial("chunk4".to_string()).await.unwrap();
let recovered = recovery.recover().await;
assert_eq!(recovered.len(), 3);
assert_eq!(recovered[0], "chunk2"); }
}