Skip to main content

ferrum_kv/cache/
prefix.rs

1//! Prefix caching for shared prompt optimization
2
3use ferrum_types::{FerrumError, Result, TokenId};
4use parking_lot::RwLock;
5use std::collections::HashMap;
6use std::sync::Arc;
7use tracing::{debug, trace};
8
9/// Prefix identifier
10#[derive(Debug, Clone, PartialEq, Eq, Hash)]
11pub struct PrefixId(Vec<TokenId>);
12
13impl PrefixId {
14    /// Create new prefix ID from tokens
15    pub fn new(tokens: Vec<TokenId>) -> Self {
16        Self(tokens)
17    }
18
19    /// Get tokens
20    pub fn tokens(&self) -> &[TokenId] {
21        &self.0
22    }
23
24    /// Get length
25    pub fn len(&self) -> usize {
26        self.0.len()
27    }
28
29    /// Check if empty
30    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/// Cached prefix information
48#[derive(Debug, Clone)]
49pub struct CachedPrefix {
50    /// Prefix tokens
51    pub prefix_id: PrefixId,
52    /// KV cache handle for this prefix
53    pub kv_handle: Arc<dyn ferrum_interfaces::KvCacheHandle + Send + Sync>,
54    /// Last-token logits from the prefill — used to sample the first generated
55    /// token on a cache hit without re-running the executor.
56    pub last_logits: Vec<f32>,
57    /// Reference count
58    pub ref_count: usize,
59    /// Last access time
60    pub last_access: std::time::Instant,
61    /// Size in tokens
62    pub size: usize,
63}
64
65impl CachedPrefix {
66    /// Create new cached prefix
67    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    /// Add reference
84    pub fn add_ref(&mut self) {
85        self.ref_count += 1;
86        self.touch();
87    }
88
89    /// Remove reference
90    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    /// Update access time
101    pub fn touch(&mut self) {
102        self.last_access = std::time::Instant::now();
103    }
104
105    /// Check if can be evicted
106    pub fn can_evict(&self) -> bool {
107        self.ref_count == 0
108    }
109}
110
111/// Prefix cache for shared prompt optimization
112#[derive(Debug)]
113pub struct PrefixCache {
114    /// Cached prefixes by prefix ID
115    prefixes: RwLock<HashMap<PrefixId, CachedPrefix>>,
116    /// Maximum number of cached prefixes
117    max_prefixes: usize,
118    /// Minimum prefix length to cache
119    min_prefix_length: usize,
120    /// Statistics
121    hits: parking_lot::Mutex<usize>,
122    misses: parking_lot::Mutex<usize>,
123    evictions: parking_lot::Mutex<usize>,
124}
125
126impl PrefixCache {
127    /// Create new prefix cache
128    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    /// Find matching prefix for given tokens.
140    ///
141    /// Returns `(PrefixId, KvCacheHandle, last_logits)` for the longest
142    /// matching prefix, or `None` on miss.
143    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        // Find longest matching prefix
158        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            // Update access time
177            drop(prefixes); // Release read lock
178            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    /// Store prefix in cache
191    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(()); // Don't cache short prefixes
199        }
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        // Check if we need to evict
207        if prefixes.len() >= self.max_prefixes && !prefixes.contains_key(&prefix_id) {
208            self.evict_lru(&mut prefixes);
209        }
210
211        // Store or update prefix
212        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    /// Remove reference to prefix
223    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            // Remove if no more references
231            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    /// Evict least recently used prefix
244    fn evict_lru(&self, prefixes: &mut HashMap<PrefixId, CachedPrefix>) {
245        let mut oldest_id = None;
246        let mut oldest_time = None;
247
248        // First try to find least recently used prefix with ref_count == 0
249        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 no evictable prefix found, evict LRU regardless of ref_count
264        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    /// Evict up to n prefixes, returning the number actually evicted
286    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    /// Get cache statistics
302    pub fn stats(&self) -> PrefixCacheStats {
303        // Lock metrics first, then prefixes to avoid potential lock ordering issues
304        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); // Release read lock as soon as possible
312
313        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    /// Clear all cached prefixes
330    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    /// Get configuration
340    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) // Default: cache up to 100 prefixes, minimum 8 tokens
348    }
349}
350
351/// Prefix cache statistics
352#[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    // Mock KV cache handle for testing
367    #[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        // Store prefix
474        cache
475            .store_prefix(&tokens, handle.clone(), vec![0.1; 10])
476            .unwrap();
477
478        // Should find exact match
479        let result = cache.find_prefix(&tokens);
480        assert!(result.is_some());
481
482        // Should find prefix for longer sequence
483        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); // Minimum 5 tokens
498
499        let short_tokens = vec![TokenId::new(1), TokenId::new(2)]; // Too short
500        let handle = Arc::new(MockKvHandle::new(2));
501
502        // Should not store short prefix
503        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); // Max 2 prefixes
514
515        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        // Store 2 prefixes
522        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        // Access first one to make it more recent
530        cache.find_prefix(&tokens1);
531
532        // Store third - should evict tokens2 (LRU)
533        cache
534            .store_prefix(&tokens3, handle.clone(), vec![0.1; 10])
535            .unwrap();
536
537        // tokens1 and tokens3 should exist, tokens2 should be evicted
538        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        // Hit
552        cache.find_prefix(&tokens);
553
554        // Miss
555        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}