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#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct CachedPrompt {
17 pub prompt_hash: String,
19 pub original_prompt: String,
21 pub optimized_prompt: String,
23 pub model_used: String,
25 pub tokens_saved: Option<u32>,
27 pub quality_score: Option<f64>,
29 pub created_at: u64,
31 pub last_used: u64,
33 pub usage_count: u32,
35}
36
37#[derive(Debug, Clone, Serialize, Deserialize)]
39pub struct PromptCacheConfig {
40 pub enabled: bool,
42 pub cache_dir: PathBuf,
44 pub max_cache_size: usize,
46 pub max_age_days: u64,
48 pub enable_auto_cleanup: bool,
50 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 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 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
94pub struct PromptCache {
96 config: PromptCacheConfig,
97 cache: HashMap<String, CachedPrompt>,
98 dirty: bool,
99}
100
101impl PromptCache {
102 pub async fn new() -> Self {
104 Self::with_config(PromptCacheConfig::default()).await
105 }
106
107 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 if cache.config.enabled {
117 let _ = cache.load_cache().await;
118
119 if cache.config.enable_auto_cleanup {
121 let _ = cache.cleanup_expired();
122 }
123 }
124
125 cache
126 }
127
128 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 pub fn put(&mut self, entry: CachedPrompt) -> Result<(), PromptCacheError> {
143 if !self.config.enabled {
144 return Ok(());
145 }
146 if entry
148 .quality_score
149 .is_some_and(|quality| quality < self.config.min_quality_threshold)
150 {
151 return Ok(()); }
153
154 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 pub fn contains(&self, prompt_hash: &str) -> bool {
167 self.config.enabled && self.cache.contains_key(prompt_hash)
168 }
169
170 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 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 pub fn hash_prompt(prompt: &str) -> String {
204 vtcode_commons::utils::calculate_sha256(prompt.as_bytes())
205 }
206
207 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 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 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 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 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 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#[derive(Debug, Clone, Serialize, Deserialize)]
315pub struct CacheStats {
316 pub total_entries: usize,
318 pub total_usage: u32,
320 pub total_tokens_saved: u32,
322 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#[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
350pub struct PromptOptimizer {
352 cache: PromptCache,
353 llm_provider: Box<dyn crate::llm::provider::LLMProvider>,
354}
355
356impl PromptOptimizer {
357 pub async fn new(llm_provider: Box<dyn crate::llm::provider::LLMProvider>) -> Self {
359 Self { cache: PromptCache::new().await, llm_provider }
360 }
361
362 pub fn with_cache(mut self, cache: PromptCache) -> Self {
364 self.cache = cache;
365 self
366 }
367
368 pub async fn save_cache(&self) -> Result<(), PromptCacheError> {
370 self.cache.save_cache().await
371 }
372
373 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 if let Some(cached) = self.cache.get(&prompt_hash) {
384 return Ok(cached.optimized_prompt);
385 }
386
387 let optimized = self.generate_optimized_prompt(original_prompt, target_model, context).await?;
389
390 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 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), created_at: current_timestamp(),
404 last_used: current_timestamp(),
405 usage_count: 1,
406 };
407
408 self.cache.put(entry)?;
410
411 Ok(optimized)
412 }
413
414 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 pub fn cache_stats(&self) -> CacheStats {
477 self.cache.stats()
478 }
479
480 pub async fn clear_cache(&mut self) -> Result<(), PromptCacheError> {
482 self.cache.clear().await
483 }
484}
485
486#[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}