use crate::{Client, ClientError};
use futures::channel::oneshot::{Sender, channel};
use std::{
collections::BTreeMap,
fmt,
sync::{Arc, Mutex},
};
#[derive(Clone, Default)]
pub struct RequestDeduplicationClient<C> {
inner: C,
requests: Arc<Mutex<BTreeMap<DeduplicationKey, Vec<Sender<Response>>>>>,
}
impl<C: fmt::Debug> fmt::Debug for RequestDeduplicationClient<C> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RequestDeduplicationClient")
.field("inner", &self.inner)
.field("requests", &self.requests.lock().unwrap().len())
.finish()
}
}
impl<C> RequestDeduplicationClient<C> {
pub fn new(inner: C) -> Self {
Self {
inner,
requests: Arc::new(Mutex::new(BTreeMap::new())),
}
}
}
impl<C: Client> Client for RequestDeduplicationClient<C> {
async fn get(
&self,
url: &url::Url,
params: &[(&str, &str)],
) -> Result<(u16, std::sync::Arc<String>), ClientError> {
let response = self
.try_get_or_wait(url, params, async {
match self.inner.get(url, params).await {
Ok((status, data)) => Response::String(status, data),
Err(err) => Response::Error(err),
}
})
.await;
match response {
Response::String(status, data) => Ok((status, data)),
Response::Binary(_, _) => panic!(),
Response::Error(client_error) => Err(client_error),
}
}
async fn get_bin(
&self,
url: &url::Url,
params: &[(&str, &str)],
) -> Result<(u16, std::sync::Arc<Vec<u8>>), ClientError> {
let response = self
.try_get_or_wait(url, params, async {
match self.inner.get_bin(url, params).await {
Ok((status, data)) => Response::Binary(status, data),
Err(err) => Response::Error(err),
}
})
.await;
match response {
Response::String(_, _) => panic!(),
Response::Binary(status, data) => Ok((status, data)),
Response::Error(client_error) => Err(client_error),
}
}
}
impl<C: Client> RequestDeduplicationClient<C> {
async fn try_get_or_wait(
&self,
url: &url::Url,
params: &[(&str, &str)],
getter: impl Future<Output = Response>,
) -> Response {
let key = DeduplicationKey {
url: url.clone(),
params: params
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect(),
};
let (rx, request) = {
let mut requests = self.requests.lock().unwrap();
let (tx, rx) = channel::<Response>();
match requests.get_mut(&key) {
Some(ongoing_requests) => {
ongoing_requests.push(tx);
tracing::trace!(
"Deduplicated request to {}. Queue length: {}",
key,
ongoing_requests.len()
);
(rx, None)
}
None => {
tracing::debug!("A fresh request to: {}", key);
requests.insert(key.clone(), vec![tx]);
let request = async move {
let response = getter.await;
let mut requests = self.requests.lock().unwrap();
let senders = requests
.remove(&key)
.expect("Key no longer exist. This is a bug");
tracing::debug!(
"Sending response of {} to {} requesters",
key,
senders.len()
);
for tx in senders {
tx.send(response.clone()).unwrap();
}
};
(rx, Some(request))
}
}
};
if let Some(request) = request {
request.await;
}
rx.await
.expect("Dedup channel closed instead of answered. This is a bug")
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
struct DeduplicationKey {
url: url::Url,
params: Vec<(String, String)>,
}
impl fmt::Display for DeduplicationKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "DeduplicationKey: {}:{:?}", self.url, self.params)
}
}
#[derive(Debug, Clone)]
enum Response {
String(u16, std::sync::Arc<String>),
Binary(u16, std::sync::Arc<Vec<u8>>),
Error(ClientError),
}