use std::sync::Arc;
use std::time::Duration;
use anyhow::Result;
use async_lock::{Barrier, Mutex};
use smol::Task;
use tracing::{trace, warn};
use crate::codec::message::RequestInfo;
use crate::specs::message::Message;
use crate::{cache, config};
pub struct CacheFetch {
pub request_info: RequestInfo,
pub result_barrier: Arc<Barrier>,
pub result: Arc<Mutex<Option<Result<Option<Message>>>>>,
}
pub struct CacheStore {
pub request_info: RequestInfo,
pub response: Message,
}
pub enum CacheMsg {
Fetch(CacheFetch),
Store(CacheStore),
}
pub fn start_cache(config: &config::Config) -> Result<(async_channel::Sender<CacheMsg>, Task<()>)> {
let (cache_tx, cache_rx): (
async_channel::Sender<CacheMsg>,
async_channel::Receiver<CacheMsg>,
) = async_channel::bounded(32);
let redis_url = config.redis.trim();
let mut cache: Box<dyn cache::DnsCache + Send> = if redis_url.is_empty() {
Box::new(cache::memory::Cache::new(config.cache_size))
} else {
Box::new(cache::redis::Cache::new(
&redis_url,
Duration::from_millis(1000),
)?)
};
let task = smol::spawn(async move {
trace!("Cache task waiting for requests");
while let Ok(msg) = cache_rx.recv().await {
match msg {
CacheMsg::Fetch(fetch) => {
trace!("Cache fetch: {}", fetch.request_info.name);
let result = cache.fetch(fetch.request_info).await;
fetch.result.lock().await.replace(result);
fetch.result_barrier.wait().await;
}
CacheMsg::Store(store) => {
trace!("Cache store: {}", store.request_info.name);
if let Err(e) = cache.store(store.request_info, store.response).await {
warn!("Cache store failed: {}", e);
}
}
}
}
});
Ok((cache_tx, task))
}