use super::{WarmupEntry, WarmupLoader, WarmupReport};
use crate::backend::CacheReader;
use crate::cache::ChainCache;
use crate::error::OxCacheResult;
use crate::i18n::messages::{MSG_DETAIL_WARMUP_TTL_LOOKUP_FAILED, MSG_PANIC_WARMUP_SEMAPHORE, t};
use std::collections::HashSet;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::sync::Semaphore;
const DEFAULT_MAX_VALUE_BYTES: usize = crate::core::constants::MAX_JSON_SIZE;
const DEFAULT_MAX_ENTRIES: usize = 100_000;
#[derive(Clone)]
pub struct Warmup {
chain: Arc<ChainCache>,
concurrency: usize,
max_value_bytes: usize,
max_entries: usize,
}
pub struct WarmupBuilder {
chain: Arc<ChainCache>,
concurrency: usize,
max_value_bytes: usize,
max_entries: usize,
}
impl WarmupBuilder {
pub fn concurrency(mut self, concurrency: usize) -> Self {
self.concurrency = concurrency.max(1);
self
}
pub fn max_value_bytes(mut self, max_value_bytes: usize) -> Self {
self.max_value_bytes = max_value_bytes;
self
}
pub fn max_entries(mut self, max_entries: usize) -> Self {
self.max_entries = max_entries;
self
}
pub fn build(self) -> Warmup {
Warmup {
chain: self.chain,
concurrency: self.concurrency,
max_value_bytes: self.max_value_bytes,
max_entries: self.max_entries,
}
}
}
type BackfillItem = (String, Vec<u8>, Option<std::time::Duration>);
impl Warmup {
pub fn builder(chain: Arc<ChainCache>) -> WarmupBuilder {
WarmupBuilder {
chain,
concurrency: 8,
max_value_bytes: DEFAULT_MAX_VALUE_BYTES,
max_entries: DEFAULT_MAX_ENTRIES,
}
}
pub async fn run(&self, loader: Arc<dyn WarmupLoader>) -> OxCacheResult<WarmupReport> {
let entries = loader.load_hot_keys().await?;
let loader_entries = entries.len();
let mut seen: HashSet<String> = HashSet::with_capacity(entries.len());
let mut deduped_entries: Vec<WarmupEntry> = entries
.into_iter()
.filter(|e| seen.insert(e.key.clone()))
.collect();
let deduped = deduped_entries.len();
let mut report = WarmupReport {
loader_entries,
deduped,
..WarmupReport::default()
};
if self.max_entries > 0 && deduped_entries.len() > self.max_entries {
report.dropped_over_entry_cap = deduped_entries.len() - self.max_entries;
deduped_entries.truncate(self.max_entries);
}
let mut direct: Vec<BackfillItem> = Vec::with_capacity(deduped_entries.len());
let mut key_only: Vec<String> = Vec::new();
for entry in deduped_entries {
match entry.value {
Some(value) => {
if self.max_value_bytes > 0
&& crate::infra::serialization::utils::check_data_size(
&value,
self.max_value_bytes,
"warmup entry value",
)
.is_err()
{
report.dropped_value_too_large += 1;
continue;
}
direct.push((entry.key, value, entry.ttl));
}
None => key_only.push(entry.key),
}
}
let (warmed, mut failures, peak) = self.backfill_batch(direct).await;
report.warmed = warmed;
report.peak_concurrency = peak;
report.failed += failures.len();
report.failures.append(&mut failures);
if !key_only.is_empty() {
let key_refs: Vec<&str> = key_only.iter().map(String::as_str).collect();
let fetched = self.chain.iter_entries(&key_refs).await;
let mut promote_items: Vec<BackfillItem> = Vec::new();
for (key, value) in fetched {
match value {
Some(value) => {
let ttl = match self.chain.ttl(&key).await {
Ok(ttl) => ttl,
Err(e) => {
report.failed += 1;
report.failures.push((
key,
t(
MSG_DETAIL_WARMUP_TTL_LOOKUP_FAILED,
&[("err", e.to_string())],
),
));
continue;
}
};
promote_items.push((key, value, ttl));
}
None => report.missing += 1,
}
}
let (promoted, mut promote_failures, promote_peak) =
self.backfill_batch(promote_items).await;
report.promoted = promoted;
report.peak_concurrency = report.peak_concurrency.max(promote_peak);
report.failed += promote_failures.len();
report.failures.append(&mut promote_failures);
}
Ok(report)
}
async fn backfill_batch(
&self,
items: Vec<BackfillItem>,
) -> (usize, Vec<(String, String)>, usize) {
let mut warmed = 0usize;
let mut failures: Vec<(String, String)> = Vec::new();
let in_flight = Arc::new(AtomicUsize::new(0));
let peak = Arc::new(AtomicUsize::new(0));
let permits = Arc::new(Semaphore::new(self.concurrency.max(1)));
let mut set = tokio::task::JoinSet::new();
for (key, value, ttl) in items {
let permit = permits
.clone()
.acquire_owned()
.await
.unwrap_or_else(|e| panic!("{}: {e:?}", t(MSG_PANIC_WARMUP_SEMAPHORE, &[])));
let chain = self.chain.clone();
let in_flight = in_flight.clone();
let peak = peak.clone();
set.spawn(async move {
let now = in_flight.fetch_add(1, Ordering::Relaxed) + 1;
peak.fetch_max(now, Ordering::Relaxed);
let result = chain
.set(&key, value, ttl)
.await
.map_err(|e| (key, e.to_string()));
in_flight.fetch_sub(1, Ordering::Relaxed);
drop(permit);
result
});
}
while let Some(joined) = set.join_next().await {
match joined {
Ok(Ok(())) => warmed += 1,
Ok(Err((key, err))) => failures.push((key, err)),
Err(join_err) => failures.push(("<join>".to_string(), join_err.to_string())),
}
}
(warmed, failures, peak.load(Ordering::Relaxed))
}
}