use dashmap::DashMap;
use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Instant;
use super::globals::ScriptNeeds;
pub struct ScriptCache {
entries: DashMap<String, CacheEntry>,
needs_cache: DashMap<String, ScriptNeeds>,
max_size: usize,
hits: AtomicU64,
misses: AtomicU64,
evictions: AtomicU64,
}
struct CacheEntry {
bytecode: Arc<Vec<u8>>,
last_access: Instant,
access_count: AtomicUsize,
code_hash: u64,
}
impl ScriptCache {
pub fn new(max_size: usize) -> Self {
Self {
entries: DashMap::with_capacity(max_size),
needs_cache: DashMap::new(),
max_size,
hits: AtomicU64::new(0),
misses: AtomicU64::new(0),
evictions: AtomicU64::new(0),
}
}
pub fn with_default_size() -> Self {
Self::new(1000)
}
pub fn cache_key(script_key: &str, code: &str) -> String {
let code_hash = Self::hash_code(code);
format!("{}:{:016x}", script_key, code_hash)
}
fn hash_code(code: &str) -> u64 {
let mut hasher = DefaultHasher::new();
code.hash(&mut hasher);
hasher.finish()
}
pub fn get_or_compile<F>(
&self,
script_key: &str,
code: &str,
compile_fn: F,
) -> Result<Vec<u8>, mlua::Error>
where
F: FnOnce(&str) -> Result<Vec<u8>, mlua::Error>,
{
let code_hash = Self::hash_code(code);
let key = format!("{}:{:016x}", script_key, code_hash);
if let Some(mut entry) = self.entries.get_mut(&key) {
if entry.code_hash == code_hash {
entry.last_access = Instant::now();
entry.access_count.fetch_add(1, Ordering::Relaxed);
self.hits.fetch_add(1, Ordering::Relaxed);
return Ok((*entry.bytecode).clone());
}
}
self.misses.fetch_add(1, Ordering::Relaxed);
let bytecode = compile_fn(code)?;
if self.entries.len() >= self.max_size {
self.evict_lru();
}
let bytecode_arc = Arc::new(bytecode.clone());
self.entries.insert(
key,
CacheEntry {
bytecode: bytecode_arc,
last_access: Instant::now(),
access_count: AtomicUsize::new(1),
code_hash,
},
);
Ok(bytecode)
}
pub fn get(&self, script_key: &str, code: &str) -> Option<Vec<u8>> {
let code_hash = Self::hash_code(code);
let key = format!("{}:{:016x}", script_key, code_hash);
if let Some(mut entry) = self.entries.get_mut(&key) {
if entry.code_hash == code_hash {
entry.last_access = Instant::now();
entry.access_count.fetch_add(1, Ordering::Relaxed);
self.hits.fetch_add(1, Ordering::Relaxed);
return Some((*entry.bytecode).clone());
}
}
self.misses.fetch_add(1, Ordering::Relaxed);
None
}
pub fn insert(&self, script_key: &str, code: &str, bytecode: Vec<u8>) {
let code_hash = Self::hash_code(code);
let key = format!("{}:{:016x}", script_key, code_hash);
if self.entries.len() >= self.max_size {
self.evict_lru();
}
self.entries.insert(
key,
CacheEntry {
bytecode: Arc::new(bytecode),
last_access: Instant::now(),
access_count: AtomicUsize::new(1),
code_hash,
},
);
}
pub fn get_or_analyze_needs(&self, script_key: &str, code: &str) -> ScriptNeeds {
if let Some(needs) = self.needs_cache.get(script_key) {
return *needs;
}
let needs = ScriptNeeds::analyze(code);
self.needs_cache.insert(script_key.to_string(), needs);
needs
}
pub fn invalidate(&self, script_key: &str) {
self.entries
.retain(|k, _| !k.starts_with(&format!("{}:", script_key)));
self.needs_cache.remove(script_key);
}
pub fn clear(&self) {
self.entries.clear();
self.needs_cache.clear();
}
fn evict_lru(&self) {
let mut oldest_key: Option<String> = None;
let mut oldest_time = Instant::now();
for entry in self.entries.iter() {
if entry.last_access < oldest_time {
oldest_time = entry.last_access;
oldest_key = Some(entry.key().clone());
}
}
if let Some(key) = oldest_key {
self.entries.remove(&key);
self.evictions.fetch_add(1, Ordering::Relaxed);
}
}
pub fn stats(&self) -> CacheStats {
CacheStats {
entries: self.entries.len(),
max_size: self.max_size,
hits: self.hits.load(Ordering::Relaxed),
misses: self.misses.load(Ordering::Relaxed),
evictions: self.evictions.load(Ordering::Relaxed),
}
}
}
#[derive(Debug, Clone)]
pub struct CacheStats {
pub entries: usize,
pub max_size: usize,
pub hits: u64,
pub misses: u64,
pub evictions: u64,
}
impl CacheStats {
pub fn hit_rate(&self) -> f64 {
let total = self.hits + self.misses;
if total == 0 {
0.0
} else {
(self.hits as f64 / total as f64) * 100.0
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use mlua::Lua;
#[test]
fn test_cache_creation() {
let cache = ScriptCache::new(100);
assert_eq!(cache.max_size, 100);
let stats = cache.stats();
assert_eq!(stats.entries, 0);
assert_eq!(stats.hits, 0);
assert_eq!(stats.misses, 0);
}
#[test]
fn test_cache_key_generation() {
let key1 = ScriptCache::cache_key("script1", "return 1");
let key2 = ScriptCache::cache_key("script1", "return 2");
let key3 = ScriptCache::cache_key("script1", "return 1");
assert_ne!(key1, key2); assert_eq!(key1, key3); }
#[test]
fn test_cache_hit_miss() {
let cache = ScriptCache::new(10);
let lua = Lua::new();
let result = cache.get_or_compile("test", "return 1 + 1", |code| {
let chunk = lua.load(code);
let func = chunk.into_function()?;
Ok(func.dump(false))
});
assert!(result.is_ok());
let stats = cache.stats();
assert_eq!(stats.misses, 1);
assert_eq!(stats.hits, 0);
let result = cache.get_or_compile("test", "return 1 + 1", |_| {
panic!("Should not compile again!");
});
assert!(result.is_ok());
let stats = cache.stats();
assert_eq!(stats.misses, 1);
assert_eq!(stats.hits, 1);
}
#[test]
fn test_cache_invalidation() {
let cache = ScriptCache::new(10);
cache.insert("script1", "code1", vec![1, 2, 3]);
cache.insert("script1", "code2", vec![4, 5, 6]);
cache.insert("script2", "code1", vec![7, 8, 9]);
assert_eq!(cache.stats().entries, 3);
cache.invalidate("script1");
assert_eq!(cache.stats().entries, 1);
assert!(cache.get("script2", "code1").is_some());
assert!(cache.get("script1", "code1").is_none());
}
#[test]
fn test_cache_eviction() {
let cache = ScriptCache::new(2);
cache.insert("s1", "c1", vec![1]);
cache.insert("s2", "c2", vec![2]);
assert_eq!(cache.stats().entries, 2);
cache.insert("s3", "c3", vec![3]);
assert_eq!(cache.stats().entries, 2);
assert_eq!(cache.stats().evictions, 1);
}
#[test]
fn test_hit_rate() {
let cache = ScriptCache::new(10);
cache.insert("s1", "c1", vec![1]);
cache.get("s1", "c1");
cache.get("s1", "c1");
cache.get("s1", "c1");
cache.get("s2", "c2");
let stats = cache.stats();
assert_eq!(stats.hits, 3);
assert_eq!(stats.misses, 1);
assert_eq!(stats.hit_rate(), 75.0);
}
#[test]
fn test_compile_error_not_cached() {
let cache = ScriptCache::new(10);
let lua = Lua::new();
let result =
cache.get_or_compile("bad_script", "this is not valid lua syntax !!!", |code| {
let chunk = lua.load(code);
let func = chunk.into_function()?;
Ok(func.dump(false))
});
assert!(result.is_err());
let stats = cache.stats();
assert_eq!(stats.entries, 0, "Failed compile should not be cached");
}
#[test]
fn test_invalidate_nonexistent() {
let cache = ScriptCache::new(10);
cache.insert("s1", "c1", vec![1]);
assert_eq!(cache.stats().entries, 1);
cache.invalidate("nonexistent");
assert_eq!(cache.stats().entries, 1);
assert!(cache.get("s1", "c1").is_some());
}
#[test]
fn test_concurrent_same_script() {
use std::sync::Arc;
use std::thread;
let cache = Arc::new(ScriptCache::new(10));
let compile_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let handles: Vec<_> = (0..4)
.map(|_| {
let c = cache.clone();
let cc = compile_count.clone();
thread::spawn(move || {
c.get_or_compile("test", "return 1", |_| {
cc.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
std::thread::sleep(std::time::Duration::from_millis(10));
Ok(vec![1, 2, 3])
})
.unwrap()
})
})
.collect();
for h in handles {
h.join().unwrap();
}
assert_eq!(cache.stats().entries, 1);
}
#[test]
fn test_different_code_same_key() {
let cache = ScriptCache::new(10);
cache.insert("s1", "code_v1", vec![1, 1, 1]);
cache.insert("s1", "code_v2", vec![2, 2, 2]);
assert_eq!(cache.stats().entries, 2);
let v1 = cache.get("s1", "code_v1");
let v2 = cache.get("s1", "code_v2");
assert!(v1.is_some());
assert!(v2.is_some());
assert_eq!(v1.unwrap(), vec![1, 1, 1]);
assert_eq!(v2.unwrap(), vec![2, 2, 2]);
}
#[test]
fn test_invalidate_clears_all_versions() {
let cache = ScriptCache::new(10);
cache.insert("s1", "code_v1", vec![1]);
cache.insert("s1", "code_v2", vec![2]);
cache.insert("s2", "code_x", vec![3]);
assert_eq!(cache.stats().entries, 3);
cache.invalidate("s1");
assert_eq!(cache.stats().entries, 1);
assert!(cache.get("s1", "code_v1").is_none());
assert!(cache.get("s1", "code_v2").is_none());
assert!(cache.get("s2", "code_x").is_some());
}
#[test]
fn test_lru_eviction_order() {
let cache = ScriptCache::new(3);
cache.insert("s1", "c1", vec![1]);
std::thread::sleep(std::time::Duration::from_millis(5));
cache.insert("s2", "c2", vec![2]);
std::thread::sleep(std::time::Duration::from_millis(5));
cache.insert("s3", "c3", vec![3]);
cache.get("s1", "c1");
cache.insert("s4", "c4", vec![4]);
assert_eq!(cache.stats().entries, 3);
assert!(
cache.get("s1", "c1").is_some(),
"s1 should still exist (recently accessed)"
);
assert!(cache.get("s3", "c3").is_some(), "s3 should still exist");
assert!(
cache.get("s4", "c4").is_some(),
"s4 should exist (just added)"
);
}
#[test]
fn test_clear() {
let cache = ScriptCache::new(10);
cache.insert("s1", "c1", vec![1]);
cache.insert("s2", "c2", vec![2]);
cache.insert("s3", "c3", vec![3]);
assert_eq!(cache.stats().entries, 3);
cache.clear();
assert_eq!(cache.stats().entries, 0);
assert!(cache.get("s1", "c1").is_none());
}
#[test]
fn test_zero_hit_rate() {
let cache = ScriptCache::new(10);
let stats = cache.stats();
assert_eq!(stats.hit_rate(), 0.0);
cache.get("nonexistent", "code");
cache.get("nonexistent2", "code2");
let stats = cache.stats();
assert_eq!(stats.hits, 0);
assert_eq!(stats.misses, 2);
assert_eq!(stats.hit_rate(), 0.0);
}
}