Skip to main content

tokio_prompt_orchestrator/
prompt_optimizer.rs

1//! # Prompt Optimizer
2//!
3//! Prompt variant testing and automatic improvement tracking using bandit
4//! algorithms: UCB1, Epsilon-Greedy, Thompson Sampling, and BestFirst.
5
6#![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// ---------------------------------------------------------------------------
16// Re-exported legacy types (preserved for backward compatibility)
17// ---------------------------------------------------------------------------
18
19/// Errors returned by the legacy prompt optimizer.
20#[derive(Debug, Error)]
21pub enum PromptOptimizerError {
22    /// No variants could be generated from the base prompt.
23    #[error("variant generation produced zero variants")]
24    NoVariants,
25    /// All parallel inference calls failed.
26    #[error("all variant inferences failed: {0}")]
27    AllInferencesFailed(String),
28    /// Internal lock was poisoned.
29    #[error("internal lock poisoned")]
30    LockPoisoned,
31}
32
33/// A configurable quality metric used to score an inference response (legacy).
34#[derive(Debug, Clone, Serialize, Deserialize)]
35pub enum QualityMetric {
36    /// Reward longer responses up to `max_chars`.
37    ResponseLength { max_chars: usize },
38    /// Score 1.0 if any keyword appears.
39    KeywordPresence { keywords: Vec<String> },
40    /// Score 1.0 if the response is valid JSON.
41    JsonValidity,
42    /// Score = fraction of required_keys present.
43    JsonKeyPresence { required_keys: Vec<String> },
44    /// Score 1.0 if length is within [min_chars, max_chars].
45    LengthWindow { min_chars: usize, max_chars: usize },
46}
47
48impl QualityMetric {
49    /// Evaluate this metric against a response string.
50    #[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/// Computes a weighted aggregate quality score from multiple metrics (legacy).
83#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct ScoringEngine {
85    /// Each entry is (metric, weight).
86    pub metrics: Vec<(QualityMetric, f64)>,
87}
88
89impl ScoringEngine {
90    /// Create a new scoring engine.
91    #[must_use]
92    pub fn new(metrics: Vec<(QualityMetric, f64)>) -> Self { Self { metrics } }
93
94    /// Score a response.
95    #[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/// Strategies for generating prompt variants (legacy).
117#[derive(Debug, Clone, Serialize, Deserialize)]
118pub enum VariantStrategy {
119    /// Prepend different instructional prefixes.
120    InstructionPrefix,
121    /// Append different closing instructions.
122    ClosingSuffix,
123    /// Reframe the request.
124    Reframe,
125    /// Use a custom set of prefix strings.
126    CustomPrefixes(Vec<String>),
127}
128
129/// Generates N prompt variants from a base prompt (legacy).
130pub struct VariantGenerator {
131    strategy: VariantStrategy,
132    num_variants: usize,
133}
134
135impl VariantGenerator {
136    /// Create a new generator.
137    #[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    /// Generate variants.
143    #[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/// Result of one A/B experiment run (legacy).
174#[derive(Debug, Clone, Serialize, Deserialize)]
175pub struct AbExperimentResult {
176    /// The base prompt tested.
177    pub base_prompt: String,
178    /// Intent label.
179    pub intent: String,
180    /// Each variant and its score.
181    pub variants: Vec<VariantScore>,
182    /// Index of the winning variant.
183    pub winner_index: usize,
184    /// Score of the winning variant.
185    pub winner_score: f64,
186}
187
188/// Score for a single prompt variant (legacy).
189#[derive(Debug, Clone, Serialize, Deserialize)]
190pub struct VariantScore {
191    /// The variant text.
192    pub prompt: String,
193    /// The model's response.
194    pub response: String,
195    /// Aggregate quality score.
196    pub score: f64,
197}
198
199/// Maps intent strings to the best known prompt variant (legacy).
200#[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    /// Create a registry capped at `max_intents` entries.
213    #[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    /// Promote a variant if it beats the current best.
219    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    /// Look up the best variant for an intent.
241    #[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    /// Return all registered intents and their scores.
248    #[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/// Configuration for the legacy prompt A/B optimizer.
260#[derive(Debug, Clone, Serialize, Deserialize)]
261pub struct AbOptimizerConfig {
262    /// Strategy for generating variants.
263    pub strategy: VariantStrategy,
264    /// Number of variants to generate per experiment.
265    pub num_variants: usize,
266    /// Whether to automatically promote the winner.
267    pub auto_promote: bool,
268    /// Maximum number of intents to track.
269    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
282/// Orchestrates prompt A/B testing (legacy).
283pub struct PromptAbOptimizer {
284    config: AbOptimizerConfig,
285    scoring: ScoringEngine,
286    registry: Arc<PromoterRegistry>,
287    generator: VariantGenerator,
288}
289
290impl PromptAbOptimizer {
291    /// Create a new optimizer.
292    #[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    /// Run an A/B experiment for the given base prompt.
299    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    /// Return the best known variant for a prompt.
330    #[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// ---------------------------------------------------------------------------
338// NEW SPEC IMPLEMENTATION
339// ---------------------------------------------------------------------------
340
341/// A single prompt variant being tracked by the optimizer.
342#[derive(Debug, Clone)]
343pub struct PromptVariant {
344    /// Unique variant identifier.
345    pub id: String,
346    /// The prompt template text.
347    pub template: String,
348    /// Running performance score (mean of recorded scores).
349    pub performance_score: f64,
350    /// Number of times this variant has been evaluated.
351    pub sample_count: u64,
352    /// Average cost per evaluation.
353    pub avg_cost: f64,
354    /// Average number of output tokens.
355    pub avg_tokens_out: u64,
356    /// Unix timestamp when the variant was created.
357    pub created_at: u64,
358    // Internal accumulators (not exposed in the public interface).
359    total_cost: f64,
360    total_tokens_out: u64,
361    /// Number of successes (score > 0.5) for Thompson Sampling.
362    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/// Strategy for selecting the next variant to evaluate.
376#[derive(Debug, Clone)]
377pub enum OptimizationStrategy {
378    /// Upper Confidence Bound — balances exploration and exploitation.
379    UCB1 {
380        /// Exploration constant (higher = more exploration). Typical: 1.0–2.0.
381        exploration: f64,
382    },
383    /// Epsilon-Greedy — exploit the best variant most of the time.
384    EpsilonGreedy {
385        /// Probability of exploring a random variant (0.0–1.0).
386        epsilon: f64,
387    },
388    /// Thompson Sampling — sample from Beta posterior per variant.
389    ThompsonSampling {
390        /// Prior alpha parameter (pseudo-successes before any data).
391        alpha: f64,
392        /// Prior beta parameter (pseudo-failures before any data).
393        beta_param: f64,
394    },
395    /// Always pick the highest-scoring variant (or first unsampled).
396    BestFirst,
397}
398
399/// Tracks prompt variants and selects the best using bandit algorithms.
400pub struct PromptOptimizer {
401    /// All registered variants.
402    pub variants: Vec<PromptVariant>,
403    /// Selection strategy.
404    pub strategy: OptimizationStrategy,
405    /// Total number of selection trials performed.
406    pub total_trials: u64,
407    next_id: u64,
408}
409
410impl PromptOptimizer {
411    /// Create a new optimizer with the given strategy.
412    pub fn new(strategy: OptimizationStrategy) -> Self {
413        Self { variants: Vec::new(), strategy, total_trials: 0, next_id: 0 }
414    }
415
416    /// Add a variant and return its generated ID.
417    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    /// Select the next variant to evaluate using the configured strategy.
425    ///
426    /// `rng_seed` is used for stochastic strategies (EpsilonGreedy, ThompsonSampling).
427    /// Returns `None` if no variants are registered.
428    pub fn select(&mut self, rng_seed: u64) -> Option<&PromptVariant> {
429        if self.variants.is_empty() { return None; }
430
431        // Always prefer an unsampled variant first (UCB1 and BestFirst benefit from this too).
432        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                // LCG pseudo-random from seed.
454                let rand_val = lcg_rand(rng_seed + self.total_trials) as f64 / u64::MAX as f64;
455                if rand_val < epsilon {
456                    // Explore: random variant.
457                    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                    // Exploit: best mean.
461                    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                // Approximate Thompson Sampling: sample Beta(alpha + s, beta + f) for each variant.
471                // Beta mean = alpha/(alpha+beta); add deterministic noise based on seed+variant_idx.
472                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                        // Approximate sample: mean + scaled noise.
479                        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                // Prefer unsampled first, then highest performance_score.
490                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    /// Record feedback for a variant after evaluation.
505    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            // Online mean update.
509            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    /// Return the variant with the highest performance score.
521    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    /// Convergence ratio: best_score / 1.0 (max possible).
528    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
533/// Applies simple transformations to prompt templates.
534pub struct PromptMutator;
535
536impl PromptMutator {
537    /// Replace 2–3 common words with synonyms and return the modified template.
538    ///
539    /// Uses a small hardcoded thesaurus. Deterministic for a given seed.
540    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            // Case-insensitive search.
560            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                // Preserve capitalization of first char.
566                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    /// Prepend an instruction to the template.
582    pub fn add_instruction(template: &str, instruction: &str) -> String {
583        format!("{instruction}\n\n{template}")
584    }
585
586    /// Truncate the template to approximately `max_tokens` tokens (1.3 tokens/word).
587    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
598// ---------------------------------------------------------------------------
599// Math helpers
600// ---------------------------------------------------------------------------
601
602/// Linear congruential generator for deterministic pseudo-randomness.
603fn 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// ---------------------------------------------------------------------------
608// Tests
609// ---------------------------------------------------------------------------
610
611#[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        // With no samples, should select in insertion order (unsampled preferred).
624        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        // Prime variant A as best.
642        opt.record_feedback("variant-0", 0.9, 0.01, 100);
643        opt.record_feedback("variant-1", 0.1, 0.01, 100);
644
645        // With epsilon=1.0, always explore (random).
646        let mut selections: HashMap<String, u64> = HashMap::new();
647        for i in 0..100u64 {
648            // Reset trial count to avoid total_trials dominating.
649            if let Some(v) = opt.select(i * 1000) {
650                *selections.entry(v.id.clone()).or_default() += 1;
651            }
652        }
653        // Both variants should have been selected.
654        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        // Force-feed scores without going through select.
679        opt.record_feedback("variant-0", 0.2, 0.01, 50);
680        opt.record_feedback("variant-1", 0.95, 0.01, 50);
681
682        // Next select should pick best (variant-1).
683        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        // Should differ from original (at least one word replaced).
692        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        // max_tokens=5 → max_words = floor(5/1.3) = 3
706        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    // --- Legacy tests ---
720
721    #[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// ===========================================================================
746// Round-32 additions: compression, few-shot selection, CoT injection,
747// and the unified PromptOptimizer facade.
748// ===========================================================================
749
750// ---------------------------------------------------------------------------
751// OptimizationGoal
752// ---------------------------------------------------------------------------
753
754/// High-level objective that guides how the compression optimizer applies its tools.
755#[derive(Clone, Debug)]
756pub enum CompressionGoal {
757    /// Reduce token count as aggressively as possible.
758    MinimizeTokens,
759    /// Preserve maximum clarity even at the cost of extra tokens.
760    MaximizeClarity,
761    /// Balance cost savings against output quality.
762    BalancedCostQuality,
763}
764
765// ---------------------------------------------------------------------------
766// FewShotExample / FewShotSelector
767// ---------------------------------------------------------------------------
768
769/// A single input/output example for few-shot prompting.
770#[derive(Clone, Debug)]
771pub struct FewShotExample {
772    /// The example input text.
773    pub input: String,
774    /// The expected output text.
775    pub output: String,
776    /// Estimated token cost for this example.
777    pub tokens: usize,
778    /// Relevance score assigned during selection (0.0–1.0).
779    pub relevance_score: f64,
780}
781
782/// Selects the most relevant few-shot examples that fit within a token budget.
783#[derive(Debug, Default)]
784pub struct FewShotSelector {
785    examples: Vec<FewShotExample>,
786    max_examples: usize,
787}
788
789impl FewShotSelector {
790    /// Create a selector that returns at most `max_examples` examples.
791    pub fn new(max_examples: usize) -> Self {
792        Self { examples: Vec::new(), max_examples }
793    }
794
795    /// Add a new candidate example.
796    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    /// Return up to `max_examples` examples sorted by Jaccard similarity to
806    /// `query`, limited to `token_budget` total tokens.
807    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    /// Jaccard similarity between the word sets of two strings.
831    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// ---------------------------------------------------------------------------
844// PromptCompressor
845// ---------------------------------------------------------------------------
846
847/// Applies lossless and near-lossless text compression to reduce token count.
848#[derive(Debug, Clone)]
849pub struct PromptCompressor {
850    /// Maximum tokens the compressed output should occupy (soft limit).
851    pub max_tokens: usize,
852    /// Target compression ratio (0.0–1.0). 0.8 means aim to keep 80% of chars.
853    pub compression_ratio: f64,
854}
855
856impl PromptCompressor {
857    /// Create a compressor with the given limits.
858    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    /// Compress `text` by:
863    /// 1. Collapsing redundant whitespace.
864    /// 2. Deduplicating consecutive identical sentences.
865    /// 3. Dropping common filler phrases.
866    /// 4. Abbreviating frequent verbose patterns.
867    pub fn compress(&self, text: &str) -> String {
868        // Step 1 – normalise whitespace.
869        let mut out = collapse_whitespace(text);
870
871        // Step 2 – deduplicate consecutive identical sentences.
872        out = dedup_sentences(&out);
873
874        // Step 3 – remove filler phrases (case-insensitive, whole words only,
875        // including at the end of a sentence).
876        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        // Step 4 – abbreviate common patterns.
885        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        // Final whitespace cleanup.
892        collapse_whitespace(&out)
893    }
894
895    /// Fraction of characters saved: `(orig - compressed) / orig`.
896    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
905/// Filler phrases dropped by [`PromptCompressor::compress`]. Longer phrases
906/// come first so "please note that" wins over "please".
907fn 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// ---------------------------------------------------------------------------
952// ChainOfThoughtInjector
953// ---------------------------------------------------------------------------
954
955/// Wraps a prompt with chain-of-thought scaffolding.
956#[derive(Debug, Clone)]
957pub struct ChainOfThoughtInjector {
958    /// Text inserted before the user prompt to invoke step-by-step reasoning.
959    pub cot_prefix: String,
960    /// Text appended after the prompt to elicit the final answer.
961    pub cot_suffix: String,
962}
963
964impl ChainOfThoughtInjector {
965    /// Create an injector with the canonical CoT framing.
966    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    /// Inject CoT scaffolding around `prompt`.
974    pub fn inject(&self, prompt: &str) -> String {
975        format!("{}\n\n{}\n{}", self.cot_prefix, prompt, self.cot_suffix)
976    }
977
978    /// Inject a numbered step scaffold with `n_steps` placeholders.
979    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
993// ---------------------------------------------------------------------------
994// PromptOptimizer (unified facade)
995// ---------------------------------------------------------------------------
996
997/// Unified facade that combines compression, few-shot selection, and CoT
998/// injection into a single optimisation pipeline.
999pub struct PromptCompressionOptimizer {
1000    /// Text compressor.
1001    pub compressor: PromptCompressor,
1002    /// Few-shot example selector.
1003    pub few_shot: FewShotSelector,
1004    /// Chain-of-thought injector.
1005    pub cot: ChainOfThoughtInjector,
1006    /// High-level optimisation objective.
1007    pub goal: CompressionGoal,
1008}
1009
1010impl PromptCompressionOptimizer {
1011    /// Create an optimizer with sensible defaults.
1012    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    /// Optimise a single prompt within the given `token_budget`.
1022    ///
1023    /// * Always applies compression.
1024    /// * Injects few-shot examples if the budget allows.
1025    /// * Adds CoT scaffolding when the goal is `MaximizeClarity` or
1026    ///   `BalancedCostQuality`.
1027    pub fn optimize(&self, prompt: &str, token_budget: usize) -> String {
1028        // 1. Compress.
1029        let compressed = self.compressor.compress(prompt);
1030
1031        // 2. Rough token estimate: 1 token ≈ 4 chars.
1032        let base_tokens = compressed.len() / 4 + 1;
1033        let remaining = token_budget.saturating_sub(base_tokens);
1034
1035        // 3. Optionally prepend few-shot examples.
1036        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        // 4. Optionally inject CoT.
1048        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    /// Optimise a batch of prompts, each within the same `token_budget`.
1059    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        // Budget of 15 should only fit one example.
1098        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}