use crate::cache::CacheMetrics;
use crate::error::{Error, Result};
use dashmap::DashMap;
use lru::LruCache;
use sha2::{Digest, Sha256};
use std::cell::RefCell;
use std::collections::HashMap;
use std::num::NonZeroUsize;
use std::path::PathBuf;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, RwLock};
use std::time::Instant;
use tracing::info;
const MIN_PREFIX_LENGTH: usize = 100;
const MAX_TRACKED_PREFIXES: usize = 1000;
thread_local! {
static LOCAL_PREFIX_CACHE: RefCell<LruCache<String, Arc<SessionData>>> =
RefCell::new(LruCache::new(NonZeroUsize::new(100).unwrap()));
}
#[derive(Debug)]
pub struct SessionData {
pub prefix: String,
pub token_count: usize,
pub file_path: Option<PathBuf>,
pub memory_state: Option<Vec<u8>>,
pub created_at: Instant,
pub last_accessed: AtomicU64,
pub access_count: AtomicUsize,
}
impl SessionData {
pub fn new(prefix: String, token_count: usize) -> Self {
Self {
prefix,
token_count,
file_path: None,
memory_state: None,
created_at: Instant::now(),
last_accessed: AtomicU64::new(0),
access_count: AtomicUsize::new(0),
}
}
pub fn touch(&self) {
self.access_count.fetch_add(1, Ordering::Relaxed);
let now = Instant::now().elapsed().as_secs();
self.last_accessed.store(now, Ordering::Relaxed);
}
pub fn age_seconds(&self) -> u64 {
self.created_at.elapsed().as_secs()
}
}
#[derive(Default)]
struct TrieNode {
children: HashMap<i32, Box<TrieNode>>,
session: Option<Arc<SessionData>>,
#[allow(dead_code)]
frequency: usize,
}
pub struct PrefixDetector {
root: RwLock<TrieNode>,
frequency_map: DashMap<String, usize>,
frequency_threshold: usize,
max_tracked: usize,
}
impl PrefixDetector {
pub fn new(frequency_threshold: usize) -> Self {
Self {
root: RwLock::new(TrieNode::default()),
frequency_map: DashMap::new(),
frequency_threshold,
max_tracked: MAX_TRACKED_PREFIXES,
}
}
pub fn analyze_tokens(&self, tokens: &[i32]) -> Option<usize> {
if tokens.len() < MIN_PREFIX_LENGTH {
return None;
}
let prefix_lengths = [
MIN_PREFIX_LENGTH,
MIN_PREFIX_LENGTH * 2,
tokens.len() / 2,
tokens.len() * 3 / 4,
];
let mut best_prefix_len = None;
let mut best_frequency = 0;
for &len in &prefix_lengths {
if len > tokens.len() {
continue;
}
let prefix_key = Self::tokens_to_key(&tokens[..len]);
let mut count = self.frequency_map.entry(prefix_key.clone()).or_insert(0);
*count += 1;
if *count >= self.frequency_threshold && *count > best_frequency {
best_frequency = *count;
best_prefix_len = Some(len);
}
}
if self.frequency_map.len() > self.max_tracked {
self.cleanup_infrequent();
}
best_prefix_len
}
pub fn find_longest_prefix(&self, tokens: &[i32]) -> Option<(usize, Arc<SessionData>)> {
let root = self.root.read().unwrap();
let mut node = &*root;
let mut best_match = None;
for (idx, &token) in tokens.iter().enumerate() {
match node.children.get(&token) {
Some(child) => {
node = child;
if let Some(session) = &node.session {
best_match = Some((idx + 1, session.clone()));
}
}
None => break,
}
}
best_match
}
pub fn insert_prefix(&self, tokens: &[i32], session: Arc<SessionData>) {
if tokens.len() < MIN_PREFIX_LENGTH {
return;
}
let mut root = self.root.write().unwrap();
let mut node = &mut *root;
for &token in tokens {
node = node
.children
.entry(token)
.or_insert_with(|| Box::new(TrieNode::default()));
}
node.session = Some(session);
}
fn tokens_to_key(tokens: &[i32]) -> String {
let mut hasher = Sha256::new();
for &token in tokens {
hasher.update(token.to_le_bytes());
}
format!("{:x}", hasher.finalize())
}
fn cleanup_infrequent(&self) {
let threshold = self.frequency_threshold / 2;
self.frequency_map.retain(|_, count| *count >= threshold);
}
}
pub struct PrefixCache {
sessions: Arc<DashMap<String, Arc<SessionData>>>,
detector: Arc<PrefixDetector>,
max_sessions: usize,
ttl_seconds: u64,
metrics: Arc<CacheMetrics>,
session_dir: Option<PathBuf>,
}
impl PrefixCache {
pub fn new(
max_sessions: usize,
ttl_seconds: u64,
frequency_threshold: usize,
session_dir: Option<PathBuf>,
) -> Result<Self> {
if let Some(ref dir) = session_dir {
std::fs::create_dir_all(dir).map_err(|e| Error::ConfigurationError {
message: format!("Failed to create session directory: {e}"),
})?;
}
Ok(Self {
sessions: Arc::new(DashMap::new()),
detector: Arc::new(PrefixDetector::new(frequency_threshold)),
max_sessions,
ttl_seconds,
metrics: Arc::new(CacheMetrics::default()),
session_dir,
})
}
pub fn find_prefix_session(
&self,
text: &str,
tokens: &[i32],
) -> Option<(usize, Arc<SessionData>)> {
let key = Self::compute_key(text);
let local_hit = LOCAL_PREFIX_CACHE.with(|cache| cache.borrow_mut().get(&key).cloned());
if let Some(session) = local_hit {
self.metrics.hits.fetch_add(1, Ordering::Relaxed);
session.touch();
return Some((session.token_count, session));
}
if let Some((prefix_len, session)) = self.detector.find_longest_prefix(tokens) {
if session.age_seconds() < self.ttl_seconds {
self.metrics.hits.fetch_add(1, Ordering::Relaxed);
session.touch();
LOCAL_PREFIX_CACHE.with(|cache| {
cache.borrow_mut().put(key, session.clone());
});
return Some((prefix_len, session));
}
}
self.metrics.misses.fetch_add(1, Ordering::Relaxed);
None
}
pub fn register_prefix(&self, text: &str, tokens: &[i32], session_data: Vec<u8>) -> Result<()> {
if tokens.len() < MIN_PREFIX_LENGTH {
return Err(Error::InvalidInput {
message: format!(
"Prefix too short: {} tokens (minimum: {})",
tokens.len(),
MIN_PREFIX_LENGTH
),
});
}
if self.sessions.len() >= self.max_sessions {
self.evict_oldest();
}
let key = Self::compute_key(text);
let mut session = SessionData::new(text.to_string(), tokens.len());
session.memory_state = Some(session_data);
if let Some(ref dir) = self.session_dir {
let file_path = dir.join(format!("{key}.session"));
session.file_path = Some(file_path);
}
let session = Arc::new(session);
self.sessions.insert(key.clone(), session.clone());
self.detector.insert_prefix(tokens, session.clone());
info!(
"Registered prefix cache for {} tokens (key: {})",
tokens.len(),
&key[..8]
);
Ok(())
}
pub fn analyze(&self, tokens: &[i32]) -> Option<usize> {
self.detector.analyze_tokens(tokens)
}
pub fn clear(&self) {
self.sessions.clear();
LOCAL_PREFIX_CACHE.with(|cache| cache.borrow_mut().clear());
*self.detector.root.write().unwrap() = TrieNode::default();
self.metrics
.evictions
.store(self.sessions.len() as u64, Ordering::Relaxed);
}
pub fn stats(&self) -> PrefixCacheStats {
PrefixCacheStats {
session_count: self.sessions.len(),
total_hits: self.metrics.hits.load(Ordering::Relaxed),
total_misses: self.metrics.misses.load(Ordering::Relaxed),
total_evictions: self.metrics.evictions.load(Ordering::Relaxed),
memory_usage_bytes: self.estimate_memory_usage(),
}
}
fn compute_key(text: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(text.as_bytes());
format!("{:x}", hasher.finalize())
}
fn evict_oldest(&self) {
let mut oldest_key = None;
let mut oldest_time = u64::MAX;
for entry in self.sessions.iter() {
let last_accessed = entry.value().last_accessed.load(Ordering::Relaxed);
if last_accessed < oldest_time {
oldest_time = last_accessed;
oldest_key = Some(entry.key().clone());
}
}
if let Some(key) = oldest_key {
self.sessions.remove(&key);
self.metrics.evictions.fetch_add(1, Ordering::Relaxed);
}
}
fn estimate_memory_usage(&self) -> u64 {
let mut total = 0u64;
for entry in self.sessions.iter() {
let session = entry.value();
total += session.prefix.len() as u64;
if let Some(ref state) = session.memory_state {
total += state.len() as u64;
}
total += 256; }
total
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct PrefixCacheStats {
pub session_count: usize,
pub total_hits: u64,
pub total_misses: u64,
pub total_evictions: u64,
pub memory_usage_bytes: u64,
}
impl PrefixCacheStats {
pub fn hit_rate(&self) -> f64 {
let total = self.total_hits + self.total_misses;
if total == 0 {
0.0
} else {
#[allow(clippy::cast_precision_loss)]
{
(self.total_hits as f64 / total as f64) * 100.0
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_prefix_detector() {
let detector = PrefixDetector::new(3);
let tokens1: Vec<i32> = (0..150).collect();
let tokens2: Vec<i32> = (0..150).collect();
let tokens3: Vec<i32> = (0..150).collect();
assert_eq!(detector.analyze_tokens(&tokens1), None);
assert_eq!(detector.analyze_tokens(&tokens2), None);
let result = detector.analyze_tokens(&tokens3);
assert!(result.is_some(), "Expected Some(_), got None");
assert_eq!(result.unwrap(), MIN_PREFIX_LENGTH);
}
#[test]
fn test_prefix_trie() {
let detector = PrefixDetector::new(1);
let tokens: Vec<i32> = (0..200).collect();
let session = Arc::new(SessionData::new("test".to_string(), 150));
detector.insert_prefix(&tokens[..150], session.clone());
let result = detector.find_longest_prefix(&tokens);
assert!(result.is_some());
let (len, found_session) = result.unwrap();
assert_eq!(len, 150);
assert_eq!(found_session.token_count, 150);
}
#[test]
fn test_cache_operations() {
let cache = PrefixCache::new(10, 3600, 2, None).unwrap();
let text = "This is a test prefix that is long enough to be cached";
let tokens: Vec<i32> = (0..150).collect();
assert!(cache.find_prefix_session(text, &tokens).is_none());
let session_data = vec![1, 2, 3, 4, 5];
cache.register_prefix(text, &tokens, session_data).unwrap();
let result = cache.find_prefix_session(text, &tokens);
assert!(result.is_some());
let stats = cache.stats();
assert_eq!(stats.session_count, 1);
assert_eq!(stats.total_hits, 1);
}
}