tokio_prompt_orchestrator/
cache_warmer.rs1use std::sync::atomic::{AtomicU64, Ordering};
14use std::sync::Mutex;
15use std::time::Instant;
16
17#[derive(Debug, Clone)]
21pub enum WarmingStrategy {
22 TopK { k: usize },
24 AllTemplates,
26 ScheduledPatterns,
28 UserDefined(Vec<String>),
30}
31
32#[derive(Debug, Clone)]
34pub enum WarmingStatus {
35 Pending,
37 Running,
39 Completed {
41 cached_at: Instant,
43 },
44 Failed(String),
46 Skipped,
48}
49
50#[derive(Debug, Clone)]
54pub struct WarmingJob {
55 pub id: u64,
57 pub query: String,
59 pub model: String,
61 pub priority: u8,
63 pub scheduled_at: Instant,
65 pub status: WarmingStatus,
67}
68
69#[derive(Debug, Clone)]
87pub struct WarmingPattern {
88 pub template: String,
90 pub variables: Vec<(String, Vec<String>)>,
92 pub frequency_weight: f64,
94}
95
96impl WarmingPattern {
97 pub fn expand(&self) -> Vec<String> {
101 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#[derive(Debug, Clone, Default)]
134pub struct WarmingStats {
135 pub total_jobs: u64,
137 pub completed: u64,
139 pub failed: u64,
141 pub cache_entries_added: u64,
143 pub time_spent_ms: u64,
145}
146
147pub struct CacheWarmer {
149 pub patterns: Vec<WarmingPattern>,
151 pub strategy: WarmingStrategy,
153 jobs: Mutex<Vec<WarmingJob>>,
155 next_id: AtomicU64,
157 stat_total: AtomicU64,
159 stat_completed: AtomicU64,
160 stat_failed: AtomicU64,
161 stat_time_ms: AtomicU64,
162}
163
164impl CacheWarmer {
165 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 pub fn add_pattern(&mut self, pattern: WarmingPattern) {
181 self.patterns.push(pattern);
182 }
183
184 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 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 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 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 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 pub fn record_time_ms(&self, ms: u64) {
279 self.stat_time_ms.fetch_add(ms, Ordering::Relaxed);
280 }
281
282 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}