use crate::kv_error::{CacheValidationError, ValidationResult, KvError, KvResult};
use crate::runtime::transformer::KvCache;
#[derive(Clone)]
pub struct CachedPrefix {
pub kv: KvCache,
pub last_logits: Vec<f32>,
}
#[derive(Clone)]
pub struct CacheLookup {
pub matched_len: usize,
pub kv: KvCache,
pub last_logits: Option<Vec<f32>>,
}
pub struct PrefixCache {
entries: Vec<(Vec<u32>, CachedPrefix)>,
max_entries: usize,
}
impl PrefixCache {
pub fn new(max_entries: usize) -> Self {
Self {
entries: Vec::new(),
max_entries: max_entries.max(1),
}
}
pub fn lookup(&self, tokens: &[u32]) -> KvResult<CacheLookup> {
Self::validate_tokens(tokens)?;
let mut best: Option<(&Vec<u32>, &CachedPrefix)> = None;
for (seq, cached) in &self.entries {
if tokens.starts_with(seq)
&& best.is_none_or(|(b, _)| seq.len() > b.len())
{
best = Some((seq, cached));
}
}
let (seq, cached) = best.ok_or_else(|| {
if self.is_empty() {
KvError::EmptyCache("lookup".to_string())
} else {
KvError::InvalidTokenSequence {
reason: "No matching prefix found".to_string(),
position: None,
}
}
})?;
let matched_len = seq.len();
let last_logits = (matched_len == tokens.len() && !cached.last_logits.is_empty())
.then(|| cached.last_logits.clone());
if matched_len > 0 && cached.kv.layers.is_empty() {
return Err(KvError::InvalidCacheState(
"Matched prefix has empty KV layers".to_string()
));
}
Ok(CacheLookup {
matched_len,
kv: cached.kv.clone(),
last_logits,
})
}
pub fn insert(&mut self, tokens: Vec<u32>, cached: CachedPrefix) -> ValidationResult<()> {
Self::validate_tokens(&tokens)?;
if !cached.kv.layers.is_empty() {
let kv_length = cached.kv.layers[0].as_ref().map(|l| l.len()).unwrap_or(0);
if kv_length > 0 && !tokens.is_empty() && kv_length < tokens.len() / 2 {
return Err(CacheValidationError::DimensionMismatch {
layer: 0,
expected: tokens.len(),
actual: kv_length,
});
}
}
if tokens.len() > self.max_entries && !self.entries.iter().any(|(k, _)| *k == tokens) {
return Err(CacheValidationError::SequenceTooLong {
length: tokens.len(),
max_length: self.max_entries,
});
}
if let Some(pos) = self.entries.iter().position(|(k, _)| *k == tokens) {
self.entries.remove(pos);
}
self.entries.insert(0, (tokens, cached));
while self.entries.len() > self.max_entries {
let last = self.entries.len() - 1;
self.entries.remove(last);
}
Ok(())
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn capacity(&self) -> usize {
self.max_entries
}
pub fn clear(&mut self) {
self.entries.clear();
}
pub fn remove(&mut self, tokens: &[u32]) -> bool {
if let Some(pos) = self.entries.iter().position(|(k, _)| *k == tokens) {
self.entries.remove(pos);
true
} else {
false
}
}
fn validate_tokens(tokens: &[u32]) -> ValidationResult<()> {
if tokens.is_empty() {
return Err(CacheValidationError::EmptySequence);
}
const MAX_REASONABLE_TOKEN_ID: u32 = 1_000_000;
for (pos, &token) in tokens.iter().enumerate() {
if token > MAX_REASONABLE_TOKEN_ID {
return Err(CacheValidationError::InvalidTokenId {
token_id: token,
position: pos,
});
}
}
Ok(())
}
pub fn stats(&self) -> PrefixCacheStats {
let total_tokens: usize = self.entries.iter()
.map(|(tokens, _)| tokens.len())
.sum();
let total_kv_bytes: usize = self.entries.iter()
.map(|(_, cached)| {
cached.kv.layers.iter()
.filter_map(|layer| layer.as_ref())
.map(|layer| layer.len() * 2 * 4) .sum::<usize>()
})
.sum();
PrefixCacheStats {
entries: self.entries.len(),
capacity: self.max_entries,
total_tokens,
total_kv_bytes,
utilization: if self.max_entries > 0 {
self.entries.len() as f64 / self.max_entries as f64
} else {
0.0
},
}
}
}
#[derive(Debug, Clone)]
pub struct PrefixCacheStats {
pub entries: usize,
pub capacity: usize,
pub total_tokens: usize,
pub total_kv_bytes: usize,
pub utilization: f64,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::runtime::transformer::{KvCache, KvLayer};
fn dummy_cache(n_positions: usize) -> KvCache {
let mut kv = KvCache::new(1);
let mut layer = KvLayer::new_fp32();
for _ in 0..n_positions {
layer.append(&[1.0, 2.0], &[3.0, 4.0]);
}
kv.layers[0] = Some(layer);
kv
}
#[test]
fn lookup_miss_returns_error() {
let pc = PrefixCache::new(4);
let result = pc.lookup(&[1, 2, 3]);
assert!(result.is_err());
assert!(matches!(result, Err(KvError::EmptyCache(_))));
assert!(pc.is_empty());
}
#[test]
fn exact_match_returns_logits_and_skips_reprocess() {
let mut pc = PrefixCache::new(4);
let tokens = vec![1, 2, 3];
pc.insert(
tokens.clone(),
CachedPrefix {
kv: dummy_cache(3),
last_logits: vec![0.5, 0.2],
},
).expect("insert should succeed");
let look = pc.lookup(&tokens).expect("exact match");
assert_eq!(look.matched_len, 3);
assert_eq!(look.last_logits, Some(vec![0.5, 0.2]));
}
#[test]
fn prefix_match_returns_cloned_kv_without_logits() {
let mut pc = PrefixCache::new(4);
pc.insert(
vec![1, 2],
CachedPrefix {
kv: dummy_cache(2),
last_logits: vec![0.9],
},
).expect("insert should succeed");
let look = pc.lookup(&[1, 2, 3, 4]).expect("prefix match");
assert_eq!(look.matched_len, 2);
assert!(look.last_logits.is_none(), "partial match must not expose logits");
let layer = look.kv.layers[0].as_ref().unwrap();
assert_eq!(layer.len(), 2);
}
#[test]
fn longest_prefix_wins() {
let mut pc = PrefixCache::new(4);
pc.insert(
vec![1],
CachedPrefix { kv: dummy_cache(1), last_logits: vec![] },
).expect("insert should succeed");
pc.insert(
vec![1, 2, 3],
CachedPrefix { kv: dummy_cache(3), last_logits: vec![] },
).expect("insert should succeed");
let look = pc.lookup(&[1, 2, 3, 4]).expect("match");
assert_eq!(look.matched_len, 3, "longer cached prefix should win");
}
#[test]
fn non_prefix_entry_is_ignored() {
let mut pc = PrefixCache::new(4);
pc.insert(
vec![9, 9],
CachedPrefix { kv: dummy_cache(2), last_logits: vec![] },
).expect("insert should succeed");
let result = pc.lookup(&[1, 2, 9, 9]);
assert!(result.is_err(), "entry must be a *prefix* of the query");
}
#[test]
fn insert_refreshes_existing_key() {
let mut pc = PrefixCache::new(4);
pc.insert(
vec![1, 2],
CachedPrefix { kv: dummy_cache(2), last_logits: vec![0.1] },
).expect("insert should succeed");
pc.insert(
vec![1, 2],
CachedPrefix { kv: dummy_cache(2), last_logits: vec![0.9] },
).expect("insert should succeed");
assert_eq!(pc.len(), 1, "refresh should not duplicate");
let look = pc.lookup(&[1, 2]).unwrap();
assert_eq!(look.last_logits, Some(vec![0.9]), "refresh should overwrite");
}
#[test]
fn lru_eviction_drops_oldest() {
let mut pc = PrefixCache::new(2);
pc.insert(vec![1], CachedPrefix { kv: dummy_cache(1), last_logits: vec![] }).expect("insert should succeed");
pc.insert(vec![2], CachedPrefix { kv: dummy_cache(1), last_logits: vec![] }).expect("insert should succeed");
pc.insert(vec![3], CachedPrefix { kv: dummy_cache(1), last_logits: vec![] }).expect("insert should succeed");
assert_eq!(pc.len(), 2);
assert!(pc.lookup(&[1, 2]).is_err());
assert!(pc.lookup(&[2]).is_ok());
assert!(pc.lookup(&[3]).is_ok());
}
#[test]
fn capacity_minimum_is_one() {
let pc = PrefixCache::new(0);
assert_eq!(pc.capacity(), 1);
}
#[test]
fn empty_tokens_rejected() {
let mut pc = PrefixCache::new(4);
let result = pc.insert(vec![], CachedPrefix {
kv: dummy_cache(0),
last_logits: vec![]
});
assert!(matches!(result, Err(CacheValidationError::EmptySequence)));
}
#[test]
fn invalid_token_id_rejected() {
let mut pc = PrefixCache::new(4);
let invalid_tokens = vec![1, 2, 999999999]; let result = pc.insert(invalid_tokens, CachedPrefix {
kv: dummy_cache(3),
last_logits: vec![]
});
assert!(matches!(result, Err(CacheValidationError::InvalidTokenId { .. })));
}
#[test]
fn dimension_mismatch_rejected() {
let mut pc = PrefixCache::new(4);
let result = pc.insert(vec![1, 2, 3], CachedPrefix {
kv: dummy_cache(2), last_logits: vec![]
});
assert!(result.is_ok(), "Lenient validation should allow this case");
let extreme_result = pc.insert(vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10], CachedPrefix {
kv: dummy_cache(1), last_logits: vec![]
});
assert!(extreme_result.is_err(), "Extreme mismatch should be rejected");
}
#[test]
fn cache_stats_work() {
let mut pc = PrefixCache::new(4);
pc.insert(vec![1, 2], CachedPrefix {
kv: dummy_cache(2),
last_logits: vec![]
}).expect("insert should succeed");
pc.insert(vec![3, 4, 5], CachedPrefix {
kv: dummy_cache(3),
last_logits: vec![]
}).expect("insert should succeed");
let stats = pc.stats();
assert_eq!(stats.entries, 2);
assert_eq!(stats.capacity, 4);
assert_eq!(stats.total_tokens, 5);
assert!(stats.total_kv_bytes > 0);
assert!((stats.utilization - 0.5).abs() < 0.01); }
#[test]
fn remove_existing_entry() {
let mut pc = PrefixCache::new(4);
pc.insert(vec![1, 2], CachedPrefix {
kv: dummy_cache(2),
last_logits: vec![]
}).expect("insert should succeed");
assert_eq!(pc.len(), 1);
assert!(pc.remove(&[1, 2]));
assert_eq!(pc.len(), 0);
}
#[test]
fn remove_nonexistent_entry() {
let mut pc = PrefixCache::new(4);
assert!(!pc.remove(&[1, 2]));
}
#[test]
fn clear_empties_cache() {
let mut pc = PrefixCache::new(4);
pc.insert(vec![1, 2], CachedPrefix {
kv: dummy_cache(2),
last_logits: vec![]
}).expect("insert should succeed");
assert_eq!(pc.len(), 1);
pc.clear();
assert_eq!(pc.len(), 0);
assert!(pc.is_empty());
}
}