use std::sync::{Arc, Weak};
use std::time::Duration;
use async_trait::async_trait;
use dashmap::DashMap;
use crate::error::Result;
use crate::time::{Clock, Timestamp};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RecoveryAction {
Set,
Remove,
Expire,
}
#[derive(Debug, Clone)]
pub struct RecoveryItem {
pub key: Arc<str>,
pub action: RecoveryAction,
pub timestamp: Timestamp,
pub expires_at: Timestamp,
pub remaining_retries: Option<u32>,
}
#[derive(Debug, Clone)]
pub struct RecoveryConfig {
pub enabled: bool,
pub delay: Duration,
pub max_items: Option<usize>,
pub max_retries: Option<u32>,
}
impl Default for RecoveryConfig {
fn default() -> Self {
Self {
enabled: true,
delay: Duration::from_secs(2),
max_items: None,
max_retries: None,
}
}
}
#[async_trait]
pub trait RecoveryExecutor: Send + Sync {
async fn replay(&self, item: &RecoveryItem) -> Result<()>;
}
pub struct AutoRecoveryService {
config: RecoveryConfig,
clock: Arc<dyn Clock>,
queue: DashMap<Arc<str>, RecoveryItem>,
executor: std::sync::OnceLock<Weak<dyn RecoveryExecutor>>,
}
impl AutoRecoveryService {
#[must_use]
pub fn new(config: RecoveryConfig, clock: Arc<dyn Clock>) -> Arc<Self> {
Arc::new(Self {
config,
clock,
queue: DashMap::new(),
executor: std::sync::OnceLock::new(),
})
}
pub fn set_executor(&self, executor: Weak<dyn RecoveryExecutor>) {
let _ = self.executor.set(executor);
}
#[must_use]
pub fn len(&self) -> usize {
self.queue.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.queue.is_empty()
}
pub fn enqueue(&self, item: RecoveryItem) {
if !self.config.enabled {
return;
}
match self.queue.get(&item.key) {
Some(existing) if existing.timestamp >= item.timestamp => return,
_ => {}
}
if let Some(max) = self.config.max_items
&& self.queue.len() >= max
&& !self.queue.contains_key(&item.key)
{
self.evict_one_before(item.expires_at);
if self.queue.len() >= max {
return; }
}
let mut item = item;
if item.remaining_retries.is_none() {
item.remaining_retries = self.config.max_retries;
}
self.queue.insert(Arc::clone(&item.key), item);
}
fn evict_one_before(&self, bound: Timestamp) {
let victim = self
.queue
.iter()
.min_by_key(|e| e.value().expires_at.ticks())
.map(|e| (e.key().clone(), e.value().expires_at));
if let Some((key, expires_at)) = victim
&& expires_at < bound
{
self.queue.remove(&key);
}
}
pub fn spawn(self: &Arc<Self>) {
if !self.config.enabled {
return;
}
let service = Arc::clone(self);
let delay = self.config.delay;
tokio::spawn(async move {
let mut ticker = tokio::time::interval(delay.max(Duration::from_millis(50)));
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
ticker.tick().await;
let Some(executor) = service.executor.get().and_then(Weak::upgrade) else {
if service.executor.get().is_some() {
break;
}
continue;
};
service.drain_once(executor.as_ref()).await;
}
});
}
pub async fn drain_once(&self, executor: &dyn RecoveryExecutor) {
let now = self.clock.now();
let keys: Vec<Arc<str>> = self.queue.iter().map(|e| e.key().clone()).collect();
for key in keys {
let Some(item) = self.queue.get(&key).map(|e| e.value().clone()) else {
continue;
};
if now >= item.expires_at {
self.queue.remove(&key);
continue;
}
match executor.replay(&item).await {
Ok(()) => {
self.queue.remove(&key);
}
Err(_) => self.record_failure(&key),
}
}
}
fn record_failure(&self, key: &Arc<str>) {
let mut drop_it = false;
if let Some(mut entry) = self.queue.get_mut(key) {
match entry.remaining_retries {
Some(0) => drop_it = true,
Some(n) => entry.remaining_retries = Some(n - 1),
None => {}
}
}
if drop_it {
self.queue.remove(key);
}
}
}