use crate::cache::{CacheMetrics, CacheStats, CacheStore};
use dashmap::DashMap;
use llama_cpp_2::token::LlamaToken;
use lru::LruCache;
use sha2::{Digest, Sha256};
use std::cell::RefCell;
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use tracing::trace;
thread_local! {
static LOCAL_TOKEN_CACHE: RefCell<LruCache<String, Vec<LlamaToken>>> =
RefCell::new(LruCache::new(NonZeroUsize::new(1000).unwrap()));
}
pub struct TokenCache {
shared: Arc<DashMap<String, Vec<LlamaToken>>>,
metrics: Arc<CacheMetrics>,
max_size: usize,
#[allow(dead_code)]
ttl_seconds: Option<u64>,
}
impl TokenCache {
pub fn new(max_size: usize) -> Self {
Self::with_ttl(max_size, None)
}
pub fn with_ttl(max_size: usize, ttl_seconds: Option<u64>) -> Self {
trace!(
"Creating TokenCache with max_size: {}, ttl: {:?}",
max_size, ttl_seconds
);
Self {
shared: Arc::new(DashMap::with_capacity(max_size.min(100_000))),
metrics: Arc::new(CacheMetrics::new()),
max_size,
ttl_seconds,
}
}
pub fn compute_key(text: &str, model_name: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(text.as_bytes());
hasher.update(model_name.as_bytes());
format!("{:x}", hasher.finalize())
}
pub fn metrics(&self) -> &CacheMetrics {
&self.metrics
}
}
impl CacheStore<String, Vec<LlamaToken>> for TokenCache {
fn get(&self, key: &String) -> Option<Vec<LlamaToken>> {
let result = LOCAL_TOKEN_CACHE.with(|cache| {
let mut cache = cache.borrow_mut();
if let Some(tokens) = cache.get(key) {
trace!("Token cache hit (thread-local) for key: {}", key);
self.metrics.record_hit();
return Some(tokens.clone());
}
None
});
if result.is_some() {
return result;
}
if let Some(entry) = self.shared.get(key) {
let tokens = entry.clone();
trace!("Token cache hit (shared) for key: {}", key);
LOCAL_TOKEN_CACHE.with(|cache| {
let mut cache = cache.borrow_mut();
cache.put(key.clone(), tokens.clone());
});
self.metrics.record_hit();
Some(tokens)
} else {
trace!("Token cache miss for key: {}", key);
self.metrics.record_miss();
None
}
}
fn insert(&self, key: String, value: Vec<LlamaToken>) {
let token_count = value.len();
trace!(
"Inserting {} tokens into cache with key: {}",
token_count, key
);
LOCAL_TOKEN_CACHE.with(|cache| {
let mut cache = cache.borrow_mut();
cache.put(key.clone(), value.clone());
});
if self.shared.len() < self.max_size {
self.shared.insert(key, value);
let memory_bytes = (token_count * 4 + 64) as u64;
self.metrics
.memory_bytes
.fetch_add(memory_bytes, Ordering::Relaxed);
} else {
trace!("Shared cache full, skipping insertion");
}
}
fn clear(&self) {
self.shared.clear();
LOCAL_TOKEN_CACHE.with(|cache| cache.borrow_mut().clear());
self.metrics.reset();
}
fn stats(&self) -> CacheStats {
CacheStats::from_metrics(
&self.metrics,
self.shared.len().try_into().unwrap_or(u64::MAX),
)
}
fn len(&self) -> usize {
self.shared.len()
}
}
impl TokenCache {
pub fn evict_oldest(&self, count: usize) {
LOCAL_TOKEN_CACHE.with(|cache| {
cache.borrow_mut().clear();
});
if count == 0 {
return;
}
let current_count = self.shared.len();
if current_count == 0 {
return;
}
if count >= current_count / 2 {
self.clear();
} else {
let keys_to_remove: Vec<_> = self
.shared
.iter()
.take(count)
.map(|entry| entry.key().clone())
.collect();
let removed = keys_to_remove.len();
for key in keys_to_remove {
self.shared.remove(&key);
}
self.metrics.record_eviction(removed as u64);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use llama_cpp_2::token::LlamaToken;
fn create_test_tokens() -> Vec<LlamaToken> {
vec![LlamaToken(1), LlamaToken(2), LlamaToken(3)]
}
#[test]
fn test_token_cache_creation() {
let cache = TokenCache::new(1000);
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn test_compute_key() {
let key1 = TokenCache::compute_key("hello", "model1");
let key2 = TokenCache::compute_key("hello", "model1");
let key3 = TokenCache::compute_key("hello", "model2");
let key4 = TokenCache::compute_key("world", "model1");
assert_eq!(key1, key2);
assert_ne!(key1, key3); assert_ne!(key1, key4); }
#[test]
fn test_cache_miss() {
let cache = TokenCache::new(100);
let key = "test_key".to_string();
let result = cache.get(&key);
assert!(result.is_none());
assert_eq!(cache.metrics.misses.load(Ordering::Relaxed), 1);
assert_eq!(cache.metrics.hits.load(Ordering::Relaxed), 0);
}
#[test]
fn test_cache_hit() {
let cache = TokenCache::new(100);
let key = "test_key".to_string();
let tokens = create_test_tokens();
cache.insert(key.clone(), tokens.clone());
let result = cache.get(&key);
assert_eq!(result, Some(tokens));
assert_eq!(cache.metrics.hits.load(Ordering::Relaxed), 1);
}
#[test]
fn test_thread_local_caching() {
let cache = TokenCache::new(100);
let key = "test_key".to_string();
let tokens = create_test_tokens();
cache.insert(key.clone(), tokens.clone());
let result1 = cache.get(&key);
assert_eq!(result1, Some(tokens.clone()));
let result2 = cache.get(&key);
assert_eq!(result2, Some(tokens));
assert_eq!(cache.metrics.hits.load(Ordering::Relaxed), 2);
}
#[test]
fn test_cache_clear() {
let cache = TokenCache::new(100);
let key = "test_key".to_string();
let tokens = create_test_tokens();
cache.insert(key.clone(), tokens);
assert_eq!(cache.len(), 1);
cache.clear();
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
assert_eq!(cache.metrics.hits.load(Ordering::Relaxed), 0);
assert_eq!(cache.metrics.misses.load(Ordering::Relaxed), 0);
let result = cache.get(&key);
assert!(result.is_none());
assert_eq!(cache.metrics.misses.load(Ordering::Relaxed), 1);
}
#[test]
fn test_cache_stats() {
let cache = TokenCache::new(100);
let key1 = "key1".to_string();
let key2 = "key2".to_string();
let tokens = create_test_tokens();
cache.insert(key1.clone(), tokens.clone());
cache.insert(key2, tokens);
let _ = cache.get(&key1); let _ = cache.get(&"missing".to_string());
let stats = cache.stats();
assert_eq!(stats.hits, 1);
assert_eq!(stats.misses, 1);
assert_eq!(stats.entry_count, 2);
assert_eq!(stats.hit_rate, 0.5);
}
#[test]
fn test_ttl_cache_creation() {
let cache = TokenCache::with_ttl(100, Some(3600));
assert_eq!(cache.ttl_seconds, Some(3600));
assert!(cache.is_empty());
}
}