1#![allow(dead_code)]
7
8use std::collections::HashMap;
9use std::sync::{Arc, Mutex};
10
11use serde::{Deserialize, Serialize};
12use thiserror::Error;
13use tracing::{debug, info, warn};
14
15#[derive(Debug, Error)]
21pub enum PromptOptimizerError {
22 #[error("variant generation produced zero variants")]
24 NoVariants,
25 #[error("all variant inferences failed: {0}")]
27 AllInferencesFailed(String),
28 #[error("internal lock poisoned")]
30 LockPoisoned,
31}
32
33#[derive(Debug, Clone, Serialize, Deserialize)]
35pub enum QualityMetric {
36 ResponseLength { max_chars: usize },
38 KeywordPresence { keywords: Vec<String> },
40 JsonValidity,
42 JsonKeyPresence { required_keys: Vec<String> },
44 LengthWindow { min_chars: usize, max_chars: usize },
46}
47
48impl QualityMetric {
49 #[must_use]
51 pub fn score(&self, response: &str) -> f64 {
52 match self {
53 QualityMetric::ResponseLength { max_chars } => {
54 if *max_chars == 0 { return 0.0; }
55 (response.len() as f64 / *max_chars as f64).min(1.0)
56 }
57 QualityMetric::KeywordPresence { keywords } => {
58 let lower = response.to_lowercase();
59 if keywords.iter().any(|kw| lower.contains(kw.to_lowercase().as_str())) { 1.0 } else { 0.0 }
60 }
61 QualityMetric::JsonValidity => {
62 if serde_json::from_str::<serde_json::Value>(response).is_ok() { 1.0 } else { 0.0 }
63 }
64 QualityMetric::JsonKeyPresence { required_keys } => {
65 if required_keys.is_empty() { return 1.0; }
66 match serde_json::from_str::<serde_json::Value>(response) {
67 Ok(serde_json::Value::Object(map)) => {
68 let found = required_keys.iter().filter(|k| map.contains_key(k.as_str())).count();
69 found as f64 / required_keys.len() as f64
70 }
71 _ => 0.0,
72 }
73 }
74 QualityMetric::LengthWindow { min_chars, max_chars } => {
75 let len = response.len();
76 if len >= *min_chars && len <= *max_chars { 1.0 } else { 0.0 }
77 }
78 }
79 }
80}
81
82#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct ScoringEngine {
85 pub metrics: Vec<(QualityMetric, f64)>,
87}
88
89impl ScoringEngine {
90 #[must_use]
92 pub fn new(metrics: Vec<(QualityMetric, f64)>) -> Self { Self { metrics } }
93
94 #[must_use]
96 pub fn score(&self, response: &str) -> f64 {
97 if self.metrics.is_empty() { return 0.0; }
98 let total_weight: f64 = self.metrics.iter().map(|(_, w)| w.abs()).sum();
99 if total_weight == 0.0 { return 0.0; }
100 let weighted_sum: f64 = self.metrics.iter()
101 .map(|(metric, weight)| metric.score(response) * weight.abs())
102 .sum();
103 (weighted_sum / total_weight).clamp(0.0, 1.0)
104 }
105}
106
107impl Default for ScoringEngine {
108 fn default() -> Self {
109 Self::new(vec![
110 (QualityMetric::ResponseLength { max_chars: 2000 }, 0.6),
111 (QualityMetric::KeywordPresence { keywords: vec!["answer".to_string(), "result".to_string()] }, 0.4),
112 ])
113 }
114}
115
116#[derive(Debug, Clone, Serialize, Deserialize)]
118pub enum VariantStrategy {
119 InstructionPrefix,
121 ClosingSuffix,
123 Reframe,
125 CustomPrefixes(Vec<String>),
127}
128
129pub struct VariantGenerator {
131 strategy: VariantStrategy,
132 num_variants: usize,
133}
134
135impl VariantGenerator {
136 #[must_use]
138 pub fn new(strategy: VariantStrategy, num_variants: usize) -> Self {
139 Self { strategy, num_variants: num_variants.clamp(1, 32) }
140 }
141
142 #[must_use]
144 pub fn generate(&self, base_prompt: &str) -> Vec<String> {
145 match &self.strategy {
146 VariantStrategy::InstructionPrefix => {
147 let prefixes = ["", "Please answer concisely. ", "Think step by step. ",
148 "Provide a detailed explanation. ", "Answer as an expert. ",
149 "Be direct and precise. ", "Use bullet points. ", "Explain to a beginner. "];
150 prefixes.iter().take(self.num_variants).map(|p| format!("{p}{base_prompt}")).collect()
151 }
152 VariantStrategy::ClosingSuffix => {
153 let suffixes = ["", "\nBe concise.", "\nProvide examples.", "\nBe thorough.", "\nSummarise at the end."];
154 suffixes.iter().take(self.num_variants).map(|s| format!("{base_prompt}{s}")).collect()
155 }
156 VariantStrategy::Reframe => {
157 let frames = [
158 base_prompt.to_string(),
159 format!("Regarding the following: {base_prompt}\nWhat is the best answer?"),
160 format!("I need help with: {base_prompt}"),
161 format!("Question: {base_prompt}\nAnswer:"),
162 format!("Context: {base_prompt}\nProvide a clear response."),
163 ];
164 frames.into_iter().take(self.num_variants).collect()
165 }
166 VariantStrategy::CustomPrefixes(prefixes) => {
167 prefixes.iter().take(self.num_variants).map(|p| format!("{p}{base_prompt}")).collect()
168 }
169 }
170 }
171}
172
173#[derive(Debug, Clone, Serialize, Deserialize)]
175pub struct AbExperimentResult {
176 pub base_prompt: String,
178 pub intent: String,
180 pub variants: Vec<VariantScore>,
182 pub winner_index: usize,
184 pub winner_score: f64,
186}
187
188#[derive(Debug, Clone, Serialize, Deserialize)]
190pub struct VariantScore {
191 pub prompt: String,
193 pub response: String,
195 pub score: f64,
197}
198
199#[derive(Debug)]
201pub struct PromoterRegistry {
202 inner: Mutex<PromoterInner>,
203}
204
205#[derive(Debug)]
206struct PromoterInner {
207 map: HashMap<String, (String, f64)>,
208 max_intents: usize,
209}
210
211impl PromoterRegistry {
212 #[must_use]
214 pub fn new(max_intents: usize) -> Self {
215 Self { inner: Mutex::new(PromoterInner { map: HashMap::new(), max_intents: max_intents.max(1) }) }
216 }
217
218 pub fn promote(&self, intent: &str, variant: String, score: f64) {
220 let Ok(mut inner) = self.inner.lock() else {
221 warn!("PromoterRegistry lock poisoned; skipping promote");
222 return;
223 };
224 let should_insert = match inner.map.get(intent) {
225 Some((_, existing_score)) => score > *existing_score,
226 None => {
227 if inner.map.len() >= inner.max_intents {
228 warn!(max = inner.max_intents, "PromoterRegistry at capacity");
229 return;
230 }
231 true
232 }
233 };
234 if should_insert {
235 info!(intent, score, "promoting new best prompt variant");
236 inner.map.insert(intent.to_string(), (variant, score));
237 }
238 }
239
240 #[must_use]
242 pub fn best_variant(&self, intent: &str) -> Option<String> {
243 let Ok(inner) = self.inner.lock() else { return None; };
244 inner.map.get(intent).map(|(v, _)| v.clone())
245 }
246
247 #[must_use]
249 pub fn snapshot(&self) -> Vec<(String, String, f64)> {
250 let Ok(inner) = self.inner.lock() else { return vec![]; };
251 inner.map.iter().map(|(intent, (variant, score))| (intent.clone(), variant.clone(), *score)).collect()
252 }
253}
254
255impl Default for PromoterRegistry {
256 fn default() -> Self { Self::new(10_000) }
257}
258
259#[derive(Debug, Clone, Serialize, Deserialize)]
261pub struct AbOptimizerConfig {
262 pub strategy: VariantStrategy,
264 pub num_variants: usize,
266 pub auto_promote: bool,
268 pub max_intents: usize,
270}
271
272impl Default for AbOptimizerConfig {
273 fn default() -> Self {
274 Self { strategy: VariantStrategy::InstructionPrefix, num_variants: 4, auto_promote: true, max_intents: 10_000 }
275 }
276}
277
278fn derive_intent(prompt: &str) -> String {
279 prompt.chars().take(64).collect::<String>().to_lowercase().trim().to_string()
280}
281
282pub struct PromptAbOptimizer {
284 config: AbOptimizerConfig,
285 scoring: ScoringEngine,
286 registry: Arc<PromoterRegistry>,
287 generator: VariantGenerator,
288}
289
290impl PromptAbOptimizer {
291 #[must_use]
293 pub fn new(config: AbOptimizerConfig, scoring: ScoringEngine, registry: Arc<PromoterRegistry>) -> Self {
294 let generator = VariantGenerator::new(config.strategy.clone(), config.num_variants);
295 Self { config, scoring, registry, generator }
296 }
297
298 pub async fn run<F, Fut>(&self, base_prompt: &str, infer_fn: F) -> Result<AbExperimentResult, PromptOptimizerError>
300 where
301 F: Fn(String) -> Fut,
302 Fut: std::future::Future<Output = Result<String, String>>,
303 {
304 let variants = self.generator.generate(base_prompt);
305 if variants.is_empty() { return Err(PromptOptimizerError::NoVariants); }
306 let mut scored: Vec<VariantScore> = Vec::with_capacity(variants.len());
307 let mut errors: Vec<String> = Vec::new();
308 for variant in variants {
309 match infer_fn(variant.clone()).await {
310 Ok(response) => {
311 let score = self.scoring.score(&response);
312 scored.push(VariantScore { prompt: variant, response, score });
313 }
314 Err(e) => { warn!(error = %e, "variant inference failed"); errors.push(e); }
315 }
316 }
317 if scored.is_empty() { return Err(PromptOptimizerError::AllInferencesFailed(errors.join("; "))); }
318 let winner_index = scored.iter().enumerate()
319 .max_by(|(_, a), (_, b)| a.score.partial_cmp(&b.score).unwrap_or(std::cmp::Ordering::Equal))
320 .map(|(i, _)| i).unwrap_or(0);
321 let winner_score = scored[winner_index].score;
322 let intent = derive_intent(base_prompt);
323 if self.config.auto_promote {
324 self.registry.promote(&intent, scored[winner_index].prompt.clone(), winner_score);
325 }
326 Ok(AbExperimentResult { base_prompt: base_prompt.to_string(), intent, variants: scored, winner_index, winner_score })
327 }
328
329 #[must_use]
331 pub fn best_variant_for(&self, prompt: &str) -> Option<String> {
332 let intent = derive_intent(prompt);
333 self.registry.best_variant(&intent)
334 }
335}
336
337#[derive(Debug, Clone)]
343pub struct PromptVariant {
344 pub id: String,
346 pub template: String,
348 pub performance_score: f64,
350 pub sample_count: u64,
352 pub avg_cost: f64,
354 pub avg_tokens_out: u64,
356 pub created_at: u64,
358 total_cost: f64,
360 total_tokens_out: u64,
361 successes: u64,
363}
364
365impl PromptVariant {
366 fn new(id: String, template: String) -> Self {
367 Self {
368 id, template, performance_score: 0.0, sample_count: 0,
369 avg_cost: 0.0, avg_tokens_out: 0, created_at: 0,
370 total_cost: 0.0, total_tokens_out: 0, successes: 0,
371 }
372 }
373}
374
375#[derive(Debug, Clone)]
377pub enum OptimizationStrategy {
378 UCB1 {
380 exploration: f64,
382 },
383 EpsilonGreedy {
385 epsilon: f64,
387 },
388 ThompsonSampling {
390 alpha: f64,
392 beta_param: f64,
394 },
395 BestFirst,
397}
398
399pub struct PromptOptimizer {
401 pub variants: Vec<PromptVariant>,
403 pub strategy: OptimizationStrategy,
405 pub total_trials: u64,
407 next_id: u64,
408}
409
410impl PromptOptimizer {
411 pub fn new(strategy: OptimizationStrategy) -> Self {
413 Self { variants: Vec::new(), strategy, total_trials: 0, next_id: 0 }
414 }
415
416 pub fn add_variant(&mut self, template: &str) -> String {
418 let id = format!("variant-{}", self.next_id);
419 self.next_id += 1;
420 self.variants.push(PromptVariant::new(id.clone(), template.to_string()));
421 id
422 }
423
424 pub fn select(&mut self, rng_seed: u64) -> Option<&PromptVariant> {
429 if self.variants.is_empty() { return None; }
430
431 let unsampled = self.variants.iter().position(|v| v.sample_count == 0);
433
434 let selected_idx = match &self.strategy {
435 OptimizationStrategy::UCB1 { exploration } => {
436 if let Some(idx) = unsampled { idx }
437 else {
438 let total_ln = (self.total_trials as f64).ln().max(0.0);
439 let exploration = *exploration;
440 self.variants.iter().enumerate()
441 .map(|(i, v)| {
442 let mean = v.performance_score;
443 let ucb = mean + exploration * (2.0 * total_ln / v.sample_count as f64).sqrt();
444 (i, ucb)
445 })
446 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
447 .map(|(i, _)| i)
448 .unwrap_or(0)
449 }
450 }
451 OptimizationStrategy::EpsilonGreedy { epsilon } => {
452 let epsilon = *epsilon;
453 let rand_val = lcg_rand(rng_seed + self.total_trials) as f64 / u64::MAX as f64;
455 if rand_val < epsilon {
456 let rand_idx = lcg_rand(rng_seed.wrapping_add(self.total_trials).wrapping_add(1337));
458 (rand_idx as usize) % self.variants.len()
459 } else {
460 self.variants.iter().enumerate()
462 .max_by(|(_, a), (_, b)| a.performance_score.partial_cmp(&b.performance_score)
463 .unwrap_or(std::cmp::Ordering::Equal))
464 .map(|(i, _)| i).unwrap_or(0)
465 }
466 }
467 OptimizationStrategy::ThompsonSampling { alpha, beta_param } => {
468 let alpha = *alpha;
469 let beta_p = *beta_param;
470 self.variants.iter().enumerate()
473 .map(|(i, v)| {
474 let s = v.successes as f64;
475 let f = (v.sample_count - v.successes) as f64;
476 let a = alpha + s;
477 let b = beta_p + f;
478 let mean = a / (a + b);
480 let noise_seed = lcg_rand(rng_seed.wrapping_add(self.total_trials).wrapping_add(i as u64));
481 let noise = (noise_seed as f64 / u64::MAX as f64 - 0.5) * 0.1;
482 let sample = (mean + noise).clamp(0.0, 1.0);
483 (i, sample)
484 })
485 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
486 .map(|(i, _)| i).unwrap_or(0)
487 }
488 OptimizationStrategy::BestFirst => {
489 if let Some(idx) = unsampled { idx }
491 else {
492 self.variants.iter().enumerate()
493 .max_by(|(_, a), (_, b)| a.performance_score.partial_cmp(&b.performance_score)
494 .unwrap_or(std::cmp::Ordering::Equal))
495 .map(|(i, _)| i).unwrap_or(0)
496 }
497 }
498 };
499
500 self.total_trials += 1;
501 self.variants.get(selected_idx)
502 }
503
504 pub fn record_feedback(&mut self, variant_id: &str, score: f64, cost: f64, tokens_out: u64) {
506 if let Some(v) = self.variants.iter_mut().find(|v| v.id == variant_id) {
507 let n = v.sample_count as f64;
508 v.performance_score = (v.performance_score * n + score) / (n + 1.0);
510 v.sample_count += 1;
511 v.total_cost += cost;
512 v.avg_cost = v.total_cost / v.sample_count as f64;
513 v.total_tokens_out += tokens_out;
514 v.avg_tokens_out = v.total_tokens_out / v.sample_count;
515 if score > 0.5 { v.successes += 1; }
516 debug!(variant_id, score, "prompt_optimizer: feedback recorded");
517 }
518 }
519
520 pub fn best_variant(&self) -> Option<&PromptVariant> {
522 self.variants.iter()
523 .filter(|v| v.sample_count > 0)
524 .max_by(|a, b| a.performance_score.partial_cmp(&b.performance_score).unwrap_or(std::cmp::Ordering::Equal))
525 }
526
527 pub fn convergence_ratio(&self) -> f64 {
529 self.best_variant().map(|v| v.performance_score).unwrap_or(0.0).clamp(0.0, 1.0)
530 }
531}
532
533pub struct PromptMutator;
535
536impl PromptMutator {
537 pub fn paraphrase(template: &str, seed: u64) -> String {
541 let thesaurus: HashMap<&str, Vec<&str>> = [
542 ("quickly", vec!["rapidly", "swiftly", "promptly"]),
543 ("help", vec!["assist", "support", "aid"]),
544 ("make", vec!["create", "generate", "produce"]),
545 ("show", vec!["display", "present", "demonstrate"]),
546 ("use", vec!["utilize", "employ", "apply"]),
547 ("good", vec!["excellent", "effective", "optimal"]),
548 ("bad", vec!["poor", "ineffective", "suboptimal"]),
549 ("big", vec!["large", "substantial", "significant"]),
550 ("small", vec!["minimal", "compact", "concise"]),
551 ("important", vec!["critical", "essential", "key"]),
552 ].iter().cloned().collect();
553
554 let mut result = template.to_string();
555 let mut replacements = 0;
556 let mut rng = seed;
557 for (word, synonyms) in &thesaurus {
558 if replacements >= 3 { break; }
559 let lower = result.to_lowercase();
561 if let Some(pos) = lower.find(word) {
562 rng = lcg_rand(rng);
563 let syn_idx = (rng as usize) % synonyms.len();
564 let synonym = synonyms[syn_idx];
565 let original_char = result.chars().nth(pos).unwrap_or('a');
567 let replacement = if original_char.is_uppercase() {
568 let mut s = synonym.to_string();
569 if let Some(c) = s.get_mut(0..1) { c.make_ascii_uppercase(); }
570 s
571 } else {
572 synonym.to_string()
573 };
574 result = format!("{}{}{}", &result[..pos], replacement, &result[pos + word.len()..]);
575 replacements += 1;
576 }
577 }
578 result
579 }
580
581 pub fn add_instruction(template: &str, instruction: &str) -> String {
583 format!("{instruction}\n\n{template}")
584 }
585
586 pub fn trim_to_budget(template: &str, max_tokens: usize) -> String {
588 let max_words = ((max_tokens as f64) / 1.3) as usize;
589 let words: Vec<&str> = template.split_whitespace().collect();
590 if words.len() <= max_words {
591 template.to_string()
592 } else {
593 words[..max_words].join(" ")
594 }
595 }
596}
597
598fn lcg_rand(seed: u64) -> u64 {
604 seed.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1_442_695_040_888_963_407)
605}
606
607#[cfg(test)]
612#[allow(clippy::unwrap_used, clippy::expect_used)]
613mod tests {
614 use super::*;
615
616 #[test]
617 fn ucb1_selects_unsampled_first() {
618 let mut opt = PromptOptimizer::new(OptimizationStrategy::UCB1 { exploration: 1.414 });
619 let id_a = opt.add_variant("template A");
620 let id_b = opt.add_variant("template B");
621 let id_c = opt.add_variant("template C");
622
623 let selected = opt.select(42).unwrap();
625 assert_eq!(selected.id, id_a);
626 opt.record_feedback(&id_a, 0.5, 0.01, 100);
627
628 let selected = opt.select(42).unwrap();
629 assert_eq!(selected.id, id_b);
630 opt.record_feedback(&id_b, 0.5, 0.01, 100);
631
632 let selected = opt.select(42).unwrap();
633 assert_eq!(selected.id, id_c);
634 }
635
636 #[test]
637 fn epsilon_greedy_explores_at_epsilon_rate() {
638 let mut opt = PromptOptimizer::new(OptimizationStrategy::EpsilonGreedy { epsilon: 1.0 });
639 let _id_a = opt.add_variant("template A");
640 let _id_b = opt.add_variant("template B");
641 opt.record_feedback("variant-0", 0.9, 0.01, 100);
643 opt.record_feedback("variant-1", 0.1, 0.01, 100);
644
645 let mut selections: HashMap<String, u64> = HashMap::new();
647 for i in 0..100u64 {
648 if let Some(v) = opt.select(i * 1000) {
650 *selections.entry(v.id.clone()).or_default() += 1;
651 }
652 }
653 assert!(selections.contains_key("variant-0") || selections.contains_key("variant-1"));
655 }
656
657 #[test]
658 fn best_variant_returns_highest_score() {
659 let mut opt = PromptOptimizer::new(OptimizationStrategy::BestFirst);
660 let _id_a = opt.add_variant("template A");
661 let _id_b = opt.add_variant("template B");
662 let _id_c = opt.add_variant("template C");
663
664 opt.record_feedback("variant-0", 0.3, 0.01, 100);
665 opt.record_feedback("variant-1", 0.9, 0.01, 100);
666 opt.record_feedback("variant-2", 0.6, 0.01, 100);
667
668 let best = opt.best_variant().unwrap();
669 assert_eq!(best.id, "variant-1", "variant-1 has the highest score");
670 }
671
672 #[test]
673 fn best_first_selects_best_after_sampling() {
674 let mut opt = PromptOptimizer::new(OptimizationStrategy::BestFirst);
675 let _id_a = opt.add_variant("low score template");
676 let _id_b = opt.add_variant("high score template");
677
678 opt.record_feedback("variant-0", 0.2, 0.01, 50);
680 opt.record_feedback("variant-1", 0.95, 0.01, 50);
681
682 let selected = opt.select(0).unwrap();
684 assert_eq!(selected.id, "variant-1");
685 }
686
687 #[test]
688 fn mutator_paraphrase_changes_output() {
689 let template = "Please help me quickly make a good solution";
690 let paraphrased = PromptMutator::paraphrase(template, 42);
691 assert_ne!(paraphrased, template, "paraphrase should modify the template");
693 }
694
695 #[test]
696 fn mutator_add_instruction_prepends() {
697 let result = PromptMutator::add_instruction("Do the task.", "Be concise.");
698 assert!(result.starts_with("Be concise.\n\n"), "instruction should be prepended");
699 assert!(result.contains("Do the task."));
700 }
701
702 #[test]
703 fn mutator_trim_to_budget_truncates() {
704 let template = "one two three four five six seven eight nine ten";
705 let trimmed = PromptMutator::trim_to_budget(template, 5);
707 let word_count = trimmed.split_whitespace().count();
708 assert!(word_count <= 4, "trimmed template should have at most 4 words, got {word_count}");
709 }
710
711 #[test]
712 fn convergence_ratio_reflects_best_score() {
713 let mut opt = PromptOptimizer::new(OptimizationStrategy::BestFirst);
714 let _id = opt.add_variant("template");
715 opt.record_feedback("variant-0", 0.75, 0.01, 100);
716 assert!((opt.convergence_ratio() - 0.75).abs() < 0.001);
717 }
718
719 #[test]
722 fn scoring_engine_length_metric() {
723 let engine = ScoringEngine::new(vec![(QualityMetric::ResponseLength { max_chars: 100 }, 1.0)]);
724 assert!((engine.score("hello") - 0.05).abs() < 0.01);
725 assert!((engine.score(&"x".repeat(100)) - 1.0).abs() < 0.001);
726 }
727
728 #[test]
729 fn scoring_engine_keyword_metric() {
730 let engine = ScoringEngine::new(vec![(QualityMetric::KeywordPresence { keywords: vec!["answer".to_string()] }, 1.0)]);
731 assert_eq!(engine.score("The answer is 42"), 1.0);
732 assert_eq!(engine.score("I don't know"), 0.0);
733 }
734
735 #[test]
736 fn promoter_registry_promotes_better_score() {
737 let reg = Arc::new(PromoterRegistry::default());
738 reg.promote("test intent", "variant A".to_string(), 0.5);
739 reg.promote("test intent", "variant B".to_string(), 0.9);
740 reg.promote("test intent", "variant C".to_string(), 0.3);
741 assert_eq!(reg.best_variant("test intent").as_deref(), Some("variant B"));
742 }
743}
744
745#[derive(Clone, Debug)]
756pub enum CompressionGoal {
757 MinimizeTokens,
759 MaximizeClarity,
761 BalancedCostQuality,
763}
764
765#[derive(Clone, Debug)]
771pub struct FewShotExample {
772 pub input: String,
774 pub output: String,
776 pub tokens: usize,
778 pub relevance_score: f64,
780}
781
782#[derive(Debug, Default)]
784pub struct FewShotSelector {
785 examples: Vec<FewShotExample>,
786 max_examples: usize,
787}
788
789impl FewShotSelector {
790 pub fn new(max_examples: usize) -> Self {
792 Self { examples: Vec::new(), max_examples }
793 }
794
795 pub fn add_example(&mut self, input: impl Into<String>, output: impl Into<String>, tokens: usize) {
797 self.examples.push(FewShotExample {
798 input: input.into(),
799 output: output.into(),
800 tokens,
801 relevance_score: 0.0,
802 });
803 }
804
805 pub fn select_relevant<'a>(&'a self, query: &str, token_budget: usize) -> Vec<&'a FewShotExample> {
808 let mut scored: Vec<(f64, &FewShotExample)> = self
809 .examples
810 .iter()
811 .map(|ex| (Self::jaccard(query, &ex.input), ex))
812 .collect();
813 scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
814
815 let mut result = Vec::new();
816 let mut used_tokens = 0usize;
817 for (_, ex) in scored {
818 if result.len() >= self.max_examples {
819 break;
820 }
821 if used_tokens + ex.tokens > token_budget {
822 continue;
823 }
824 used_tokens += ex.tokens;
825 result.push(ex);
826 }
827 result
828 }
829
830 fn jaccard(a: &str, b: &str) -> f64 {
832 let words_a: std::collections::HashSet<&str> = a.split_whitespace().collect();
833 let words_b: std::collections::HashSet<&str> = b.split_whitespace().collect();
834 if words_a.is_empty() && words_b.is_empty() {
835 return 1.0;
836 }
837 let intersection = words_a.intersection(&words_b).count();
838 let union = words_a.union(&words_b).count();
839 if union == 0 { 0.0 } else { intersection as f64 / union as f64 }
840 }
841}
842
843#[derive(Debug, Clone)]
849pub struct PromptCompressor {
850 pub max_tokens: usize,
852 pub compression_ratio: f64,
854}
855
856impl PromptCompressor {
857 pub fn new(max_tokens: usize, compression_ratio: f64) -> Self {
859 Self { max_tokens, compression_ratio: compression_ratio.clamp(0.0, 1.0) }
860 }
861
862 pub fn compress(&self, text: &str) -> String {
868 let mut out = collapse_whitespace(text);
870
871 out = dedup_sentences(&out);
873
874 if let Some(re) = filler_regex() {
877 out = re.replace_all(&out, " ").into_owned();
878 out = collapse_whitespace(&out);
879 if let Some(space_punct) = space_before_punct_regex() {
880 out = space_punct.replace_all(&out, "$1").into_owned();
881 }
882 }
883
884 out = out.replace("in order to", "to")
886 .replace("due to the fact that", "because")
887 .replace("for the purpose of", "for")
888 .replace("at this point in time", "now")
889 .replace("in the event that", "if");
890
891 collapse_whitespace(&out)
893 }
894
895 pub fn estimate_savings(&self, original: &str, compressed: &str) -> f64 {
897 if original.is_empty() {
898 return 0.0;
899 }
900 let saved = original.len().saturating_sub(compressed.len());
901 saved as f64 / original.len() as f64
902 }
903}
904
905fn filler_regex() -> Option<&'static regex::Regex> {
908 static RE: std::sync::OnceLock<Option<regex::Regex>> = std::sync::OnceLock::new();
909 RE.get_or_init(|| {
910 regex::Regex::new(
911 r"(?i)\b(?:please note that|it is worth noting that|it should be noted that|as previously mentioned|as mentioned above|note that|please|kindly)\b",
912 )
913 .ok()
914 })
915 .as_ref()
916}
917
918fn space_before_punct_regex() -> Option<&'static regex::Regex> {
919 static RE: std::sync::OnceLock<Option<regex::Regex>> = std::sync::OnceLock::new();
920 RE.get_or_init(|| regex::Regex::new(r" +([.,;:!?])").ok()).as_ref()
921}
922
923fn collapse_whitespace(s: &str) -> String {
924 let mut out = String::with_capacity(s.len());
925 let mut prev_space = false;
926 for ch in s.chars() {
927 if ch.is_whitespace() {
928 if !prev_space {
929 out.push(' ');
930 }
931 prev_space = true;
932 } else {
933 out.push(ch);
934 prev_space = false;
935 }
936 }
937 out.trim().to_string()
938}
939
940fn dedup_sentences(s: &str) -> String {
941 let sentences: Vec<&str> = s.split(". ").collect();
942 let mut result: Vec<&str> = Vec::with_capacity(sentences.len());
943 for sentence in &sentences {
944 if result.last().is_none_or(|last| *last != *sentence) {
945 result.push(sentence);
946 }
947 }
948 result.join(". ")
949}
950
951#[derive(Debug, Clone)]
957pub struct ChainOfThoughtInjector {
958 pub cot_prefix: String,
960 pub cot_suffix: String,
962}
963
964impl ChainOfThoughtInjector {
965 pub fn new() -> Self {
967 Self {
968 cot_prefix: "Let's think step by step:".to_string(),
969 cot_suffix: "Therefore, the answer is:".to_string(),
970 }
971 }
972
973 pub fn inject(&self, prompt: &str) -> String {
975 format!("{}\n\n{}\n{}", self.cot_prefix, prompt, self.cot_suffix)
976 }
977
978 pub fn inject_numbered_steps(&self, prompt: &str, n_steps: u32) -> String {
980 let steps: Vec<String> = (1..=n_steps)
981 .map(|i| format!("Step {}: [reasoning here]", i))
982 .collect();
983 format!("{}\n\n{}\n\n{}\n{}", self.cot_prefix, prompt, steps.join("\n"), self.cot_suffix)
984 }
985}
986
987impl Default for ChainOfThoughtInjector {
988 fn default() -> Self {
989 Self::new()
990 }
991}
992
993pub struct PromptCompressionOptimizer {
1000 pub compressor: PromptCompressor,
1002 pub few_shot: FewShotSelector,
1004 pub cot: ChainOfThoughtInjector,
1006 pub goal: CompressionGoal,
1008}
1009
1010impl PromptCompressionOptimizer {
1011 pub fn new(goal: CompressionGoal) -> Self {
1013 Self {
1014 compressor: PromptCompressor::new(4096, 0.8),
1015 few_shot: FewShotSelector::new(3),
1016 cot: ChainOfThoughtInjector::new(),
1017 goal,
1018 }
1019 }
1020
1021 pub fn optimize(&self, prompt: &str, token_budget: usize) -> String {
1028 let compressed = self.compressor.compress(prompt);
1030
1031 let base_tokens = compressed.len() / 4 + 1;
1033 let remaining = token_budget.saturating_sub(base_tokens);
1034
1035 let examples = self.few_shot.select_relevant(&compressed, remaining);
1037 let mut result = if !examples.is_empty() {
1038 let shots: Vec<String> = examples
1039 .iter()
1040 .map(|ex| format!("Input: {}\nOutput: {}", ex.input, ex.output))
1041 .collect();
1042 format!("{}\n\n{}", shots.join("\n\n"), compressed)
1043 } else {
1044 compressed
1045 };
1046
1047 match self.goal {
1049 CompressionGoal::MinimizeTokens => {}
1050 CompressionGoal::MaximizeClarity | CompressionGoal::BalancedCostQuality => {
1051 result = self.cot.inject(&result);
1052 }
1053 }
1054
1055 result
1056 }
1057
1058 pub fn optimize_batch(&self, prompts: &[String], token_budget: usize) -> Vec<String> {
1060 prompts.iter().map(|p| self.optimize(p, token_budget)).collect()
1061 }
1062}
1063
1064#[cfg(test)]
1065mod round32_tests {
1066 use super::*;
1067
1068 #[test]
1069 fn compressor_removes_fillers() {
1070 let c = PromptCompressor::new(4096, 0.8);
1071 let out = c.compress("Please summarise this kindly.");
1072 assert!(!out.to_lowercase().contains("please"));
1073 assert!(!out.to_lowercase().contains("kindly"));
1074 assert_eq!(out, "summarise this.");
1075 }
1076
1077 #[test]
1078 fn compressor_keeps_words_containing_fillers() {
1079 let c = PromptCompressor::new(4096, 0.8);
1080 assert_eq!(c.compress("This will displease unkindly users."), "This will displease unkindly users.");
1081 }
1082
1083 #[test]
1084 fn compressor_savings_positive() {
1085 let c = PromptCompressor::new(4096, 0.8);
1086 let orig = "Please please please note that in order to do this please kindly proceed.";
1087 let comp = c.compress(orig);
1088 assert!(c.estimate_savings(orig, &comp) >= 0.0);
1089 }
1090
1091 #[test]
1092 fn few_shot_respects_budget() {
1093 let mut sel = FewShotSelector::new(5);
1094 sel.add_example("what is 2+2", "4", 10);
1095 sel.add_example("what is 3+3", "6", 10);
1096 sel.add_example("what is 4+4", "8", 10);
1097 let chosen = sel.select_relevant("what is 5+5", 15);
1099 assert!(chosen.len() <= 1);
1100 }
1101
1102 #[test]
1103 fn cot_inject_contains_prefix_and_suffix() {
1104 let inj = ChainOfThoughtInjector::new();
1105 let out = inj.inject("What is 2+2?");
1106 assert!(out.contains("Let's think step by step:"));
1107 assert!(out.contains("Therefore, the answer is:"));
1108 assert!(out.contains("What is 2+2?"));
1109 }
1110
1111 #[test]
1112 fn cot_numbered_steps() {
1113 let inj = ChainOfThoughtInjector::new();
1114 let out = inj.inject_numbered_steps("Solve x+1=5", 3);
1115 assert!(out.contains("Step 1:"));
1116 assert!(out.contains("Step 2:"));
1117 assert!(out.contains("Step 3:"));
1118 }
1119
1120 #[test]
1121 fn optimizer_minimize_no_cot() {
1122 let opt = PromptCompressionOptimizer::new(CompressionGoal::MinimizeTokens);
1123 let out = opt.optimize("Please kindly answer this question.", 1000);
1124 assert!(!out.contains("Let's think step by step:"));
1125 }
1126
1127 #[test]
1128 fn optimizer_batch() {
1129 let opt = PromptCompressionOptimizer::new(CompressionGoal::BalancedCostQuality);
1130 let prompts = vec!["Hello world".to_string(), "Please answer this kindly.".to_string()];
1131 let results = opt.optimize_batch(&prompts, 500);
1132 assert_eq!(results.len(), 2);
1133 }
1134}