Skip to main content

tokio_prompt_orchestrator/
ab_test.rs

1//! # A/B Testing Framework
2//!
3//! Statistically rigorous A/B testing for LLM model and prompt experiments.
4//! Supports weighted traffic splits, deterministic user assignment via FNV hash,
5//! two-proportion z-test significance testing, and automatic experiment completion.
6
7#![allow(dead_code)]
8
9use std::collections::HashMap;
10use std::sync::Arc;
11
12use dashmap::DashMap;
13use tracing::{debug, info, warn};
14
15use crate::templates::PromptTemplate;
16
17// ---------------------------------------------------------------------------
18// Re-export legacy types so existing lib.rs re-exports keep compiling
19// ---------------------------------------------------------------------------
20
21/// Which A/B variant a request is assigned to (legacy two-variant API).
22#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
23pub enum Variant {
24    /// The control variant.
25    A,
26    /// The treatment variant.
27    B,
28}
29
30/// The metric used to evaluate variant quality (legacy API).
31#[derive(Clone)]
32pub enum SuccessMetric {
33    /// Higher output length is better.
34    OutputLength,
35    /// Lower latency is better.
36    Latency,
37    /// A caller-supplied scoring function.
38    CustomFn(Arc<dyn Fn(&str) -> f64 + Send + Sync>),
39    /// Explicit numeric rating supplied by the caller.
40    UserRating,
41}
42
43impl std::fmt::Debug for SuccessMetric {
44    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
45        match self {
46            Self::OutputLength => write!(f, "OutputLength"),
47            Self::Latency => write!(f, "Latency"),
48            Self::CustomFn(_) => write!(f, "CustomFn(...)"),
49            Self::UserRating => write!(f, "UserRating"),
50        }
51    }
52}
53
54/// Configuration for a single legacy two-variant A/B experiment.
55#[derive(Debug, Clone)]
56pub struct AbTestConfig {
57    /// Unique experiment name.
58    pub name: String,
59    /// The control prompt template.
60    pub variant_a: PromptTemplate,
61    /// The treatment prompt template.
62    pub variant_b: PromptTemplate,
63    /// Fraction of traffic routed to variant A (0.0–1.0).
64    pub traffic_split: f64,
65    /// Metric used to score each response.
66    pub success_metric: SuccessMetric,
67    /// Minimum observations per variant before analysis.
68    pub min_samples: usize,
69}
70
71/// Statistical result of a legacy two-variant A/B experiment.
72#[derive(Debug, Clone)]
73pub struct AbTestResult {
74    /// The winning variant, or `None` when not yet significant.
75    pub winner: Option<Variant>,
76    /// Confidence level derived from the p-value.
77    pub confidence: f64,
78    /// Two-tailed p-value.
79    pub p_value: f64,
80    /// Effect size (Cohen's d).
81    pub effect_size: f64,
82    /// Number of observations for variant A.
83    pub samples_a: usize,
84    /// Number of observations for variant B.
85    pub samples_b: usize,
86    /// Sample mean score for variant A.
87    pub mean_a: f64,
88    /// Sample mean score for variant B.
89    pub mean_b: f64,
90}
91
92/// Per-variant running statistics.
93#[derive(Debug, Default)]
94struct LegacyVariantStats {
95    scores: Vec<f64>,
96}
97
98impl LegacyVariantStats {
99    fn push(&mut self, score: f64) { self.scores.push(score); }
100    fn n(&self) -> usize { self.scores.len() }
101    fn mean(&self) -> f64 {
102        if self.scores.is_empty() { return 0.0; }
103        self.scores.iter().sum::<f64>() / self.scores.len() as f64
104    }
105    fn variance(&self) -> f64 {
106        let n = self.scores.len();
107        if n < 2 { return 0.0; }
108        let mean = self.mean();
109        let ss: f64 = self.scores.iter().map(|x| (x - mean).powi(2)).sum();
110        ss / (n - 1) as f64
111    }
112}
113
114struct LegacyExperiment {
115    config: AbTestConfig,
116    stats_a: LegacyVariantStats,
117    stats_b: LegacyVariantStats,
118}
119
120/// Thread-safe legacy two-variant A/B test runner.
121#[derive(Clone, Default)]
122pub struct AbTestRunner {
123    experiments: Arc<DashMap<String, parking_lot::Mutex<LegacyExperiment>>>,
124}
125
126impl AbTestRunner {
127    /// Create an empty runner.
128    pub fn new() -> Self { Self::default() }
129
130    /// Register a new experiment.
131    pub fn register(&self, config: AbTestConfig) {
132        info!(name = %config.name, "ab_test: registering experiment");
133        let name = config.name.clone();
134        self.experiments.insert(name, parking_lot::Mutex::new(LegacyExperiment {
135            config,
136            stats_a: LegacyVariantStats::default(),
137            stats_b: LegacyVariantStats::default(),
138        }));
139    }
140
141    /// Remove an experiment.
142    pub fn delete(&self, name: &str) -> bool {
143        let removed = self.experiments.remove(name).is_some();
144        if removed { info!(name, "ab_test: experiment deleted"); }
145        else { warn!(name, "ab_test: delete requested for unknown experiment"); }
146        removed
147    }
148
149    /// Assign a user deterministically to a variant.
150    pub fn assign(&self, name: &str, user_id: &str) -> Option<Variant> {
151        let exp = self.experiments.get(name)?;
152        let guard = exp.lock();
153        let split = guard.config.traffic_split.clamp(0.0, 1.0);
154        let hash = fnv1a_64_pair(name.as_bytes(), user_id.as_bytes());
155        let bucket = (hash as f64) / (u64::MAX as f64);
156        let variant = if bucket < split { Variant::A } else { Variant::B };
157        debug!(experiment = name, user = user_id, bucket, split, ?variant, "ab_test: assignment");
158        Some(variant)
159    }
160
161    /// Record a metric observation for the given variant.
162    pub fn record_observation(&self, name: &str, variant: Variant, score: f64) {
163        if let Some(exp) = self.experiments.get(name) {
164            let mut guard = exp.lock();
165            match variant {
166                Variant::A => guard.stats_a.push(score),
167                Variant::B => guard.stats_b.push(score),
168            }
169        }
170    }
171
172    /// Run Welch's t-test on the accumulated observations.
173    pub fn analyse(&self, name: &str) -> Option<AbTestResult> {
174        let exp = self.experiments.get(name)?;
175        let guard = exp.lock();
176        let min = guard.config.min_samples;
177        if guard.stats_a.n() < min || guard.stats_b.n() < min { return None; }
178        let result = welch_t_test(&guard.stats_a, &guard.stats_b);
179        info!(experiment = name, ?result.winner, p_value = result.p_value, "ab_test: analysis complete");
180        Some(result)
181    }
182
183    /// Return sample counts for both variants.
184    pub fn sample_counts(&self, name: &str) -> Option<(usize, usize)> {
185        let exp = self.experiments.get(name)?;
186        let guard = exp.lock();
187        Some((guard.stats_a.n(), guard.stats_b.n()))
188    }
189
190    /// Return names of all registered experiments.
191    pub fn experiment_names(&self) -> Vec<String> {
192        self.experiments.iter().map(|e| e.key().clone()).collect()
193    }
194
195    /// Return current means for both variants.
196    pub fn current_means(&self, name: &str) -> Option<(f64, f64)> {
197        let exp = self.experiments.get(name)?;
198        let guard = exp.lock();
199        Some((guard.stats_a.mean(), guard.stats_b.mean()))
200    }
201}
202
203fn welch_t_test(a: &LegacyVariantStats, b: &LegacyVariantStats) -> AbTestResult {
204    let na = a.n() as f64;
205    let nb = b.n() as f64;
206    let mean_a = a.mean();
207    let mean_b = b.mean();
208    let var_a = a.variance();
209    let var_b = b.variance();
210    let se_sq_a = var_a / na;
211    let se_sq_b = var_b / nb;
212    let se = (se_sq_a + se_sq_b).sqrt();
213    let (p_value, winner) = if se < f64::EPSILON {
214        (1.0_f64, None)
215    } else {
216        let t = (mean_a - mean_b) / se;
217        let df_num = (se_sq_a + se_sq_b).powi(2);
218        let df_den = se_sq_a.powi(2) / (na - 1.0) + se_sq_b.powi(2) / (nb - 1.0);
219        let df = if df_den < f64::EPSILON { 1.0 } else { df_num / df_den };
220        let p = two_tailed_p_legacy(t, df);
221        let w = if p <= 0.05 {
222            if mean_a > mean_b { Some(Variant::A) } else { Some(Variant::B) }
223        } else { None };
224        (p, w)
225    };
226    let pooled_var = ((na - 1.0) * var_a + (nb - 1.0) * var_b) / (na + nb - 2.0);
227    let pooled_sd = pooled_var.sqrt();
228    let effect_size = if pooled_sd < f64::EPSILON { 0.0 } else { (mean_a - mean_b) / pooled_sd };
229    let confidence = (1.0 - p_value).clamp(0.0, 1.0);
230    AbTestResult { winner, confidence, p_value, effect_size, samples_a: a.n(), samples_b: b.n(), mean_a, mean_b }
231}
232
233fn two_tailed_p_legacy(t: f64, df: f64) -> f64 {
234    let t2 = t * t;
235    let x = df / (df + t2);
236    let p_one_tail = regularised_incomplete_beta_legacy(x, df / 2.0, 0.5) / 2.0;
237    (2.0 * p_one_tail).clamp(0.0, 1.0)
238}
239
240fn regularised_incomplete_beta_legacy(x: f64, a: f64, b: f64) -> f64 {
241    if x <= 0.0 { return 0.0; }
242    if x >= 1.0 { return 1.0; }
243    if x > (a + 1.0) / (a + b + 2.0) {
244        return 1.0 - regularised_incomplete_beta_legacy(1.0 - x, b, a);
245    }
246    let ln_beta_val = ln_gamma_legacy(a) + ln_gamma_legacy(b) - ln_gamma_legacy(a + b);
247    let front = (x.ln() * a + (1.0 - x).ln() * b - ln_beta_val).exp() / a;
248    let mut c = 1.0_f64;
249    let mut d = 1.0 - (a + b) * x / (a + 1.0);
250    d = if d.abs() < 1e-30 { 1e-30 } else { 1.0 / d };
251    let mut f = d;
252    for m in 1_u32..=200 {
253        let mf = m as f64;
254        let num_even = mf * (b - mf) * x / ((a + 2.0 * mf - 1.0) * (a + 2.0 * mf));
255        d = 1.0 + num_even * d;
256        d = if d.abs() < 1e-30 { 1e-30 } else { 1.0 / d };
257        c = 1.0 + num_even / c;
258        c = if c.abs() < 1e-30 { 1e-30 } else { c };
259        f *= c * d;
260        let num_odd = -(a + mf) * (a + b + mf) * x / ((a + 2.0 * mf) * (a + 2.0 * mf + 1.0));
261        d = 1.0 + num_odd * d;
262        d = if d.abs() < 1e-30 { 1e-30 } else { 1.0 / d };
263        c = 1.0 + num_odd / c;
264        c = if c.abs() < 1e-30 { 1e-30 } else { c };
265        let delta = c * d;
266        f *= delta;
267        if (delta - 1.0).abs() < 1e-10 { break; }
268    }
269    front * f
270}
271
272fn ln_gamma_legacy(x: f64) -> f64 {
273    const G: f64 = 7.0;
274    const C: [f64; 9] = [
275        0.999_999_999_999_809_9, 676.520_368_121_885_1, -1_259.139_216_722_402_8,
276        771.323_428_777_653_1, -176.615_029_162_140_6, 12.507_343_278_686_905,
277        -0.138_571_095_265_720_1, 9.984_369_578_019_572e-6, 1.505_632_735_149_312_4e-7,
278    ];
279    if x < 0.5 {
280        return std::f64::consts::PI.ln() - (std::f64::consts::PI * x).sin().ln() - ln_gamma_legacy(1.0 - x);
281    }
282    let z = x - 1.0;
283    let mut sum = C[0];
284    for (i, &ci) in C.iter().enumerate().skip(1) { sum += ci / (z + i as f64); }
285    let t = z + G + 0.5;
286    0.5 * std::f64::consts::TAU.ln() + (z + 0.5) * t.ln() - t + sum.ln()
287}
288
289// ---------------------------------------------------------------------------
290// NEW SPEC IMPLEMENTATION
291// ---------------------------------------------------------------------------
292
293/// A single variant in an experiment.
294#[derive(Debug, Clone)]
295pub struct ExperimentVariantSpec {
296    /// Unique variant identifier.
297    pub id: String,
298    /// Human-readable name.
299    pub name: String,
300    /// Relative weight for traffic split (e.g. 1.0 = equal share).
301    pub weight: f64,
302    /// Optional model override for this variant.
303    pub model_id: Option<String>,
304    /// Optional prompt template override.
305    pub prompt_template: Option<String>,
306    /// Arbitrary config key-value pairs.
307    pub config: HashMap<String, String>,
308}
309
310/// Status of an experiment lifecycle.
311#[derive(Debug, Clone, PartialEq, Eq)]
312pub enum ExperimentStatus {
313    /// Not yet started.
314    Draft,
315    /// Currently collecting samples.
316    Running,
317    /// Temporarily paused.
318    Paused,
319    /// Finished normally.
320    Completed,
321    /// Stopped early.
322    Aborted,
323}
324
325/// Condition that triggers automatic experiment completion.
326#[derive(Debug, Clone)]
327pub enum EndCondition {
328    /// Stop after a fixed number of total samples.
329    FixedSamples(u64),
330    /// Stop after a fixed duration in seconds.
331    FixedDuration(u64),
332    /// Never auto-stop — must be stopped manually.
333    ManualStop,
334    /// Stop when statistical significance is reached.
335    StatisticalSignificance {
336        /// Minimum samples per variant before checking.
337        min_samples: u64,
338        /// Required confidence level (e.g. 0.95).
339        confidence: f64,
340    },
341}
342
343/// A multi-variant experiment definition.
344#[derive(Debug, Clone)]
345pub struct Experiment {
346    /// Unique experiment identifier.
347    pub id: String,
348    /// Human-readable name.
349    pub name: String,
350    /// The variants to test.
351    pub variants: Vec<ExperimentVariantSpec>,
352    /// Current lifecycle status.
353    pub status: ExperimentStatus,
354    /// Unix timestamp when the experiment was created.
355    pub created_at: u64,
356    /// Condition that triggers automatic completion.
357    pub end_condition: EndCondition,
358}
359
360/// Records that a user was assigned to a specific variant.
361#[derive(Debug, Clone)]
362pub struct Assignment {
363    /// Experiment this assignment belongs to.
364    pub experiment_id: String,
365    /// Variant the user was assigned to.
366    pub variant_id: String,
367    /// User identifier.
368    pub user_id: String,
369    /// Unix timestamp of assignment.
370    pub assigned_at: u64,
371}
372
373/// Aggregated results for one variant.
374#[derive(Debug, Clone)]
375pub struct ExperimentResult {
376    /// Variant this result belongs to.
377    pub variant_id: String,
378    /// Total number of observations.
379    pub samples: u64,
380    /// Number of successful outcomes.
381    pub successes: u64,
382    /// Total cost accumulated.
383    pub total_cost: f64,
384    /// Average latency in milliseconds.
385    pub avg_latency_ms: f64,
386    /// Fraction of samples that were successful.
387    pub conversion_rate: f64,
388}
389
390/// Per-variant mutable state stored inside `AbTestManager`.
391#[derive(Debug, Default, Clone)]
392struct VariantState {
393    samples: u64,
394    successes: u64,
395    total_cost: f64,
396    total_latency_ms: u64,
397}
398
399struct ExperimentState {
400    experiment: Experiment,
401    variants: HashMap<String, VariantState>,
402    /// Total samples across all variants (for FixedSamples end condition).
403    total_samples: u64,
404}
405
406/// Manages multi-variant A/B experiments.
407pub struct AbTestManager {
408    experiments: HashMap<String, ExperimentState>,
409}
410
411impl AbTestManager {
412    /// Create an empty manager.
413    pub fn new() -> Self {
414        Self { experiments: HashMap::new() }
415    }
416
417    /// Register a new experiment.
418    pub fn create_experiment(&mut self, experiment: Experiment) {
419        let mut variants: HashMap<String, VariantState> = HashMap::new();
420        for v in &experiment.variants {
421            variants.insert(v.id.clone(), VariantState::default());
422        }
423        self.experiments.insert(experiment.id.clone(), ExperimentState {
424            experiment,
425            variants,
426            total_samples: 0,
427        });
428    }
429
430    /// Deterministically assign a user to a variant using weighted FNV hash.
431    ///
432    /// Returns `None` if the experiment doesn't exist or is not Running.
433    pub fn assign(&self, experiment_id: &str, user_id: &str) -> Option<Assignment> {
434        let state = self.experiments.get(experiment_id)?;
435        if state.experiment.status != ExperimentStatus::Running {
436            return None;
437        }
438        let variants = &state.experiment.variants;
439        if variants.is_empty() {
440            return None;
441        }
442        let total_weight: f64 = variants.iter().map(|v| v.weight.max(0.0)).sum();
443        if total_weight <= 0.0 {
444            return None;
445        }
446
447        // FNV-1a hash of user_id + experiment_id
448        let hash = fnv1a_64_pair(user_id.as_bytes(), experiment_id.as_bytes());
449        // Map to [0.0, total_weight)
450        let bucket = (hash as f64 / u64::MAX as f64) * total_weight;
451
452        let mut cumulative = 0.0;
453        let mut selected_variant_id = &variants[0].id;
454        for v in variants {
455            cumulative += v.weight.max(0.0);
456            if bucket < cumulative {
457                selected_variant_id = &v.id;
458                break;
459            }
460        }
461
462        Some(Assignment {
463            experiment_id: experiment_id.to_string(),
464            variant_id: selected_variant_id.clone(),
465            user_id: user_id.to_string(),
466            assigned_at: 0, // caller fills in timestamp
467        })
468    }
469
470    /// Record an outcome for a previously made assignment.
471    pub fn record_outcome(&mut self, assignment: &Assignment, success: bool, cost: f64, latency_ms: u64) {
472        if let Some(state) = self.experiments.get_mut(&assignment.experiment_id) {
473            if let Some(vs) = state.variants.get_mut(&assignment.variant_id) {
474                vs.samples += 1;
475                if success { vs.successes += 1; }
476                vs.total_cost += cost;
477                vs.total_latency_ms += latency_ms;
478                state.total_samples += 1;
479            }
480        }
481    }
482
483    /// Return aggregated results for all variants in an experiment.
484    pub fn results(&self, experiment_id: &str) -> Vec<ExperimentResult> {
485        let Some(state) = self.experiments.get(experiment_id) else { return vec![]; };
486        state.experiment.variants.iter().map(|v| {
487            let vs = state.variants.get(&v.id).cloned().unwrap_or_default();
488            let conversion_rate = if vs.samples > 0 { vs.successes as f64 / vs.samples as f64 } else { 0.0 };
489            let avg_latency_ms = if vs.samples > 0 { vs.total_latency_ms as f64 / vs.samples as f64 } else { 0.0 };
490            ExperimentResult {
491                variant_id: v.id.clone(),
492                samples: vs.samples,
493                successes: vs.successes,
494                total_cost: vs.total_cost,
495                avg_latency_ms,
496                conversion_rate,
497            }
498        }).collect()
499    }
500
501    /// Two-proportion z-test significance.
502    ///
503    /// Returns p-value approximation using erfc: p = erfc(|z| / sqrt(2)).
504    pub fn statistical_significance(r_a: &ExperimentResult, r_b: &ExperimentResult) -> f64 {
505        let n_a = r_a.samples as f64;
506        let n_b = r_b.samples as f64;
507        if n_a < 1.0 || n_b < 1.0 { return 1.0; }
508        let s_a = r_a.successes as f64;
509        let s_b = r_b.successes as f64;
510        let p_a = s_a / n_a;
511        let p_b = s_b / n_b;
512        let p_pool = (s_a + s_b) / (n_a + n_b);
513        let denom = (p_pool * (1.0 - p_pool) * (1.0 / n_a + 1.0 / n_b)).sqrt();
514        if denom < f64::EPSILON { return 1.0; }
515        let z = (p_a - p_b) / denom;
516        erfc_approx(z.abs() / std::f64::consts::SQRT_2)
517    }
518
519    /// Return the variant ID with the highest conversion rate if statistically
520    /// significant (z > 1.96 for p < 0.05), or None.
521    pub fn winner(&self, experiment_id: &str) -> Option<String> {
522        let results = self.results(experiment_id);
523        if results.len() < 2 { return None; }
524
525        // Find variant with highest conversion rate.
526        let best = results.iter()
527            .max_by(|a, b| a.conversion_rate.partial_cmp(&b.conversion_rate).unwrap_or(std::cmp::Ordering::Equal))?;
528
529        // Check significance vs second-best.
530        let second = results.iter()
531            .filter(|r| r.variant_id != best.variant_id)
532            .max_by(|a, b| a.conversion_rate.partial_cmp(&b.conversion_rate).unwrap_or(std::cmp::Ordering::Equal))?;
533
534        let p = Self::statistical_significance(best, second);
535        if p < 0.05 {
536            Some(best.variant_id.clone())
537        } else {
538            None
539        }
540    }
541
542    /// Check end condition and mark experiment completed if met. Returns true if completed.
543    pub fn auto_complete(&mut self, experiment_id: &str, now: u64) -> bool {
544        let state = match self.experiments.get_mut(experiment_id) {
545            Some(s) => s,
546            None => return false,
547        };
548        if state.experiment.status != ExperimentStatus::Running {
549            return false;
550        }
551        let should_complete = match &state.experiment.end_condition {
552            EndCondition::FixedSamples(n) => state.total_samples >= *n,
553            EndCondition::FixedDuration(secs) => now >= state.experiment.created_at + secs,
554            EndCondition::ManualStop => false,
555            EndCondition::StatisticalSignificance { min_samples, confidence } => {
556                let min = *min_samples;
557                let required_confidence = *confidence;
558                // Check if all variants have min samples and any pair is significant.
559                let all_have_min = state.variants.values().all(|v| v.samples >= min);
560                if !all_have_min { return false; }
561                // Build results inline to avoid borrow issues.
562                let results: Vec<ExperimentResult> = state.experiment.variants.iter().map(|v| {
563                    let vs = state.variants.get(&v.id).cloned().unwrap_or_default();
564                    let conversion_rate = if vs.samples > 0 { vs.successes as f64 / vs.samples as f64 } else { 0.0 };
565                    let avg_latency_ms = if vs.samples > 0 { vs.total_latency_ms as f64 / vs.samples as f64 } else { 0.0 };
566                    ExperimentResult {
567                        variant_id: v.id.clone(),
568                        samples: vs.samples,
569                        successes: vs.successes,
570                        total_cost: vs.total_cost,
571                        avg_latency_ms,
572                        conversion_rate,
573                    }
574                }).collect();
575                // Check any pair for significance.
576                let mut found = false;
577                'outer: for i in 0..results.len() {
578                    for j in (i + 1)..results.len() {
579                        let p = Self::statistical_significance(&results[i], &results[j]);
580                        if 1.0 - p >= required_confidence {
581                            found = true;
582                            break 'outer;
583                        }
584                    }
585                }
586                found
587            }
588        };
589        if should_complete {
590            state.experiment.status = ExperimentStatus::Completed;
591            info!(experiment_id, "ab_test: experiment auto-completed");
592        }
593        should_complete
594    }
595}
596
597impl Default for AbTestManager {
598    fn default() -> Self { Self::new() }
599}
600
601// ---------------------------------------------------------------------------
602// Math helpers
603// ---------------------------------------------------------------------------
604
605/// erfc approximation using Horner's method (Abramowitz & Stegun 7.1.26).
606/// Accurate to ~1.5e-7.
607fn erfc_approx(x: f64) -> f64 {
608    if x < 0.0 { return 2.0 - erfc_approx(-x); }
609    let t = 1.0 / (1.0 + 0.3275911 * x);
610    let poly = t * (0.254_829_592
611        + t * (-0.284_496_736
612        + t * (1.421_413_741
613        + t * (-1.453_152_027
614        + t * 1.061_405_429))));
615    poly * (-x * x).exp()
616}
617
618/// FNV-1a 64-bit hash of two byte slices separated by a null byte.
619fn fnv1a_64_pair(a: &[u8], b: &[u8]) -> u64 {
620    const OFFSET_BASIS: u64 = 14_695_981_039_346_656_037;
621    const PRIME: u64 = 1_099_511_628_211;
622    let mut hash = OFFSET_BASIS;
623    for &byte in a.iter().chain(&[0u8]).chain(b.iter()) {
624        hash ^= u64::from(byte);
625        hash = hash.wrapping_mul(PRIME);
626    }
627    hash
628}
629
630// ---------------------------------------------------------------------------
631// Tests
632// ---------------------------------------------------------------------------
633
634#[cfg(test)]
635#[allow(clippy::unwrap_used, clippy::expect_used)]
636mod tests {
637    use super::*;
638
639    fn make_running_experiment(variants_weights: &[(&str, f64)]) -> Experiment {
640        let variants = variants_weights.iter().map(|(id, w)| ExperimentVariantSpec {
641            id: id.to_string(),
642            name: id.to_string(),
643            weight: *w,
644            model_id: None,
645            prompt_template: None,
646            config: HashMap::new(),
647        }).collect();
648        Experiment {
649            id: "exp-1".to_string(),
650            name: "Test Experiment".to_string(),
651            variants,
652            status: ExperimentStatus::Running,
653            created_at: 1000,
654            end_condition: EndCondition::ManualStop,
655        }
656    }
657
658    #[test]
659    fn deterministic_assignment() {
660        let mut manager = AbTestManager::new();
661        let exp = make_running_experiment(&[("a", 1.0), ("b", 1.0)]);
662        manager.create_experiment(exp);
663
664        let a1 = manager.assign("exp-1", "user-42").unwrap();
665        let a2 = manager.assign("exp-1", "user-42").unwrap();
666        assert_eq!(a1.variant_id, a2.variant_id, "assignment must be deterministic");
667    }
668
669    #[test]
670    fn weighted_split_proportional() {
671        let mut manager = AbTestManager::new();
672        // 3:1 weight ratio — expect ~75% assigned to "heavy", ~25% to "light".
673        let exp = make_running_experiment(&[("heavy", 3.0), ("light", 1.0)]);
674        manager.create_experiment(exp);
675
676        let mut counts: HashMap<String, u64> = HashMap::new();
677        for i in 0..1000 {
678            let a = manager.assign("exp-1", &format!("user-{i}")).unwrap();
679            *counts.entry(a.variant_id).or_default() += 1;
680        }
681        let heavy = *counts.get("heavy").unwrap_or(&0) as f64;
682        // Allow ±10% tolerance around 750.
683        assert!(heavy > 650.0 && heavy < 850.0,
684            "expected ~750 heavy assignments, got {heavy}");
685    }
686
687    #[test]
688    fn z_test_significance_large_difference() {
689        let r_a = ExperimentResult {
690            variant_id: "a".into(), samples: 500, successes: 400,
691            total_cost: 0.0, avg_latency_ms: 0.0, conversion_rate: 0.8,
692        };
693        let r_b = ExperimentResult {
694            variant_id: "b".into(), samples: 500, successes: 200,
695            total_cost: 0.0, avg_latency_ms: 0.0, conversion_rate: 0.4,
696        };
697        let p = AbTestManager::statistical_significance(&r_a, &r_b);
698        assert!(p < 0.001, "large difference should be highly significant, got p={p}");
699    }
700
701    #[test]
702    fn z_test_no_significance_equal_rates() {
703        let r_a = ExperimentResult {
704            variant_id: "a".into(), samples: 100, successes: 50,
705            total_cost: 0.0, avg_latency_ms: 0.0, conversion_rate: 0.5,
706        };
707        let r_b = ExperimentResult {
708            variant_id: "b".into(), samples: 100, successes: 50,
709            total_cost: 0.0, avg_latency_ms: 0.0, conversion_rate: 0.5,
710        };
711        let p = AbTestManager::statistical_significance(&r_a, &r_b);
712        assert!(p > 0.9, "equal rates should not be significant, got p={p}");
713    }
714
715    #[test]
716    fn winner_detection() {
717        let mut manager = AbTestManager::new();
718        let exp = make_running_experiment(&[("a", 1.0), ("b", 1.0)]);
719        manager.create_experiment(exp);
720
721        // Record 500 outcomes per variant: a=80% success, b=40% success.
722        let assign_a = Assignment { experiment_id: "exp-1".into(), variant_id: "a".into(), user_id: "u".into(), assigned_at: 0 };
723        let assign_b = Assignment { experiment_id: "exp-1".into(), variant_id: "b".into(), user_id: "u".into(), assigned_at: 0 };
724        for i in 0..500u64 {
725            manager.record_outcome(&assign_a, i < 400, 0.01, 50);
726            manager.record_outcome(&assign_b, i < 200, 0.01, 50);
727        }
728        let winner = manager.winner("exp-1");
729        assert_eq!(winner.as_deref(), Some("a"), "variant a should win with 80% vs 40%");
730    }
731
732    #[test]
733    fn auto_complete_on_fixed_samples() {
734        let mut manager = AbTestManager::new();
735        let mut exp = make_running_experiment(&[("a", 1.0), ("b", 1.0)]);
736        exp.end_condition = EndCondition::FixedSamples(10);
737        manager.create_experiment(exp);
738
739        let assign_a = Assignment { experiment_id: "exp-1".into(), variant_id: "a".into(), user_id: "u".into(), assigned_at: 0 };
740        for i in 0..9u64 {
741            manager.record_outcome(&assign_a, i % 2 == 0, 0.01, 50);
742        }
743        // 9 samples — not yet complete.
744        assert!(!manager.auto_complete("exp-1", 2000));
745
746        manager.record_outcome(&assign_a, true, 0.01, 50);
747        // 10 samples — should complete.
748        assert!(manager.auto_complete("exp-1", 2000));
749        assert_eq!(manager.experiments["exp-1"].experiment.status, ExperimentStatus::Completed);
750    }
751
752    // --- Legacy AbTestRunner tests kept for regression ---
753
754    fn make_runner_with_experiment(split: f64, min_samples: usize) -> AbTestRunner {
755        let runner = AbTestRunner::new();
756        runner.register(AbTestConfig {
757            name: "test-exp".into(),
758            variant_a: PromptTemplate::builder("a").body("A").build(),
759            variant_b: PromptTemplate::builder("b").body("B").build(),
760            traffic_split: split,
761            success_metric: SuccessMetric::OutputLength,
762            min_samples,
763        });
764        runner
765    }
766
767    #[test]
768    fn legacy_assignment_is_deterministic() {
769        let runner = make_runner_with_experiment(0.5, 1);
770        let v1 = runner.assign("test-exp", "user-123").unwrap();
771        let v2 = runner.assign("test-exp", "user-123").unwrap();
772        assert_eq!(v1, v2);
773    }
774
775    #[test]
776    fn legacy_split_zero_always_gives_b() {
777        let runner = make_runner_with_experiment(0.0, 1);
778        for i in 0..20 {
779            let v = runner.assign("test-exp", &format!("u{i}")).unwrap();
780            assert_eq!(v, Variant::B);
781        }
782    }
783
784    #[test]
785    fn legacy_analyse_detects_significant_difference() {
786        let runner = make_runner_with_experiment(0.5, 20);
787        for i in 0..50 {
788            let jitter = (i as f64) * 0.001;
789            runner.record_observation("test-exp", Variant::A, 10.0 + jitter);
790            runner.record_observation("test-exp", Variant::B, 1.0 + jitter);
791        }
792        let result = runner.analyse("test-exp").unwrap();
793        assert_eq!(result.winner, Some(Variant::A));
794        assert!(result.p_value < 0.05);
795    }
796}