use std::collections::HashMap;
use std::time::{Duration, Instant};
use parking_lot::RwLock;
use rand::Rng;
#[derive(Debug, Clone)]
pub struct QueryCacheConfig {
pub ttl: Duration,
pub max_entries: usize,
pub enable_null_cache: bool,
pub enable_singleflight: bool,
pub ttl_jitter: f64,
}
impl Default for QueryCacheConfig {
fn default() -> Self {
Self {
ttl: Duration::from_secs(60),
max_entries: 10000,
enable_null_cache: true,
enable_singleflight: true,
ttl_jitter: 0.1,
}
}
}
#[derive(Debug, Clone)]
struct CacheEntry {
data: Vec<u8>,
expires_at: Instant,
#[allow(dead_code)]
is_null: bool,
}
impl CacheEntry {
fn is_expired(&self) -> bool {
Instant::now() > self.expires_at
}
}
#[derive(Debug, Clone, Default)]
struct CacheStats {
hits: u64,
misses: u64,
evictions: u64,
}
pub struct QueryCache {
config: QueryCacheConfig,
entries: RwLock<HashMap<String, CacheEntry>>,
stats: RwLock<CacheStats>,
}
impl QueryCache {
pub fn new(config: QueryCacheConfig) -> Self {
Self {
config,
entries: RwLock::new(HashMap::new()),
stats: RwLock::new(CacheStats::default()),
}
}
pub fn make_key(sql: &str, params: &[&str]) -> String {
let mut key = String::with_capacity(sql.len() + params.len() * 8);
key.push_str(sql);
for p in params {
key.push('|');
key.push_str(p);
}
key
}
pub async fn get_or_query<F, Fut>(
&self,
key: &str,
query_fn: F,
) -> Result<Vec<u8>, QueryCacheError>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<Vec<u8>, QueryCacheError>>,
{
if let Some(entry) = self.entries.read().get(key) {
if !entry.is_expired() {
self.stats.write().hits += 1;
return Ok(entry.data.clone());
}
}
self.stats.write().misses += 1;
let data = query_fn().await?;
self.put(key, data.clone());
Ok(data)
}
fn put(&self, key: &str, data: Vec<u8>) {
let mut entries = self.entries.write();
if entries.len() >= self.config.max_entries {
self.evict_oldest(&mut entries);
}
let ttl = self.jitter_ttl();
let is_null = data.is_empty();
entries.insert(
key.to_string(),
CacheEntry {
data,
expires_at: Instant::now() + ttl,
is_null,
},
);
}
pub fn invalidate(&self, pattern: &str) -> usize {
let mut entries = self.entries.write();
let keys_to_remove: Vec<String> = entries
.keys()
.filter(|k| k.contains(pattern))
.cloned()
.collect();
let count = keys_to_remove.len();
for k in keys_to_remove {
entries.remove(&k);
}
count
}
pub fn clear(&self) {
self.entries.write().clear();
}
pub fn len(&self) -> usize {
self.entries.read().len()
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn hit_rate(&self) -> f64 {
let stats = self.stats.read();
let total = stats.hits + stats.misses;
if total == 0 {
0.0
} else {
stats.hits as f64 / total as f64
}
}
pub fn hits(&self) -> u64 {
self.stats.read().hits
}
pub fn misses(&self) -> u64 {
self.stats.read().misses
}
fn evict_oldest(&self, entries: &mut HashMap<String, CacheEntry>) {
if let Some((oldest_key, _)) = entries
.iter()
.min_by_key(|(_, e)| e.expires_at)
.map(|(k, _)| (k.clone(), ()))
{
entries.remove(&oldest_key);
self.stats.write().evictions += 1;
}
}
fn jitter_ttl(&self) -> Duration {
if self.config.ttl_jitter == 0.0 {
return self.config.ttl;
}
let mut rng = rand::thread_rng();
let jitter = rng.gen_range(-self.config.ttl_jitter..=self.config.ttl_jitter);
let base_ms = self.config.ttl.as_millis() as f64;
let adjusted_ms = base_ms * (1.0 + jitter);
Duration::from_millis(adjusted_ms as u64)
}
}
impl std::fmt::Debug for QueryCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"QueryCache {{ entries: {}, hits: {}, misses: {} }}",
self.len(),
self.hits(),
self.misses()
)
}
}
#[derive(Debug, thiserror::Error)]
pub enum QueryCacheError {
#[error("query failed: {0}")]
QueryFailed(String),
#[error("serialize failed: {0}")]
SerializeFailed(String),
}
#[cfg(test)]
mod tests {
use super::*;
fn make_key(sql: &str, params: &[&str]) -> String {
QueryCache::make_key(sql, params)
}
#[test]
fn test_make_key_consistency() {
let k1 = make_key("SELECT * FROM users WHERE id = ?", &["1"]);
let k2 = make_key("SELECT * FROM users WHERE id = ?", &["1"]);
assert_eq!(k1, k2);
}
#[test]
fn test_make_key_different_params() {
let k1 = make_key("SELECT * FROM users WHERE id = ?", &["1"]);
let k2 = make_key("SELECT * FROM users WHERE id = ?", &["2"]);
assert_ne!(k1, k2);
}
#[test]
fn test_make_key_different_sql() {
let k1 = make_key("SELECT * FROM users", &[]);
let k2 = make_key("SELECT * FROM orders", &[]);
assert_ne!(k1, k2);
}
#[test]
fn test_config_default() {
let config = QueryCacheConfig::default();
assert_eq!(config.ttl, Duration::from_secs(60));
assert_eq!(config.max_entries, 10000);
assert!(config.enable_null_cache);
assert!(config.enable_singleflight);
assert_eq!(config.ttl_jitter, 0.1);
}
#[test]
fn test_cache_entry_expiry() {
let entry = CacheEntry {
data: vec![1, 2, 3],
expires_at: Instant::now() + Duration::from_secs(60),
is_null: false,
};
assert!(!entry.is_expired());
}
#[test]
fn test_cache_entry_expired() {
let entry = CacheEntry {
data: vec![1, 2, 3],
expires_at: Instant::now() - Duration::from_secs(1),
is_null: false,
};
assert!(entry.is_expired());
}
#[test]
fn test_jitter_ttl() {
let config = QueryCacheConfig {
ttl: Duration::from_secs(100),
ttl_jitter: 0.1,
..Default::default()
};
let cache = QueryCache::new(config);
for _ in 0..100 {
let ttl = cache.jitter_ttl();
let ms = ttl.as_millis();
assert!(
(90_000..=110_000).contains(&ms),
"jitter TTL out of range: {ms}ms"
);
}
}
#[test]
fn test_hit_rate_zero() {
let cache = QueryCache::new(QueryCacheConfig::default());
assert_eq!(cache.hit_rate(), 0.0);
}
#[test]
fn test_invalidate() {
let cache = QueryCache::new(QueryCacheConfig::default());
cache.put("users:1", b"data1".to_vec());
cache.put("users:2", b"data2".to_vec());
cache.put("orders:1", b"data3".to_vec());
let removed = cache.invalidate("users");
assert_eq!(removed, 2);
assert_eq!(cache.len(), 1);
}
#[test]
fn test_clear() {
let cache = QueryCache::new(QueryCacheConfig::default());
cache.put("key1", b"data".to_vec());
cache.put("key2", b"data".to_vec());
assert_eq!(cache.len(), 2);
cache.clear();
assert!(cache.is_empty());
}
}