Skip to main content

tokio_prompt_orchestrator/
experiment_runner.rs

1//! A/B test runner with statistical significance testing.
2
3use std::collections::HashMap;
4
5/// A single variant in an experiment.
6#[derive(Debug, Clone)]
7pub struct ExperimentVariant {
8    pub name: String,
9    pub prompt_template: String,
10    pub weight: f64,
11}
12
13/// Configuration for an experiment.
14#[derive(Debug, Clone)]
15pub struct ExperimentConfig {
16    pub name: String,
17    pub variants: Vec<ExperimentVariant>,
18    pub min_samples: usize,
19    pub confidence_level: f64,
20}
21
22/// Recorded results for a single variant.
23#[derive(Debug, Clone)]
24pub struct VariantResult {
25    pub variant_name: String,
26    pub samples: Vec<f64>,
27    pub mean: f64,
28    pub std_dev: f64,
29    pub sample_count: usize,
30}
31
32impl VariantResult {
33    fn new(name: &str) -> Self {
34        Self {
35            variant_name: name.to_string(),
36            samples: Vec::new(),
37            mean: 0.0,
38            std_dev: 0.0,
39            sample_count: 0,
40        }
41    }
42
43    fn push(&mut self, value: f64) {
44        self.samples.push(value);
45        self.sample_count = self.samples.len();
46        self.mean = self.samples.iter().sum::<f64>() / self.sample_count as f64;
47        if self.sample_count > 1 {
48            let variance = self.samples.iter()
49                .map(|x| (x - self.mean).powi(2))
50                .sum::<f64>()
51                / (self.sample_count - 1) as f64;
52            self.std_dev = variance.sqrt();
53        } else {
54            self.std_dev = 0.0;
55        }
56    }
57}
58
59/// Result of a statistical significance test.
60#[derive(Debug, Clone)]
61pub struct SignificanceTest {
62    pub p_value: f64,
63    pub is_significant: bool,
64    pub effect_size: f64,
65    pub winner: Option<String>,
66}
67
68/// State of a single experiment.
69#[derive(Debug)]
70pub struct ExperimentState {
71    pub config: ExperimentConfig,
72    pub results: HashMap<String, VariantResult>,
73    pub started_at_ms: u64,
74}
75
76/// Manages multiple A/B experiments.
77pub struct ExperimentRunner {
78    pub experiments: HashMap<String, ExperimentState>,
79    next_id: u64,
80}
81
82impl ExperimentRunner {
83    /// Create a new `ExperimentRunner`.
84    pub fn new() -> Self {
85        Self {
86            experiments: HashMap::new(),
87            next_id: 1,
88        }
89    }
90
91    /// Register an experiment and return its generated ID.
92    pub fn create_experiment(&mut self, config: ExperimentConfig) -> String {
93        let id = format!("exp-{}", self.next_id);
94        self.next_id += 1;
95        let mut results = HashMap::new();
96        for v in &config.variants {
97            results.insert(v.name.clone(), VariantResult::new(&v.name));
98        }
99        let state = ExperimentState {
100            config,
101            results,
102            started_at_ms: current_time_ms(),
103        };
104        self.experiments.insert(id.clone(), state);
105        id
106    }
107
108    /// Consistently assign a variant to a user using FNV1a hashing.
109    pub fn assign_variant<'a>(&'a self, experiment_id: &str, user_id: &str) -> Option<&'a ExperimentVariant> {
110        let state = self.experiments.get(experiment_id)?;
111        if state.config.variants.is_empty() {
112            return None;
113        }
114        let total_weight: f64 = state.config.variants.iter().map(|v| v.weight).sum();
115        if total_weight <= 0.0 {
116            return None;
117        }
118        let hash = fnv1a(user_id);
119        // Map hash to [0, total_weight)
120        let position = (hash as f64 / u64::MAX as f64) * total_weight;
121        let mut cumulative = 0.0;
122        for variant in &state.config.variants {
123            cumulative += variant.weight;
124            if position < cumulative {
125                return Some(variant);
126            }
127        }
128        // Fallback to last variant
129        state.config.variants.last()
130    }
131
132    /// Record a metric value for a variant.
133    pub fn record_metric(&mut self, experiment_id: &str, variant_name: &str, value: f64) {
134        if let Some(state) = self.experiments.get_mut(experiment_id) {
135            let entry = state.results.entry(variant_name.to_string())
136                .or_insert_with(|| VariantResult::new(variant_name));
137            entry.push(value);
138        }
139    }
140
141    /// Run Welch's t-test between two samples. Returns the two-tailed p-value approximated
142    /// via the normal CDF.
143    pub fn welch_t_test(a: &[f64], b: &[f64]) -> f64 {
144        if a.len() < 2 || b.len() < 2 {
145            return 1.0;
146        }
147        let mean_a = a.iter().sum::<f64>() / a.len() as f64;
148        let mean_b = b.iter().sum::<f64>() / b.len() as f64;
149        let var_a = a.iter().map(|x| (x - mean_a).powi(2)).sum::<f64>() / (a.len() - 1) as f64;
150        let var_b = b.iter().map(|x| (x - mean_b).powi(2)).sum::<f64>() / (b.len() - 1) as f64;
151        let se = (var_a / a.len() as f64 + var_b / b.len() as f64).sqrt();
152        if se == 0.0 {
153            return if (mean_a - mean_b).abs() < 1e-12 { 1.0 } else { 0.0 };
154        }
155        let t = (mean_a - mean_b) / se;
156        // Approximate p-value via standard normal CDF (two-tailed)
157        let p = 2.0 * (1.0 - normal_cdf(t.abs()));
158        p.clamp(0.0, 1.0)
159    }
160
161    /// Cohen's d effect size: (mean_a - mean_b) / pooled_std
162    pub fn cohen_d(a: &[f64], b: &[f64]) -> f64 {
163        if a.len() < 2 || b.len() < 2 {
164            return 0.0;
165        }
166        let mean_a = a.iter().sum::<f64>() / a.len() as f64;
167        let mean_b = b.iter().sum::<f64>() / b.len() as f64;
168        let var_a = a.iter().map(|x| (x - mean_a).powi(2)).sum::<f64>() / (a.len() - 1) as f64;
169        let var_b = b.iter().map(|x| (x - mean_b).powi(2)).sum::<f64>() / (b.len() - 1) as f64;
170        let pooled_std = ((var_a + var_b) / 2.0).sqrt();
171        if pooled_std == 0.0 {
172            return 0.0;
173        }
174        (mean_a - mean_b) / pooled_std
175    }
176
177    /// Pairwise comparison of control (first variant) vs all others.
178    pub fn analyze_experiment(&self, experiment_id: &str) -> Option<SignificanceTest> {
179        let state = self.experiments.get(experiment_id)?;
180        if state.config.variants.len() < 2 {
181            return None;
182        }
183        let control_name = &state.config.variants[0].name;
184        let control = state.results.get(control_name)?;
185        if control.samples.len() < state.config.min_samples {
186            return None;
187        }
188
189        let alpha = 1.0 - state.config.confidence_level;
190        let mut best_p = 1.0_f64;
191        let mut best_d = 0.0_f64;
192        let mut winner: Option<String> = None;
193
194        for variant in state.config.variants.iter().skip(1) {
195            if let Some(vr) = state.results.get(&variant.name) {
196                if vr.samples.len() < state.config.min_samples {
197                    continue;
198                }
199                let p = Self::welch_t_test(&control.samples, &vr.samples);
200                let d = Self::cohen_d(&control.samples, &vr.samples);
201                if p < best_p {
202                    best_p = p;
203                    best_d = d;
204                    if p < alpha {
205                        // Positive d means control is better, negative means variant is better
206                        winner = if d > 0.0 {
207                            Some(control_name.clone())
208                        } else {
209                            Some(variant.name.clone())
210                        };
211                    }
212                }
213            }
214        }
215
216        Some(SignificanceTest {
217            p_value: best_p,
218            is_significant: best_p < alpha,
219            effect_size: best_d,
220            winner,
221        })
222    }
223
224    /// Generate a text report for an experiment.
225    pub fn experiment_report(&self, experiment_id: &str) -> Option<String> {
226        let state = self.experiments.get(experiment_id)?;
227        let mut out = format!("=== Experiment: {} (id={}) ===\n", state.config.name, experiment_id);
228        out.push_str(&format!("Confidence level: {:.0}%\n", state.config.confidence_level * 100.0));
229        out.push_str(&format!("Min samples required: {}\n\n", state.config.min_samples));
230
231        for variant in &state.config.variants {
232            if let Some(vr) = state.results.get(&variant.name) {
233                out.push_str(&format!(
234                    "Variant: {} | n={} | mean={:.4} | std_dev={:.4}\n",
235                    vr.variant_name, vr.sample_count, vr.mean, vr.std_dev
236                ));
237            } else {
238                out.push_str(&format!("Variant: {} | no data\n", variant.name));
239            }
240        }
241
242        if let Some(sig) = self.analyze_experiment(experiment_id) {
243            out.push_str(&format!(
244                "\nSignificance test: p={:.4}, significant={}, effect_size={:.4}\n",
245                sig.p_value, sig.is_significant, sig.effect_size
246            ));
247            if let Some(w) = &sig.winner {
248                out.push_str(&format!("Winner: {}\n", w));
249            } else {
250                out.push_str("Winner: (none yet)\n");
251            }
252        } else {
253            out.push_str("\nInsufficient data for significance test.\n");
254        }
255
256        Some(out)
257    }
258}
259
260impl Default for ExperimentRunner {
261    fn default() -> Self {
262        Self::new()
263    }
264}
265
266// --- helpers ---
267
268fn fnv1a(s: &str) -> u64 {
269    let mut hash: u64 = 14695981039346656037;
270    for byte in s.bytes() {
271        hash ^= byte as u64;
272        hash = hash.wrapping_mul(1099511628211);
273    }
274    hash
275}
276
277/// Approximate the normal CDF Φ(x) using the Abramowitz & Stegun rational approximation.
278fn normal_cdf(x: f64) -> f64 {
279    if x < 0.0 {
280        return 1.0 - normal_cdf(-x);
281    }
282    let t = 1.0 / (1.0 + 0.2316419 * x);
283    let poly = t * (0.319381530
284        + t * (-0.356563782
285        + t * (1.781477937
286        + t * (-1.821255978
287        + t * 1.330274429))));
288    1.0 - ((-x * x / 2.0).exp() / (2.0 * std::f64::consts::PI).sqrt()) * poly
289}
290
291fn current_time_ms() -> u64 {
292    use std::time::{SystemTime, UNIX_EPOCH};
293    SystemTime::now()
294        .duration_since(UNIX_EPOCH)
295        .map(|d| d.as_millis() as u64)
296        .unwrap_or(0)
297}
298
299#[cfg(test)]
300mod tests {
301    use super::*;
302
303    fn make_config(name: &str, weights: &[f64]) -> ExperimentConfig {
304        let variants = weights.iter().enumerate().map(|(i, &w)| ExperimentVariant {
305            name: format!("variant_{}", i),
306            prompt_template: format!("template {}", i),
307            weight: w,
308        }).collect();
309        ExperimentConfig {
310            name: name.to_string(),
311            variants,
312            min_samples: 5,
313            confidence_level: 0.95,
314        }
315    }
316
317    #[test]
318    fn test_variant_assignment_consistent() {
319        let mut runner = ExperimentRunner::new();
320        let config = make_config("test", &[0.5, 0.5]);
321        let eid = runner.create_experiment(config);
322        let v1 = runner.assign_variant(&eid, "user-abc").map(|v| v.name.clone());
323        let v2 = runner.assign_variant(&eid, "user-abc").map(|v| v.name.clone());
324        assert_eq!(v1, v2, "same user should always get same variant");
325    }
326
327    #[test]
328    fn test_record_metrics_grows_samples() {
329        let mut runner = ExperimentRunner::new();
330        let config = make_config("growth", &[0.5, 0.5]);
331        let eid = runner.create_experiment(config);
332        runner.record_metric(&eid, "variant_0", 1.0);
333        runner.record_metric(&eid, "variant_0", 2.0);
334        runner.record_metric(&eid, "variant_0", 3.0);
335        let state = runner.experiments.get(&eid).unwrap();
336        let vr = state.results.get("variant_0").unwrap();
337        assert_eq!(vr.sample_count, 3);
338        assert!((vr.mean - 2.0).abs() < 1e-9);
339    }
340
341    #[test]
342    fn test_welch_t_test_different_distributions() {
343        let a: Vec<f64> = (0..30).map(|i| i as f64 * 1.0).collect();
344        let b: Vec<f64> = (0..30).map(|i| i as f64 * 1.0 + 100.0).collect();
345        let p = ExperimentRunner::welch_t_test(&a, &b);
346        assert!(p < 0.05, "clearly different distributions should yield p < 0.05, got {}", p);
347    }
348
349    #[test]
350    fn test_cohen_d_direction() {
351        let a = vec![10.0, 10.0, 10.0, 10.0, 10.0];
352        let b = vec![5.0, 5.0, 5.0, 5.0, 5.0];
353        let d = ExperimentRunner::cohen_d(&a, &b);
354        // a > b so d should be positive... but if pooled_std is 0, d is 0
355        // Use non-constant samples to get a meaningful pooled_std
356        let a2 = vec![10.0, 11.0, 9.0, 10.5, 9.5];
357        let b2 = vec![5.0, 6.0, 4.0, 5.5, 4.5];
358        let d2 = ExperimentRunner::cohen_d(&a2, &b2);
359        assert!(d2 > 0.0, "a > b so cohen_d should be positive, got {}", d2);
360        let _ = d; // suppress unused warning
361    }
362
363    #[test]
364    fn test_experiment_report_non_empty() {
365        let mut runner = ExperimentRunner::new();
366        let config = make_config("report_test", &[0.5, 0.5]);
367        let eid = runner.create_experiment(config);
368        // Add enough samples for significance test
369        for i in 0..10 {
370            runner.record_metric(&eid, "variant_0", i as f64);
371            runner.record_metric(&eid, "variant_1", i as f64 + 50.0);
372        }
373        let report = runner.experiment_report(&eid);
374        assert!(report.is_some());
375        let r = report.unwrap();
376        assert!(!r.is_empty());
377        assert!(r.contains("Experiment"));
378    }
379}