use crate::backend::{CacheBackend, CacheConnector, CacheReader, CacheWriter};
use crate::error::{OxCacheError, OxCacheResult};
use crate::i18n::messages::{
MSG_DETAIL_ADAPTIVE_TTL_DIVISOR_MIN, MSG_DETAIL_ADAPTIVE_TTL_MIN_EXCEEDS_MAX,
MSG_DETAIL_ADAPTIVE_TTL_MULTIPLIER_FINITE, MSG_DETAIL_ADAPTIVE_TTL_TRACKED_KEYS_MIN, t,
};
use async_trait::async_trait;
use dashmap::DashMap;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
#[derive(Debug)]
struct AccessState {
freq: u64,
last_access: Instant,
last_extend: Option<Instant>,
}
#[derive(Debug, Clone)]
pub struct AdaptiveTtlConfig {
pub hot_threshold: u64,
pub hot_ttl_multiplier: f64,
pub cold_idle_after: Duration,
pub cold_ttl_divisor: u64,
pub min_ttl: Duration,
pub max_ttl: Duration,
pub adjust_interval: Duration,
pub max_tracked_keys: usize,
}
impl Default for AdaptiveTtlConfig {
fn default() -> Self {
Self {
hot_threshold: 8,
hot_ttl_multiplier: 2.0,
cold_idle_after: Duration::from_secs(300),
cold_ttl_divisor: 4,
min_ttl: Duration::from_secs(1),
max_ttl: Duration::from_secs(3600),
adjust_interval: Duration::from_secs(60),
max_tracked_keys: 65_536,
}
}
}
impl AdaptiveTtlConfig {
pub fn validate(&self) -> OxCacheResult<()> {
if self.min_ttl > self.max_ttl {
return Err(OxCacheError::InvalidInput(t(
MSG_DETAIL_ADAPTIVE_TTL_MIN_EXCEEDS_MAX,
&[
("min", format!("{:?}", self.min_ttl)),
("max", format!("{:?}", self.max_ttl)),
],
)));
}
if !(self.hot_ttl_multiplier.is_finite() && self.hot_ttl_multiplier > 0.0) {
return Err(OxCacheError::InvalidInput(t(
MSG_DETAIL_ADAPTIVE_TTL_MULTIPLIER_FINITE,
&[("value", self.hot_ttl_multiplier.to_string())],
)));
}
if self.cold_ttl_divisor == 0 {
return Err(OxCacheError::InvalidInput(t(
MSG_DETAIL_ADAPTIVE_TTL_DIVISOR_MIN,
&[],
)));
}
if self.max_tracked_keys == 0 {
return Err(OxCacheError::InvalidInput(t(
MSG_DETAIL_ADAPTIVE_TTL_TRACKED_KEYS_MIN,
&[],
)));
}
Ok(())
}
}
#[derive(Debug, Default)]
pub struct AdaptiveTtlStats {
pub hot_extensions: u64,
pub cold_shortenings: u64,
pub failed_adjustments: u64,
pub tracked_keys: u64,
}
pub struct AdaptiveTtlBackend {
inner: Arc<dyn CacheBackend>,
config: AdaptiveTtlConfig,
accesses: DashMap<String, AccessState>,
hot_extensions: AtomicU64,
cold_shortenings: AtomicU64,
failed_adjustments: AtomicU64,
}
impl AdaptiveTtlBackend {
pub fn new(inner: Arc<dyn CacheBackend>, config: AdaptiveTtlConfig) -> OxCacheResult<Self> {
config.validate()?;
Ok(Self {
inner,
config,
accesses: DashMap::new(),
hot_extensions: AtomicU64::new(0),
cold_shortenings: AtomicU64::new(0),
failed_adjustments: AtomicU64::new(0),
})
}
pub fn stats(&self) -> AdaptiveTtlStats {
AdaptiveTtlStats {
hot_extensions: self.hot_extensions.load(Ordering::Relaxed),
cold_shortenings: self.cold_shortenings.load(Ordering::Relaxed),
failed_adjustments: self.failed_adjustments.load(Ordering::Relaxed),
tracked_keys: self.accesses.len() as u64,
}
}
pub fn reset_stats(&self) {
self.accesses.clear();
self.hot_extensions.store(0, Ordering::Relaxed);
self.cold_shortenings.store(0, Ordering::Relaxed);
self.failed_adjustments.store(0, Ordering::Relaxed);
}
fn record_access(&self, key: &str) -> Option<(u64, Instant)> {
let now = Instant::now();
match self.accesses.get_mut(key) {
Some(mut state) => {
state.freq = state.freq.saturating_add(1);
let last = state.last_access;
state.last_access = now;
Some((state.freq, last))
}
None => {
if self.accesses.len() >= self.config.max_tracked_keys {
return None;
}
self.accesses.insert(
key.to_string(),
AccessState {
freq: 1,
last_access: now,
last_extend: None,
},
);
Some((1, now))
}
}
}
fn adjusted_set_ttl(&self, ttl: Duration, freq: u64, last_access: Instant) -> Duration {
let now = Instant::now();
let clamp = |d: Duration| d.clamp(self.config.min_ttl, self.config.max_ttl);
if freq >= self.config.hot_threshold {
let scaled = self.scale_ttl(ttl, self.config.hot_ttl_multiplier);
self.hot_extensions.fetch_add(1, Ordering::Relaxed);
return clamp(scaled);
}
if now.duration_since(last_access) > self.config.cold_idle_after {
let divisor = self.config.cold_ttl_divisor.max(1) as f64;
let shortened = self.scale_ttl(ttl, 1.0 / divisor);
self.cold_shortenings.fetch_add(1, Ordering::Relaxed);
return clamp(shortened);
}
ttl
}
fn scale_ttl(&self, ttl: Duration, factor: f64) -> Duration {
let ms = ttl.as_millis();
let scaled = (ms as f64 * factor).max(1.0);
if scaled >= u64::MAX as f64 {
Duration::MAX
} else {
Duration::from_millis(scaled as u64)
}
}
async fn maybe_extend_hot(&self, key: &str, freq: u64) {
if freq < self.config.hot_threshold {
return;
}
let now = Instant::now();
let should_extend = match self.accesses.get_mut(key) {
Some(mut state) => {
let due = match state.last_extend {
Some(last) => now.duration_since(last) >= self.config.adjust_interval,
None => true,
};
if due {
state.last_extend = Some(now);
}
due
}
None => false,
};
if !should_extend {
return;
}
if let Ok(Some(remaining)) = self.inner.ttl(key).await {
let scaled = self.scale_ttl(remaining, self.config.hot_ttl_multiplier);
let target = scaled.clamp(self.config.min_ttl, self.config.max_ttl);
if target != remaining {
match self.inner.expire(key, target).await {
Ok(_) => {
self.hot_extensions.fetch_add(1, Ordering::Relaxed);
}
Err(_) => {
self.failed_adjustments.fetch_add(1, Ordering::Relaxed);
}
}
}
}
}
}
#[async_trait]
impl CacheReader for AdaptiveTtlBackend {
async fn get(&self, key: &str) -> OxCacheResult<Option<Vec<u8>>> {
let value = self.inner.get(key).await?;
if value.is_some()
&& let Some((freq, _)) = self.record_access(key)
{
self.maybe_extend_hot(key, freq).await;
}
Ok(value)
}
async fn exists(&self, key: &str) -> OxCacheResult<bool> {
self.inner.exists(key).await
}
async fn ttl(&self, key: &str) -> OxCacheResult<Option<Duration>> {
self.inner.ttl(key).await
}
async fn len(&self) -> OxCacheResult<u64> {
self.inner.len().await
}
async fn capacity(&self) -> OxCacheResult<u64> {
self.inner.capacity().await
}
async fn stats(&self) -> OxCacheResult<HashMap<String, String>> {
let mut stats = self.inner.stats().await?;
stats.insert(
"adaptive_tracked_keys".to_string(),
self.accesses.len().to_string(),
);
stats.insert(
"adaptive_hot_extensions".to_string(),
self.hot_extensions.load(Ordering::Relaxed).to_string(),
);
stats.insert(
"adaptive_cold_shortenings".to_string(),
self.cold_shortenings.load(Ordering::Relaxed).to_string(),
);
stats.insert(
"adaptive_failed_adjustments".to_string(),
self.failed_adjustments.load(Ordering::Relaxed).to_string(),
);
Ok(stats)
}
async fn keys(&self, pattern: &str) -> OxCacheResult<Vec<String>> {
self.inner.keys(pattern).await
}
}
#[async_trait]
impl CacheWriter for AdaptiveTtlBackend {
async fn set(
&self,
key: Arc<str>,
value: Arc<Vec<u8>>,
ttl: Option<Duration>,
) -> OxCacheResult<()> {
let Some(ttl) = ttl else {
return self.inner.set(key, value, None).await;
};
let (freq, last_access) = match self.record_access(&key) {
Some(pair) => pair,
None => return self.inner.set(key, value, Some(ttl)).await,
};
let adjusted = self.adjusted_set_ttl(ttl, freq, last_access);
self.inner.set(key, value, Some(adjusted)).await
}
async fn delete(&self, key: &str) -> OxCacheResult<()> {
self.inner.delete(key).await
}
async fn clear(&self) -> OxCacheResult<()> {
self.reset_stats();
self.inner.clear().await
}
async fn expire(&self, key: &str, ttl: Duration) -> OxCacheResult<bool> {
self.inner.expire(key, ttl).await
}
async fn set_many(&self, items: &[crate::backend::CacheSetItem]) -> OxCacheResult<()> {
let mut adjusted: Vec<crate::backend::CacheSetItem> = Vec::with_capacity(items.len());
for item in items {
let (key, value, ttl) = (&item.0, &item.1, &item.2);
let (freq, last_access) = match self.record_access(key) {
Some(pair) => pair,
None => {
adjusted.push((key.clone(), value.clone(), *ttl));
continue;
}
};
match *ttl {
Some(base) => {
let adjusted_ttl = self.adjusted_set_ttl(base, freq, last_access);
adjusted.push((key.clone(), value.clone(), Some(adjusted_ttl)));
}
None => adjusted.push((key.clone(), value.clone(), None)),
}
}
self.inner.set_many(&adjusted).await
}
async fn delete_many(&self, keys: &[String]) -> OxCacheResult<()> {
self.inner.delete_many(keys).await
}
}
#[async_trait]
impl CacheConnector for AdaptiveTtlBackend {
async fn health_check(&self) -> OxCacheResult<()> {
self.inner.health_check().await
}
async fn shutdown(&self) {
self.inner.shutdown().await;
}
fn backend_kind(&self) -> crate::backend::BackendKind {
self.inner.backend_kind()
}
}
#[cfg(all(test, feature = "adaptive-ttl"))]
mod tests {
use super::*;
use crate::backend::{MockBackend, MokaMemoryBackend};
use crate::error::OxCacheError;
fn hot_config() -> AdaptiveTtlConfig {
AdaptiveTtlConfig {
hot_threshold: 2,
hot_ttl_multiplier: 100.0,
min_ttl: Duration::from_secs(1),
max_ttl: Duration::from_secs(5),
..AdaptiveTtlConfig::default()
}
}
fn cold_config() -> AdaptiveTtlConfig {
AdaptiveTtlConfig {
hot_threshold: 1_000,
cold_idle_after: Duration::ZERO,
cold_ttl_divisor: 1_000,
min_ttl: Duration::from_secs(7),
max_ttl: Duration::from_secs(3_600),
..AdaptiveTtlConfig::default()
}
}
async fn store(backend: &AdaptiveTtlBackend, key: &str, ttl: Option<Duration>) {
backend
.set(Arc::from(key), Arc::new(b"v".to_vec()), ttl)
.await
.unwrap();
}
#[tokio::test]
async fn hot_extension_clamped_to_max_ttl() {
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), hot_config())
.unwrap();
store(&backend, "hot-key", Some(Duration::from_secs(60))).await;
store(&backend, "hot-key", Some(Duration::from_secs(60))).await;
let ttl = backend.ttl("hot-key").await.unwrap().unwrap();
assert!(
ttl <= Duration::from_secs(5),
"adjusted ttl must clamp to max_ttl=5s, got {ttl:?}"
);
assert_eq!(backend.stats().hot_extensions, 1, "one hot adjustment");
}
#[tokio::test]
async fn cold_shortening_clamped_to_min_ttl() {
let backend = AdaptiveTtlBackend::new(
Arc::new(MokaMemoryBackend::builder().build()),
cold_config(),
)
.unwrap();
store(&backend, "cold-key", Some(Duration::from_secs(60))).await;
store(&backend, "cold-key", Some(Duration::from_secs(60))).await;
let ttl = backend.ttl("cold-key").await.unwrap().unwrap();
assert!(
ttl >= Duration::from_millis(6_900),
"adjusted ttl must clamp to min_ttl=7s, got {ttl:?}"
);
assert_eq!(backend.stats().cold_shortenings, 2);
}
#[tokio::test]
async fn reset_stats_clears_everything() {
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), hot_config())
.unwrap();
store(&backend, "k1", Some(Duration::from_secs(60))).await;
store(&backend, "k1", Some(Duration::from_secs(60))).await;
assert!(backend.stats().tracked_keys > 0);
assert!(backend.stats().hot_extensions > 0);
backend.reset_stats();
let stats = backend.stats();
assert_eq!(stats.tracked_keys, 0);
assert_eq!(stats.hot_extensions, 0);
assert_eq!(stats.cold_shortenings, 0);
store(&backend, "k1", Some(Duration::from_secs(60))).await;
let ttl = backend.ttl("k1").await.unwrap().unwrap();
assert!(
ttl > Duration::from_secs(50),
"after reset the key must be treated as plain again: {ttl:?}"
);
}
#[tokio::test]
async fn get_hits_extend_hot_entry_ttl() {
let config = AdaptiveTtlConfig {
hot_threshold: 2,
hot_ttl_multiplier: 10.0,
adjust_interval: Duration::from_secs(3600),
min_ttl: Duration::from_secs(1),
max_ttl: Duration::from_secs(30),
..AdaptiveTtlConfig::default()
};
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), config)
.unwrap();
store(&backend, "get-hot", Some(Duration::from_secs(2))).await;
backend.get("get-hot").await.unwrap();
assert_eq!(backend.stats().hot_extensions, 1);
let ttl = backend.ttl("get-hot").await.unwrap().unwrap();
assert!(
ttl > Duration::from_secs(15),
"extension must land near 20s, got {ttl:?}"
);
backend.get("get-hot").await.unwrap();
assert_eq!(backend.stats().hot_extensions, 1);
}
#[tokio::test]
async fn none_ttl_passes_through_unadjusted() {
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), hot_config())
.unwrap();
store(&backend, "eternal", None).await;
store(&backend, "eternal", None).await;
assert_eq!(backend.stats().hot_extensions, 0);
assert_eq!(backend.ttl("eternal").await.unwrap(), None);
}
#[tokio::test]
async fn tracked_keys_cap_prevents_growth() {
let config = AdaptiveTtlConfig {
max_tracked_keys: 1,
..hot_config()
};
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), config)
.unwrap();
store(&backend, "tracked", Some(Duration::from_secs(60))).await;
store(&backend, "untracked", Some(Duration::from_secs(60))).await;
assert_eq!(backend.stats().tracked_keys, 1);
assert_eq!(backend.stats().hot_extensions, 0);
}
#[tokio::test]
async fn plain_key_passes_base_ttl_through() {
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), hot_config())
.unwrap();
store(&backend, "plain", Some(Duration::from_secs(60))).await;
let ttl = backend.ttl("plain").await.unwrap().unwrap();
assert!(
ttl > Duration::from_secs(50),
"plain key must keep base ttl, got {ttl:?}"
);
assert_eq!(backend.stats().hot_extensions, 0);
}
#[tokio::test]
async fn rejects_min_ttl_above_max_ttl_at_construction() {
let config = AdaptiveTtlConfig {
min_ttl: Duration::from_secs(30),
max_ttl: Duration::from_secs(1),
..hot_config()
};
let err =
match AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), config) {
Err(e) => e,
Ok(_) => panic!("min_ttl > max_ttl must be rejected at construction"),
};
assert!(
matches!(
&err,
OxCacheError::InvalidInput(m) if m.contains("min_ttl") && m.contains("max_ttl")
),
"rejection must name both bounds: {err:?}"
);
}
#[test]
fn config_validate_rejects_degenerate_numeric_bounds() {
for bad in [0.0, -2.0, f64::NAN, f64::INFINITY] {
let config = AdaptiveTtlConfig {
hot_ttl_multiplier: bad,
..hot_config()
};
let err = config
.validate()
.expect_err("non-finite/non-positive multiplier must be rejected");
assert!(
matches!(&err, OxCacheError::InvalidInput(m) if m.contains("hot_ttl_multiplier")),
"multiplier {bad}: {err:?}"
);
}
let config = AdaptiveTtlConfig {
cold_ttl_divisor: 0,
..hot_config()
};
let err = config
.validate()
.expect_err("cold_ttl_divisor = 0 must be rejected");
assert!(
matches!(&err, OxCacheError::InvalidInput(m) if m.contains("cold_ttl_divisor")),
"{err:?}"
);
}
#[test]
fn config_validate_rejects_zero_tracked_capacity() {
let config = AdaptiveTtlConfig {
max_tracked_keys: 0,
..hot_config()
};
let err = config
.validate()
.expect_err("max_tracked_keys = 0 must be rejected");
assert!(
matches!(&err, OxCacheError::InvalidInput(m) if m.contains("max_tracked_keys")),
"{err:?}"
);
}
#[tokio::test]
async fn expire_fault_keeps_hit_and_counts_failure() {
let backend = AdaptiveTtlBackend::new(
Arc::new(MockBackend::new("adaptive-ttl-fault", 50, false).with_fail_expire()),
hot_config(),
)
.unwrap();
store(&backend, "fault", Some(Duration::from_secs(60))).await;
let hit = backend.get("fault").await.unwrap();
assert_eq!(
hit.as_deref(),
Some(b"v".as_slice()),
"expire failure must not block the hit"
);
let stats = backend.stats();
assert_eq!(
stats.failed_adjustments, 1,
"failure must be observable in stats"
);
assert_eq!(
stats.hot_extensions, 0,
"a failed adjustment must not count as an extension"
);
}
#[tokio::test]
async fn get_hits_clamp_down_over_max_ttl() {
let backend =
AdaptiveTtlBackend::new(Arc::new(MokaMemoryBackend::builder().build()), hot_config())
.unwrap();
store(&backend, "over-max", Some(Duration::from_secs(60))).await;
backend.get("over-max").await.unwrap();
let ttl = backend.ttl("over-max").await.unwrap().unwrap();
assert!(
ttl <= Duration::from_secs(5),
"hot key ttl must clamp down to max_ttl, got {ttl:?}"
);
}
}