1use ferrum_types::{FerrumError, Result, TokenId};
4use parking_lot::RwLock;
5use std::collections::HashMap;
6use std::sync::Arc;
7use tracing::{debug, trace};
8
9#[derive(Debug, Clone, PartialEq, Eq, Hash)]
11pub struct PrefixId(Vec<TokenId>);
12
13impl PrefixId {
14 pub fn new(tokens: Vec<TokenId>) -> Self {
16 Self(tokens)
17 }
18
19 pub fn tokens(&self) -> &[TokenId] {
21 &self.0
22 }
23
24 pub fn len(&self) -> usize {
26 self.0.len()
27 }
28
29 pub fn is_empty(&self) -> bool {
31 self.0.is_empty()
32 }
33}
34
35impl From<Vec<TokenId>> for PrefixId {
36 fn from(tokens: Vec<TokenId>) -> Self {
37 Self::new(tokens)
38 }
39}
40
41impl From<&[TokenId]> for PrefixId {
42 fn from(tokens: &[TokenId]) -> Self {
43 Self::new(tokens.to_vec())
44 }
45}
46
47#[derive(Debug, Clone)]
49pub struct CachedPrefix {
50 pub prefix_id: PrefixId,
52 pub kv_handle: Arc<dyn ferrum_interfaces::KvCacheHandle + Send + Sync>,
54 pub last_logits: Vec<f32>,
57 pub ref_count: usize,
59 pub last_access: std::time::Instant,
61 pub size: usize,
63}
64
65impl CachedPrefix {
66 pub fn new(
68 prefix_id: PrefixId,
69 kv_handle: Arc<dyn ferrum_interfaces::KvCacheHandle + Send + Sync>,
70 last_logits: Vec<f32>,
71 ) -> Self {
72 let size = prefix_id.len();
73 Self {
74 prefix_id,
75 kv_handle,
76 last_logits,
77 ref_count: 1,
78 last_access: std::time::Instant::now(),
79 size,
80 }
81 }
82
83 pub fn add_ref(&mut self) {
85 self.ref_count += 1;
86 self.touch();
87 }
88
89 pub fn remove_ref(&mut self) -> Result<()> {
91 if self.ref_count == 0 {
92 return Err(FerrumError::invalid_parameter(
93 "Cannot remove ref from zero-ref prefix",
94 ));
95 }
96 self.ref_count -= 1;
97 Ok(())
98 }
99
100 pub fn touch(&mut self) {
102 self.last_access = std::time::Instant::now();
103 }
104
105 pub fn can_evict(&self) -> bool {
107 self.ref_count == 0
108 }
109}
110
111#[derive(Debug)]
113pub struct PrefixCache {
114 prefixes: RwLock<HashMap<PrefixId, CachedPrefix>>,
116 max_prefixes: usize,
118 min_prefix_length: usize,
120 hits: parking_lot::Mutex<usize>,
122 misses: parking_lot::Mutex<usize>,
123 evictions: parking_lot::Mutex<usize>,
124}
125
126impl PrefixCache {
127 pub fn new(max_prefixes: usize, min_prefix_length: usize) -> Self {
129 Self {
130 prefixes: RwLock::new(HashMap::new()),
131 max_prefixes,
132 min_prefix_length,
133 hits: parking_lot::Mutex::new(0),
134 misses: parking_lot::Mutex::new(0),
135 evictions: parking_lot::Mutex::new(0),
136 }
137 }
138
139 pub fn find_prefix(
144 &self,
145 tokens: &[TokenId],
146 ) -> Option<(
147 PrefixId,
148 Arc<dyn ferrum_interfaces::KvCacheHandle + Send + Sync>,
149 Vec<f32>,
150 )> {
151 if tokens.len() < self.min_prefix_length {
152 return None;
153 }
154
155 let prefixes = self.prefixes.read();
156
157 let mut best_match = None;
159 let mut best_len = 0;
160
161 for (prefix_id, cached_prefix) in prefixes.iter() {
162 if tokens.starts_with(prefix_id.tokens()) && prefix_id.len() > best_len {
163 best_match = Some((
164 prefix_id.clone(),
165 cached_prefix.kv_handle.clone(),
166 cached_prefix.last_logits.clone(),
167 ));
168 best_len = prefix_id.len();
169 }
170 }
171
172 if let Some(ref match_info) = best_match {
173 *self.hits.lock() += 1;
174 trace!("Prefix cache hit: {} tokens", best_len);
175
176 drop(prefixes); let mut prefixes = self.prefixes.write();
179 if let Some(cached_prefix) = prefixes.get_mut(&match_info.0) {
180 cached_prefix.touch();
181 }
182 } else {
183 *self.misses.lock() += 1;
184 trace!("Prefix cache miss for {} tokens", tokens.len());
185 }
186
187 best_match
188 }
189
190 pub fn store_prefix(
192 &self,
193 prefix_tokens: &[TokenId],
194 kv_handle: Arc<dyn ferrum_interfaces::KvCacheHandle + Send + Sync>,
195 last_logits: Vec<f32>,
196 ) -> Result<()> {
197 if prefix_tokens.len() < self.min_prefix_length {
198 return Ok(()); }
200
201 let prefix_id = PrefixId::from(prefix_tokens);
202 let cached_prefix = CachedPrefix::new(prefix_id.clone(), kv_handle, last_logits);
203
204 let mut prefixes = self.prefixes.write();
205
206 if prefixes.len() >= self.max_prefixes && !prefixes.contains_key(&prefix_id) {
208 self.evict_lru(&mut prefixes);
209 }
210
211 if let Some(existing) = prefixes.get_mut(&prefix_id) {
213 existing.add_ref();
214 } else {
215 prefixes.insert(prefix_id, cached_prefix);
216 debug!("Stored new prefix: {} tokens", prefix_tokens.len());
217 }
218
219 Ok(())
220 }
221
222 pub fn remove_ref(&self, prefix_tokens: &[TokenId]) -> Result<()> {
224 let prefix_id = PrefixId::from(prefix_tokens);
225 let mut prefixes = self.prefixes.write();
226
227 if let Some(cached_prefix) = prefixes.get_mut(&prefix_id) {
228 cached_prefix.remove_ref()?;
229
230 if cached_prefix.ref_count == 0 {
232 prefixes.remove(&prefix_id);
233 debug!(
234 "Removed unreferenced prefix: {} tokens",
235 prefix_tokens.len()
236 );
237 }
238 }
239
240 Ok(())
241 }
242
243 fn evict_lru(&self, prefixes: &mut HashMap<PrefixId, CachedPrefix>) {
245 let mut oldest_id = None;
246 let mut oldest_time = None;
247
248 for (prefix_id, cached_prefix) in prefixes.iter() {
250 if cached_prefix.can_evict() {
251 if let Some(current_oldest) = oldest_time {
252 if cached_prefix.last_access < current_oldest {
253 oldest_time = Some(cached_prefix.last_access);
254 oldest_id = Some(prefix_id.clone());
255 }
256 } else {
257 oldest_time = Some(cached_prefix.last_access);
258 oldest_id = Some(prefix_id.clone());
259 }
260 }
261 }
262
263 if oldest_id.is_none() {
265 for (prefix_id, cached_prefix) in prefixes.iter() {
266 if let Some(current_oldest) = oldest_time {
267 if cached_prefix.last_access < current_oldest {
268 oldest_time = Some(cached_prefix.last_access);
269 oldest_id = Some(prefix_id.clone());
270 }
271 } else {
272 oldest_time = Some(cached_prefix.last_access);
273 oldest_id = Some(prefix_id.clone());
274 }
275 }
276 }
277
278 if let Some(prefix_id) = oldest_id {
279 prefixes.remove(&prefix_id);
280 *self.evictions.lock() += 1;
281 debug!("Evicted LRU prefix: {} tokens", prefix_id.len());
282 }
283 }
284
285 pub fn evict_n(&self, n: usize) -> usize {
287 let mut prefixes = self.prefixes.write();
288 let mut evicted = 0;
289
290 for _ in 0..n {
291 if prefixes.is_empty() {
292 break;
293 }
294 self.evict_lru(&mut prefixes);
295 evicted += 1;
296 }
297
298 evicted
299 }
300
301 pub fn stats(&self) -> PrefixCacheStats {
303 let hits = *self.hits.lock();
305 let misses = *self.misses.lock();
306 let evictions = *self.evictions.lock();
307
308 let prefixes = self.prefixes.read();
309 let total_size: usize = prefixes.values().map(|p| p.size).sum();
310 let active_prefixes = prefixes.len();
311 drop(prefixes); PrefixCacheStats {
314 hits,
315 misses,
316 evictions,
317 active_prefixes,
318 total_cached_tokens: total_size,
319 hit_rate: {
320 if hits + misses > 0 {
321 hits as f32 / (hits + misses) as f32
322 } else {
323 0.0
324 }
325 },
326 }
327 }
328
329 pub fn clear(&self) {
331 let mut prefixes = self.prefixes.write();
332 prefixes.clear();
333 *self.hits.lock() = 0;
334 *self.misses.lock() = 0;
335 *self.evictions.lock() = 0;
336 debug!("Cleared prefix cache");
337 }
338
339 pub fn config(&self) -> (usize, usize) {
341 (self.max_prefixes, self.min_prefix_length)
342 }
343}
344
345impl Default for PrefixCache {
346 fn default() -> Self {
347 Self::new(100, 8) }
349}
350
351#[derive(Debug, Clone)]
353pub struct PrefixCacheStats {
354 pub hits: usize,
355 pub misses: usize,
356 pub evictions: usize,
357 pub active_prefixes: usize,
358 pub total_cached_tokens: usize,
359 pub hit_rate: f32,
360}
361
362#[cfg(test)]
363mod tests {
364 use super::*;
365
366 #[derive(Debug, Clone)]
368 struct MockKvHandle {
369 tokens: usize,
370 device: ferrum_types::Device,
371 block_table: ferrum_interfaces::BlockTable,
372 }
373
374 impl MockKvHandle {
375 fn new(tokens: usize) -> Self {
376 Self {
377 tokens,
378 device: ferrum_types::Device::CPU,
379 block_table: ferrum_interfaces::BlockTable::new(16),
380 }
381 }
382 }
383
384 impl ferrum_interfaces::KvCacheHandle for MockKvHandle {
385 fn block_table(&self) -> &ferrum_interfaces::BlockTable {
386 &self.block_table
387 }
388
389 fn block_table_mut(&mut self) -> &mut ferrum_interfaces::BlockTable {
390 &mut self.block_table
391 }
392
393 fn as_any(&self) -> &dyn std::any::Any {
394 self
395 }
396
397 fn device(&self) -> ferrum_types::Device {
398 self.device.clone()
399 }
400
401 fn num_tokens(&self) -> usize {
402 self.tokens
403 }
404
405 fn num_layers(&self) -> usize {
406 32
407 }
408
409 fn num_heads(&self) -> usize {
410 32
411 }
412
413 fn head_dim(&self) -> usize {
414 128
415 }
416
417 fn key_cache(
418 &self,
419 _layer: usize,
420 ) -> ferrum_types::Result<Option<ferrum_interfaces::TensorRef>> {
421 Ok(None)
422 }
423
424 fn value_cache(
425 &self,
426 _layer: usize,
427 ) -> ferrum_types::Result<Option<ferrum_interfaces::TensorRef>> {
428 Ok(None)
429 }
430
431 fn clone_handle(&self) -> ferrum_types::Result<Arc<dyn ferrum_interfaces::KvCacheHandle>> {
432 Ok(Arc::new(Self {
433 tokens: self.tokens,
434 device: self.device.clone(),
435 block_table: self.block_table.clone(),
436 }))
437 }
438
439 fn stats(&self) -> ferrum_interfaces::kv_cache::CacheHandleStats {
440 ferrum_interfaces::kv_cache::CacheHandleStats {
441 memory_bytes: 0,
442 blocks_allocated: 0,
443 tokens_stored: self.tokens,
444 utilization: 0.0,
445 last_access: std::time::Instant::now(),
446 }
447 }
448
449 fn is_valid(&self) -> bool {
450 true
451 }
452
453 fn cache_id(&self) -> String {
454 "mock".to_string()
455 }
456 }
457
458 #[test]
459 fn test_prefix_cache_creation() {
460 let cache = PrefixCache::new(50, 4);
461 let (max_prefixes, min_len) = cache.config();
462 assert_eq!(max_prefixes, 50);
463 assert_eq!(min_len, 4);
464 }
465
466 #[test]
467 fn test_prefix_storage_and_retrieval() {
468 let cache = PrefixCache::new(10, 2);
469
470 let tokens = vec![TokenId::new(1), TokenId::new(2), TokenId::new(3)];
471 let handle = Arc::new(MockKvHandle::new(3));
472
473 cache
475 .store_prefix(&tokens, handle.clone(), vec![0.1; 10])
476 .unwrap();
477
478 let result = cache.find_prefix(&tokens);
480 assert!(result.is_some());
481
482 let longer_tokens = vec![
484 TokenId::new(1),
485 TokenId::new(2),
486 TokenId::new(3),
487 TokenId::new(4),
488 ];
489 let result = cache.find_prefix(&longer_tokens);
490 assert!(result.is_some());
491 let (found_prefix, _, _) = result.unwrap();
492 assert_eq!(found_prefix.tokens(), &tokens);
493 }
494
495 #[test]
496 fn test_prefix_length_filtering() {
497 let cache = PrefixCache::new(10, 5); let short_tokens = vec![TokenId::new(1), TokenId::new(2)]; let handle = Arc::new(MockKvHandle::new(2));
501
502 cache
504 .store_prefix(&short_tokens, handle, vec![0.1; 10])
505 .unwrap();
506
507 let result = cache.find_prefix(&short_tokens);
508 assert!(result.is_none());
509 }
510
511 #[test]
512 fn test_lru_eviction() {
513 let cache = PrefixCache::new(2, 1); let tokens1 = vec![TokenId::new(1)];
516 let tokens2 = vec![TokenId::new(2)];
517 let tokens3 = vec![TokenId::new(3)];
518
519 let handle = Arc::new(MockKvHandle::new(1));
520
521 cache
523 .store_prefix(&tokens1, handle.clone(), vec![0.1; 10])
524 .unwrap();
525 cache
526 .store_prefix(&tokens2, handle.clone(), vec![0.1; 10])
527 .unwrap();
528
529 cache.find_prefix(&tokens1);
531
532 cache
534 .store_prefix(&tokens3, handle.clone(), vec![0.1; 10])
535 .unwrap();
536
537 assert!(cache.find_prefix(&tokens1).is_some());
539 assert!(cache.find_prefix(&tokens2).is_none());
540 assert!(cache.find_prefix(&tokens3).is_some());
541 }
542
543 #[test]
544 fn test_cache_stats() {
545 let cache = PrefixCache::new(10, 2);
546 let tokens = vec![TokenId::new(1), TokenId::new(2)];
547 let handle = Arc::new(MockKvHandle::new(2));
548
549 cache.store_prefix(&tokens, handle, vec![0.1; 10]).unwrap();
550
551 cache.find_prefix(&tokens);
553
554 let other_tokens = vec![TokenId::new(3), TokenId::new(4)];
556 cache.find_prefix(&other_tokens);
557
558 let stats = cache.stats();
559 assert_eq!(stats.hits, 1);
560 assert_eq!(stats.misses, 1);
561 assert_eq!(stats.hit_rate, 0.5);
562 assert_eq!(stats.active_prefixes, 1);
563 }
564}