Skip to main content

tokio_prompt_orchestrator/
cache_warmer.rs

1//! Cache Pre-Warming
2//!
3//! Implements a cache pre-warming system that generates and tracks warming jobs
4//! based on common query patterns, ensuring the prompt cache is populated with
5//! frequently-requested content before real traffic arrives.
6//!
7//! ## Key Types
8//!
9//! - [`WarmingStrategy`] — controls which patterns are expanded into jobs
10//! - [`WarmingPattern`] — a parameterised template with variable substitutions
11//! - [`CacheWarmer`] — orchestrates job generation and status tracking
12
13use std::sync::atomic::{AtomicU64, Ordering};
14use std::sync::Mutex;
15use std::time::Instant;
16
17// ── Enums ────────────────────────────────────────────────────────────────────
18
19/// Strategy governing which patterns the warmer expands into jobs.
20#[derive(Debug, Clone)]
21pub enum WarmingStrategy {
22    /// Warm the top-`k` patterns ranked by [`WarmingPattern::frequency_weight`].
23    TopK { k: usize },
24    /// Expand every registered pattern fully.
25    AllTemplates,
26    /// Expand only patterns whose `template` starts with `"scheduled:"`.
27    ScheduledPatterns,
28    /// Use caller-supplied literal queries (bypasses pattern expansion).
29    UserDefined(Vec<String>),
30}
31
32/// Life-cycle status of a single warming job.
33#[derive(Debug, Clone)]
34pub enum WarmingStatus {
35    /// Waiting to be executed.
36    Pending,
37    /// Currently being executed.
38    Running,
39    /// Successfully cached.
40    Completed {
41        /// Wall-clock time at which the response was stored in cache.
42        cached_at: Instant,
43    },
44    /// Execution failed with the given reason.
45    Failed(String),
46    /// Skipped (e.g. already cached or strategy filter excluded it).
47    Skipped,
48}
49
50// ── Structs ───────────────────────────────────────────────────────────────────
51
52/// A single unit of warming work.
53#[derive(Debug, Clone)]
54pub struct WarmingJob {
55    /// Unique identifier (monotonically increasing).
56    pub id: u64,
57    /// The fully-expanded query string to warm.
58    pub query: String,
59    /// Target model identifier (e.g. `"claude-3-5-sonnet"`).
60    pub model: String,
61    /// Scheduling priority — higher values are processed first.
62    pub priority: u8,
63    /// When this job was scheduled.
64    pub scheduled_at: Instant,
65    /// Current status.
66    pub status: WarmingStatus,
67}
68
69/// A template with named variable slots and their possible values.
70///
71/// # Example
72///
73/// ```
74/// use tokio_prompt_orchestrator::cache_warmer::WarmingPattern;
75///
76/// let mut p = WarmingPattern {
77///     template: "explain {topic} in simple terms".to_string(),
78///     variables: vec![
79///         ("topic".to_string(), vec!["recursion".to_string(), "ownership".to_string()]),
80///     ],
81///     frequency_weight: 1.0,
82/// };
83/// let queries = p.expand();
84/// assert_eq!(queries.len(), 2);
85/// ```
86#[derive(Debug, Clone)]
87pub struct WarmingPattern {
88    /// Template string where `{var_name}` is replaced by each possible value.
89    pub template: String,
90    /// Variable name → list of possible values for that variable.
91    pub variables: Vec<(String, Vec<String>)>,
92    /// Relative weight for `TopK` strategy ranking (higher = more important).
93    pub frequency_weight: f64,
94}
95
96impl WarmingPattern {
97    /// Expand this pattern into concrete query strings via cartesian product.
98    ///
99    /// At most 100 queries are returned to prevent combinatorial explosion.
100    pub fn expand(&self) -> Vec<String> {
101        // Start with the template as the only candidate.
102        let mut results: Vec<String> = vec![self.template.clone()];
103
104        for (var_name, values) in &self.variables {
105            if values.is_empty() {
106                continue;
107            }
108            let placeholder = format!("{{{}}}", var_name);
109            let mut next: Vec<String> = Vec::with_capacity(results.len() * values.len());
110            'outer: for base in &results {
111                for val in values {
112                    if next.len() >= 100 {
113                        break 'outer;
114                    }
115                    next.push(base.replace(&placeholder, val));
116                }
117                if next.len() >= 100 {
118                    break;
119                }
120            }
121            results = next;
122            if results.len() >= 100 {
123                results.truncate(100);
124                break;
125            }
126        }
127
128        results
129    }
130}
131
132/// Aggregate statistics for a [`CacheWarmer`] session.
133#[derive(Debug, Clone, Default)]
134pub struct WarmingStats {
135    /// Total jobs ever generated.
136    pub total_jobs: u64,
137    /// Jobs that completed successfully.
138    pub completed: u64,
139    /// Jobs that failed.
140    pub failed: u64,
141    /// Cache entries added (equal to completed).
142    pub cache_entries_added: u64,
143    /// Approximate wall-clock ms spent on completed jobs (recorded externally).
144    pub time_spent_ms: u64,
145}
146
147/// Orchestrates cache pre-warming across a set of [`WarmingPattern`]s.
148pub struct CacheWarmer {
149    /// Registered patterns.
150    pub patterns: Vec<WarmingPattern>,
151    /// Active strategy.
152    pub strategy: WarmingStrategy,
153    /// All jobs (pending, running, done).
154    jobs: Mutex<Vec<WarmingJob>>,
155    /// Monotonic job-id counter.
156    next_id: AtomicU64,
157    // Stats counters
158    stat_total: AtomicU64,
159    stat_completed: AtomicU64,
160    stat_failed: AtomicU64,
161    stat_time_ms: AtomicU64,
162}
163
164impl CacheWarmer {
165    /// Create a new warmer with the given strategy and no patterns.
166    pub fn new(strategy: WarmingStrategy) -> Self {
167        Self {
168            patterns: Vec::new(),
169            strategy,
170            jobs: Mutex::new(Vec::new()),
171            next_id: AtomicU64::new(1),
172            stat_total: AtomicU64::new(0),
173            stat_completed: AtomicU64::new(0),
174            stat_failed: AtomicU64::new(0),
175            stat_time_ms: AtomicU64::new(0),
176        }
177    }
178
179    /// Register a new pattern.
180    pub fn add_pattern(&mut self, pattern: WarmingPattern) {
181        self.patterns.push(pattern);
182    }
183
184    /// Generate warming jobs according to the current strategy.
185    ///
186    /// Jobs are appended to the internal queue and also returned to the caller.
187    pub fn generate_jobs(&self) -> Vec<WarmingJob> {
188        let queries: Vec<String> = match &self.strategy {
189            WarmingStrategy::AllTemplates => {
190                self.patterns.iter().flat_map(|p| p.expand()).collect()
191            }
192            WarmingStrategy::TopK { k } => {
193                let mut sorted = self.patterns.clone();
194                sorted.sort_by(|a, b| {
195                    b.frequency_weight
196                        .partial_cmp(&a.frequency_weight)
197                        .unwrap_or(std::cmp::Ordering::Equal)
198                });
199                sorted.iter().take(*k).flat_map(|p| p.expand()).collect()
200            }
201            WarmingStrategy::ScheduledPatterns => self
202                .patterns
203                .iter()
204                .filter(|p| p.template.starts_with("scheduled:"))
205                .flat_map(|p| p.expand())
206                .collect(),
207            WarmingStrategy::UserDefined(queries) => queries.clone(),
208        };
209
210        let now = Instant::now();
211        let new_jobs: Vec<WarmingJob> = queries
212            .into_iter()
213            .map(|query| {
214                let id = self.next_id.fetch_add(1, Ordering::Relaxed);
215                WarmingJob {
216                    id,
217                    query,
218                    model: "default".to_string(),
219                    priority: 128,
220                    scheduled_at: now,
221                    status: WarmingStatus::Pending,
222                }
223            })
224            .collect();
225
226        self.stat_total
227            .fetch_add(new_jobs.len() as u64, Ordering::Relaxed);
228
229        let mut guard = self.jobs.lock().unwrap_or_else(|e| e.into_inner());
230        guard.append(&mut new_jobs.clone());
231
232        new_jobs
233    }
234
235    /// Mark a job as successfully completed.
236    pub fn mark_completed(&self, job_id: u64) {
237        let mut guard = self.jobs.lock().unwrap_or_else(|e| e.into_inner());
238        if let Some(job) = guard.iter_mut().find(|j| j.id == job_id) {
239            job.status = WarmingStatus::Completed {
240                cached_at: Instant::now(),
241            };
242            self.stat_completed.fetch_add(1, Ordering::Relaxed);
243        }
244    }
245
246    /// Mark a job as failed with a descriptive reason.
247    pub fn mark_failed(&self, job_id: u64, reason: &str) {
248        let mut guard = self.jobs.lock().unwrap_or_else(|e| e.into_inner());
249        if let Some(job) = guard.iter_mut().find(|j| j.id == job_id) {
250            job.status = WarmingStatus::Failed(reason.to_string());
251            self.stat_failed.fetch_add(1, Ordering::Relaxed);
252        }
253    }
254
255    /// Return clones of all pending jobs.
256    pub fn pending_jobs(&self) -> Vec<WarmingJob> {
257        let guard = self.jobs.lock().unwrap_or_else(|e| e.into_inner());
258        guard
259            .iter()
260            .filter(|j| matches!(j.status, WarmingStatus::Pending))
261            .cloned()
262            .collect()
263    }
264
265    /// Snapshot current statistics.
266    pub fn stats(&self) -> WarmingStats {
267        let completed = self.stat_completed.load(Ordering::Relaxed);
268        WarmingStats {
269            total_jobs: self.stat_total.load(Ordering::Relaxed),
270            completed,
271            failed: self.stat_failed.load(Ordering::Relaxed),
272            cache_entries_added: completed,
273            time_spent_ms: self.stat_time_ms.load(Ordering::Relaxed),
274        }
275    }
276
277    /// Record additional wall-clock time spent on warming (in milliseconds).
278    pub fn record_time_ms(&self, ms: u64) {
279        self.stat_time_ms.fetch_add(ms, Ordering::Relaxed);
280    }
281
282    /// Return a default set of common warming patterns.
283    ///
284    /// Covers the four most common LLM task families: explain, summarise,
285    /// translate, and code generation.
286    pub fn default_patterns() -> Vec<WarmingPattern> {
287        vec![
288            WarmingPattern {
289                template: "explain {topic} in simple terms".to_string(),
290                variables: vec![(
291                    "topic".to_string(),
292                    vec![
293                        "recursion".to_string(),
294                        "machine learning".to_string(),
295                        "async/await".to_string(),
296                        "the TCP handshake".to_string(),
297                        "gradient descent".to_string(),
298                    ],
299                )],
300                frequency_weight: 0.9,
301            },
302            WarmingPattern {
303                template: "summarize the following text: {text}".to_string(),
304                variables: vec![(
305                    "text".to_string(),
306                    vec![
307                        "Lorem ipsum dolor sit amet".to_string(),
308                        "The quick brown fox jumps over the lazy dog".to_string(),
309                    ],
310                )],
311                frequency_weight: 0.8,
312            },
313            WarmingPattern {
314                template: "translate \"{phrase}\" to {language}".to_string(),
315                variables: vec![
316                    (
317                        "phrase".to_string(),
318                        vec!["Hello, world!".to_string(), "Thank you".to_string()],
319                    ),
320                    (
321                        "language".to_string(),
322                        vec![
323                            "Spanish".to_string(),
324                            "French".to_string(),
325                            "German".to_string(),
326                        ],
327                    ),
328                ],
329                frequency_weight: 0.7,
330            },
331            WarmingPattern {
332                template: "write code for {task} in {language}".to_string(),
333                variables: vec![
334                    (
335                        "task".to_string(),
336                        vec![
337                            "a binary search".to_string(),
338                            "a linked list".to_string(),
339                            "a REST API endpoint".to_string(),
340                        ],
341                    ),
342                    (
343                        "language".to_string(),
344                        vec![
345                            "Rust".to_string(),
346                            "Python".to_string(),
347                            "TypeScript".to_string(),
348                        ],
349                    ),
350                ],
351                frequency_weight: 0.85,
352            },
353        ]
354    }
355}
356
357#[cfg(test)]
358mod tests {
359    use super::*;
360
361    #[test]
362    fn expand_cartesian_product() {
363        let p = WarmingPattern {
364            template: "explain {topic} to a {audience}".to_string(),
365            variables: vec![
366                (
367                    "topic".to_string(),
368                    vec!["Rust".to_string(), "Python".to_string()],
369                ),
370                (
371                    "audience".to_string(),
372                    vec!["beginner".to_string(), "expert".to_string()],
373                ),
374            ],
375            frequency_weight: 1.0,
376        };
377        let results = p.expand();
378        assert_eq!(results.len(), 4);
379        assert!(results.contains(&"explain Rust to a beginner".to_string()));
380        assert!(results.contains(&"explain Python to a expert".to_string()));
381    }
382
383    #[test]
384    fn expand_respects_100_limit() {
385        let values: Vec<String> = (0..200).map(|i| i.to_string()).collect();
386        let p = WarmingPattern {
387            template: "item {n}".to_string(),
388            variables: vec![("n".to_string(), values)],
389            frequency_weight: 1.0,
390        };
391        assert!(p.expand().len() <= 100);
392    }
393
394    #[test]
395    fn generate_jobs_and_lifecycle() {
396        let mut warmer = CacheWarmer::new(WarmingStrategy::AllTemplates);
397        warmer.add_pattern(WarmingPattern {
398            template: "ping".to_string(),
399            variables: vec![],
400            frequency_weight: 1.0,
401        });
402        let jobs = warmer.generate_jobs();
403        assert_eq!(jobs.len(), 1);
404        let id = jobs[0].id;
405        warmer.mark_completed(id);
406        let stats = warmer.stats();
407        assert_eq!(stats.completed, 1);
408        assert_eq!(stats.cache_entries_added, 1);
409        assert_eq!(warmer.pending_jobs().len(), 0);
410    }
411
412    #[test]
413    fn default_patterns_non_empty() {
414        let patterns = CacheWarmer::default_patterns();
415        assert!(!patterns.is_empty());
416        for p in &patterns {
417            assert!(!p.expand().is_empty());
418        }
419    }
420}