use dashmap::DashMap;
use ferrox_errors::AppError;
use std::sync::Arc;
use tokio::sync::broadcast;
use tracing::{info, debug};
#[derive(Clone)]
pub struct Singleflight<T> {
in_flight: Arc<DashMap<String, broadcast::Sender<Result<T, String>>>>,
}
impl<T: Clone + Send + Sync + 'static> Singleflight<T> {
pub fn new() -> Self {
Self {
in_flight: Arc::new(DashMap::new()),
}
}
pub async fn execute<F, Fut>(&self, key: &str, fut: F) -> Result<T, AppError>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<T, AppError>> + Send + 'static,
{
let rx = {
if let Some(tx) = self.in_flight.get(key) {
Some(tx.subscribe())
} else {
let (tx, _rx) = broadcast::channel(1);
self.in_flight.insert(key.to_string(), tx);
None
}
};
if let Some(mut receiver) = rx {
debug!("Singleflight: Suspending execution. Waiting for in-flight result for key: {}", key);
return match receiver.recv().await {
Ok(Ok(val)) => Ok(val),
Ok(Err(e)) => Err(AppError::InternalError(e)),
Err(_) => Err(AppError::InternalError("Singleflight sender dropped".into())),
};
}
info!("Singleflight: Primary execution started for key: {}", key);
let result = fut().await;
if let Some((_, tx)) = self.in_flight.remove(key) {
let broadcast_payload = match &result {
Ok(val) => Ok(val.clone()),
Err(e) => Err(format!("{:?}", e)),
};
let _ = tx.send(broadcast_payload);
}
result
}
}