Skip to main content

vtcode_core/core/
prompt_caching.rs

1use crate::config::constants::prompt_cache;
2use crate::config::core::PromptCachingConfig;
3use crate::llm::provider::{Message, MessageContent, MessageRole};
4use hashbrown::HashMap;
5use serde::{Deserialize, Serialize};
6use std::fmt::Write;
7use std::path::{Path, PathBuf};
8use tokio::fs;
9use vtcode_commons::VtCodePaths;
10use vtcode_commons::utils::current_timestamp;
11
12use crate::utils::tokens::estimate_tokens;
13
14/// Cached prompt entry
15#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct CachedPrompt {
17    /// SHA-256 hash of the original prompt.
18    pub prompt_hash: String,
19    /// The original prompt text.
20    pub original_prompt: String,
21    /// The optimized prompt text produced by the LLM.
22    pub optimized_prompt: String,
23    /// Model identifier used for optimization.
24    pub model_used: String,
25    /// Estimated number of tokens saved by the optimization.
26    pub tokens_saved: Option<u32>,
27    /// Quality score assigned during optimization (0.0 to 1.0).
28    pub quality_score: Option<f64>,
29    /// Unix timestamp when the entry was created.
30    pub created_at: u64,
31    /// Unix timestamp of the most recent access.
32    pub last_used: u64,
33    /// Number of times this cached entry has been retrieved.
34    pub usage_count: u32,
35}
36
37/// Prompt caching configuration
38#[derive(Debug, Clone, Serialize, Deserialize)]
39pub struct PromptCacheConfig {
40    /// Whether prompt caching is enabled.
41    pub enabled: bool,
42    /// Directory where cache files are stored.
43    pub cache_dir: PathBuf,
44    /// Maximum number of entries the cache will hold.
45    pub max_cache_size: usize,
46    /// Maximum age of cache entries in days before expiration.
47    pub max_age_days: u64,
48    /// Whether expired entries are cleaned up automatically on startup.
49    pub enable_auto_cleanup: bool,
50    /// Minimum quality score required to store a cached entry.
51    pub min_quality_threshold: f64,
52}
53
54impl Default for PromptCacheConfig {
55    fn default() -> Self {
56        let cache_dir = default_cache_dir();
57        Self {
58            enabled: prompt_cache::DEFAULT_ENABLED && !cache_dir.as_os_str().is_empty(),
59            cache_dir,
60            max_cache_size: prompt_cache::DEFAULT_MAX_ENTRIES,
61            max_age_days: prompt_cache::DEFAULT_MAX_AGE_DAYS,
62            enable_auto_cleanup: prompt_cache::DEFAULT_AUTO_CLEANUP,
63            min_quality_threshold: prompt_cache::DEFAULT_MIN_QUALITY_THRESHOLD,
64        }
65    }
66}
67
68impl PromptCacheConfig {
69    /// Build runtime configuration from high-level TOML settings
70    pub fn from_settings(settings: &PromptCachingConfig, workspace_root: Option<&Path>) -> Self {
71        let cache_dir = settings.resolve_cache_dir(workspace_root);
72        Self {
73            enabled: settings.enabled && !cache_dir.as_os_str().is_empty(),
74            cache_dir,
75            max_cache_size: settings.max_entries,
76            max_age_days: settings.max_age_days,
77            enable_auto_cleanup: settings.enable_auto_cleanup,
78            min_quality_threshold: settings.min_quality_threshold,
79        }
80    }
81
82    /// Returns `true` if prompt caching is enabled.
83    pub fn is_enabled(&self) -> bool {
84        self.enabled
85    }
86}
87
88fn default_cache_dir() -> PathBuf {
89    VtCodePaths::resolve()
90        .map(|paths| paths.cache_dir().join("prompts"))
91        .unwrap_or_default()
92}
93
94/// Prompt caching system
95pub struct PromptCache {
96    config: PromptCacheConfig,
97    cache: HashMap<String, CachedPrompt>,
98    dirty: bool,
99}
100
101impl PromptCache {
102    /// Create a new prompt cache with default configuration.
103    pub async fn new() -> Self {
104        Self::with_config(PromptCacheConfig::default()).await
105    }
106
107    /// Create a new prompt cache with the given configuration.
108    pub async fn with_config(config: PromptCacheConfig) -> Self {
109        let mut config = config;
110        if config.cache_dir.as_os_str().is_empty() {
111            config.enabled = false;
112        }
113        let mut cache = Self { config, cache: HashMap::new(), dirty: false };
114
115        // Load existing cache
116        if cache.config.enabled {
117            let _ = cache.load_cache().await;
118
119            // Auto cleanup if enabled
120            if cache.config.enable_auto_cleanup {
121                let _ = cache.cleanup_expired();
122            }
123        }
124
125        cache
126    }
127
128    /// Get cached optimized prompt
129    pub fn get(&mut self, prompt_hash: &str) -> Option<CachedPrompt> {
130        if !self.config.enabled {
131            return None;
132        }
133        self.cache.get_mut(prompt_hash).map(|entry| {
134            entry.last_used = current_timestamp();
135            entry.usage_count += 1;
136            self.dirty = true;
137            entry.clone()
138        })
139    }
140
141    /// Store optimized prompt in cache
142    pub fn put(&mut self, entry: CachedPrompt) -> Result<(), PromptCacheError> {
143        if !self.config.enabled {
144            return Ok(());
145        }
146        // Check quality threshold
147        if entry
148            .quality_score
149            .is_some_and(|quality| quality < self.config.min_quality_threshold)
150        {
151            return Ok(()); // Don't cache low-quality entries
152        }
153
154        // Check cache size limit
155        if self.cache.len() >= self.config.max_cache_size {
156            self.evict_oldest()?;
157        }
158
159        self.cache.insert(entry.prompt_hash.clone(), entry);
160        self.dirty = true;
161
162        Ok(())
163    }
164
165    /// Check if prompt is cached
166    pub fn contains(&self, prompt_hash: &str) -> bool {
167        self.config.enabled && self.cache.contains_key(prompt_hash)
168    }
169
170    /// Get cache statistics
171    pub fn stats(&self) -> CacheStats {
172        if !self.config.enabled {
173            return CacheStats::default();
174        }
175        let total_entries = self.cache.len();
176        let total_usage = self.cache.values().map(|e| e.usage_count).sum::<u32>();
177        let total_tokens_saved = self.cache.values().filter_map(|e| e.tokens_saved).sum::<u32>();
178        let avg_quality = if !self.cache.is_empty() {
179            self.cache.values().filter_map(|e| e.quality_score).sum::<f64>() / self.cache.len() as f64
180        } else {
181            0.0
182        };
183
184        CacheStats {
185            total_entries,
186            total_usage,
187            total_tokens_saved,
188            avg_quality,
189        }
190    }
191
192    /// Clear all cache entries
193    pub async fn clear(&mut self) -> Result<(), PromptCacheError> {
194        if !self.config.enabled {
195            return Ok(());
196        }
197        self.cache.clear();
198        self.dirty = true;
199        self.save_cache().await
200    }
201
202    /// Generate hash for prompt
203    pub fn hash_prompt(prompt: &str) -> String {
204        vtcode_commons::utils::calculate_sha256(prompt.as_bytes())
205    }
206
207    /// Save cache to disk
208    pub async fn save_cache(&self) -> Result<(), PromptCacheError> {
209        if !self.config.enabled || self.config.cache_dir.as_os_str().is_empty() || !self.dirty {
210            return Ok(());
211        }
212
213        // Canonical VT Code cache roots are private and symlink-safe. A
214        // project/user override retains the existing create_dir_all behavior
215        // so workspace-local cache paths remain compatible.
216        let is_global_cache = VtCodePaths::resolve()
217            .ok()
218            .is_some_and(|paths| self.config.cache_dir.starts_with(paths.cache_dir()));
219        if is_global_cache {
220            VtCodePaths::ensure_user_dir(&self.config.cache_dir)
221                .map_err(|error| PromptCacheError::Io(std::io::Error::other(error)))?;
222        } else {
223            fs::create_dir_all(&self.config.cache_dir).await.map_err(PromptCacheError::Io)?;
224        }
225
226        let cache_path = self.config.cache_dir.join("prompt_cache.json");
227        let data = serde_json::to_string_pretty(&self.cache).map_err(PromptCacheError::Serialization)?;
228
229        if is_global_cache {
230            vtcode_commons::fs::write_private_file_atomic(&cache_path, data.as_bytes())
231                .await
232                .map_err(|error| PromptCacheError::Io(std::io::Error::other(error)))?;
233        } else {
234            fs::write(cache_path, data).await.map_err(PromptCacheError::Io)?;
235        }
236
237        Ok(())
238    }
239
240    /// Load cache from disk
241    async fn load_cache(&mut self) -> Result<(), PromptCacheError> {
242        if !self.config.enabled || self.config.cache_dir.as_os_str().is_empty() {
243            return Ok(());
244        }
245        let cache_path = self.config.cache_dir.join("prompt_cache.json");
246
247        if !fs::try_exists(&cache_path).await.unwrap_or(false) {
248            return Ok(());
249        }
250
251        let data = if VtCodePaths::resolve()
252            .ok()
253            .is_some_and(|paths| self.config.cache_dir.starts_with(paths.cache_dir()))
254        {
255            String::from_utf8(
256                vtcode_commons::fs::read_private_file_no_follow(&cache_path)
257                    .await
258                    .map_err(|error| PromptCacheError::Io(std::io::Error::other(error.to_string())))?,
259            )
260            .map_err(|error| PromptCacheError::Io(std::io::Error::other(error)))?
261        } else {
262            fs::read_to_string(cache_path).await.map_err(PromptCacheError::Io)?
263        };
264
265        self.cache = serde_json::from_str(&data).map_err(PromptCacheError::Serialization)?;
266
267        Ok(())
268    }
269
270    /// Clean up expired cache entries
271    fn cleanup_expired(&mut self) -> Result<(), PromptCacheError> {
272        if !self.config.enabled {
273            return Ok(());
274        }
275        let now = current_timestamp();
276        let max_age_seconds = self.config.max_age_days * 24 * 60 * 60;
277
278        self.cache.retain(|_, entry| now - entry.created_at < max_age_seconds);
279
280        self.dirty = true;
281        Ok(())
282    }
283
284    /// Evict oldest cache entries when cache is full
285    fn evict_oldest(&mut self) -> Result<(), PromptCacheError> {
286        if !self.config.enabled {
287            return Ok(());
288        }
289        if self.cache.is_empty() {
290            return Ok(());
291        }
292
293        // Find the oldest entry
294        let Some(oldest_key) = self
295            .cache
296            .iter()
297            .min_by_key(|(_, entry)| entry.last_used)
298            .map(|(key, _)| key.clone())
299        else {
300            return Ok(());
301        };
302
303        self.cache.remove(&oldest_key);
304        self.dirty = true;
305
306        Ok(())
307    }
308}
309
310// Note: Drop trait cannot be async, so we remove automatic save on drop.
311// Users must explicitly call save_cache() or use a wrapper that handles this.
312
313/// Cache statistics
314#[derive(Debug, Clone, Serialize, Deserialize)]
315pub struct CacheStats {
316    /// Total number of entries currently in the cache.
317    pub total_entries: usize,
318    /// Sum of all entries' usage counts.
319    pub total_usage: u32,
320    /// Total tokens saved across all cached prompts.
321    pub total_tokens_saved: u32,
322    /// Average quality score across all cached entries.
323    pub avg_quality: f64,
324}
325
326impl Default for CacheStats {
327    fn default() -> Self {
328        Self {
329            total_entries: 0,
330            total_usage: 0,
331            total_tokens_saved: 0,
332            avg_quality: 0.0,
333        }
334    }
335}
336
337/// Prompt cache errors
338#[derive(Debug, thiserror::Error)]
339pub enum PromptCacheError {
340    #[error("IO error: {0}")]
341    Io(#[from] std::io::Error),
342
343    #[error("Serialization error: {0}")]
344    Serialization(#[from] serde_json::Error),
345
346    #[error("Cache full")]
347    CacheFull,
348}
349
350/// Prompt optimizer that uses caching
351pub struct PromptOptimizer {
352    cache: PromptCache,
353    llm_provider: Box<dyn crate::llm::provider::LLMProvider>,
354}
355
356impl PromptOptimizer {
357    /// Create a new prompt optimizer with the given LLM provider.
358    pub async fn new(llm_provider: Box<dyn crate::llm::provider::LLMProvider>) -> Self {
359        Self { cache: PromptCache::new().await, llm_provider }
360    }
361
362    /// Replace the default cache with a custom [`PromptCache`] instance.
363    pub fn with_cache(mut self, cache: PromptCache) -> Self {
364        self.cache = cache;
365        self
366    }
367
368    /// Explicitly save the cache to disk
369    pub async fn save_cache(&self) -> Result<(), PromptCacheError> {
370        self.cache.save_cache().await
371    }
372
373    /// Optimize a prompt using caching
374    pub async fn optimize_prompt(
375        &mut self,
376        original_prompt: &str,
377        target_model: &str,
378        context: Option<&str>,
379    ) -> Result<String, PromptOptimizationError> {
380        let prompt_hash = PromptCache::hash_prompt(original_prompt);
381
382        // Check cache first
383        if let Some(cached) = self.cache.get(&prompt_hash) {
384            return Ok(cached.optimized_prompt);
385        }
386
387        // Generate optimized prompt
388        let optimized = self.generate_optimized_prompt(original_prompt, target_model, context).await?;
389
390        // Calculate tokens saved (rough estimate)
391        let original_tokens = estimate_tokens(original_prompt);
392        let optimized_tokens = estimate_tokens(&optimized);
393        let tokens_saved = original_tokens.saturating_sub(optimized_tokens);
394
395        // Create cache entry
396        let entry = CachedPrompt {
397            prompt_hash: prompt_hash.clone(),
398            original_prompt: original_prompt.to_string(),
399            optimized_prompt: optimized.clone(),
400            model_used: target_model.to_string(),
401            tokens_saved: Some(tokens_saved.min(u32::MAX as usize) as u32),
402            quality_score: Some(0.8), // Placeholder quality score
403            created_at: current_timestamp(),
404            last_used: current_timestamp(),
405            usage_count: 1,
406        };
407
408        // Store in cache
409        self.cache.put(entry)?;
410
411        Ok(optimized)
412    }
413
414    /// Generate optimized prompt using LLM
415    async fn generate_optimized_prompt(
416        &self,
417        original_prompt: &str,
418        target_model: &str,
419        context: Option<&str>,
420    ) -> Result<String, PromptOptimizationError> {
421        let system_prompt = format!(
422            "You are an expert prompt engineer. Your task is to optimize prompts for {target_model} \
423             to make them more effective, clearer, and more likely to produce high-quality responses. \
424             Focus on improving clarity, specificity, structure, and effectiveness while preserving \
425             the original intent and requirements."
426        );
427
428        let mut user_prompt = format!(
429            "Please optimize the following prompt for {target_model}:\n\nORIGINAL PROMPT:\n{original_prompt}\n\n"
430        );
431
432        if let Some(ctx) = context {
433            let _ = write!(user_prompt, "CONTEXT:\n{ctx}\n\n");
434        }
435
436        user_prompt.push_str(
437            "OPTIMIZATION REQUIREMENTS:\n\
438             1. Make the prompt clearer and more specific\n\
439             2. Improve structure and formatting\n\
440             3. Add relevant context or examples if helpful\n\
441             4. Ensure the prompt is appropriate for the target model\n\
442             5. Maintain the original intent and requirements\n\
443             6. Keep the optimized prompt concise but comprehensive\n\n\
444             Provide only the optimized prompt without any explanation or additional text.",
445        );
446
447        let request = crate::llm::provider::LLMRequest {
448            messages: std::sync::Arc::new(vec![
449                Message {
450                    role: MessageRole::System,
451                    content: MessageContent::Text(system_prompt),
452                    ..Default::default()
453                },
454                Message {
455                    role: MessageRole::User,
456                    content: MessageContent::Text(user_prompt),
457                    ..Default::default()
458                },
459            ]),
460            model: target_model.to_string(),
461            max_tokens: Some(2000),
462            temperature: Some(0.3),
463            ..Default::default()
464        };
465
466        let response = self
467            .llm_provider
468            .generate(request)
469            .await
470            .map_err(|e| PromptOptimizationError::LLMError(e.to_string()))?;
471
472        Ok(response.content.unwrap_or_else(|| original_prompt.to_string()))
473    }
474
475    /// Get cache statistics
476    pub fn cache_stats(&self) -> CacheStats {
477        self.cache.stats()
478    }
479
480    /// Clear cache
481    pub async fn clear_cache(&mut self) -> Result<(), PromptCacheError> {
482        self.cache.clear().await
483    }
484}
485
486/// Prompt optimization errors
487#[derive(Debug, thiserror::Error)]
488pub enum PromptOptimizationError {
489    #[error("LLM error: {0}")]
490    LLMError(String),
491
492    #[error("Cache error: {0}")]
493    CacheError(#[from] PromptCacheError),
494}
495
496#[cfg(test)]
497mod tests {
498    use super::*;
499
500    #[test]
501    fn test_prompt_hash() {
502        let prompt = "Test prompt";
503        let hash1 = PromptCache::hash_prompt(prompt);
504        let hash2 = PromptCache::hash_prompt(prompt);
505        assert_eq!(hash1, hash2);
506        assert!(!hash1.is_empty());
507    }
508
509    #[tokio::test]
510    async fn test_cache_operations() {
511        let mut cache = PromptCache::new().await;
512
513        let entry = CachedPrompt {
514            prompt_hash: "test_hash".to_owned(),
515            original_prompt: "original".to_owned(),
516            optimized_prompt: "optimized".to_owned(),
517            model_used: crate::config::constants::models::google::GEMINI_3_FLASH_PREVIEW.to_owned(),
518            tokens_saved: Some(100),
519            quality_score: Some(0.9),
520            created_at: 1000,
521            last_used: 1000,
522            usage_count: 0,
523        };
524
525        cache.put(entry).unwrap();
526        assert!(cache.contains("test_hash"));
527
528        let retrieved = cache.get("test_hash");
529        assert!(retrieved.is_some());
530        assert_eq!(retrieved.unwrap().usage_count, 1);
531    }
532
533    #[tokio::test]
534    async fn disabled_cache_config_is_no_op() {
535        let settings = PromptCachingConfig {
536            enabled: false,
537            cache_dir: "relative/cache".to_owned(),
538            ..PromptCachingConfig::default()
539        };
540        let cfg = PromptCacheConfig::from_settings(&settings, None);
541        assert!(!cfg.is_enabled());
542
543        let mut cache = PromptCache::with_config(cfg).await;
544        assert!(!cache.contains("missing"));
545        assert_eq!(cache.stats().total_entries, 0);
546
547        let entry = CachedPrompt {
548            prompt_hash: "noop".to_owned(),
549            original_prompt: "original".to_owned(),
550            optimized_prompt: "optimized".to_owned(),
551            model_used: crate::config::constants::models::google::GEMINI_3_FLASH_PREVIEW.to_owned(),
552            tokens_saved: Some(10),
553            quality_score: Some(0.9),
554            created_at: 1,
555            last_used: 1,
556            usage_count: 0,
557        };
558
559        cache.put(entry).unwrap();
560        assert!(!cache.contains("noop"));
561        assert_eq!(cache.stats().total_entries, 0);
562    }
563}