tokio_prompt_orchestrator/
experiment_runner.rs1use std::collections::HashMap;
4
5#[derive(Debug, Clone)]
7pub struct ExperimentVariant {
8 pub name: String,
9 pub prompt_template: String,
10 pub weight: f64,
11}
12
13#[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#[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#[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#[derive(Debug)]
70pub struct ExperimentState {
71 pub config: ExperimentConfig,
72 pub results: HashMap<String, VariantResult>,
73 pub started_at_ms: u64,
74}
75
76pub struct ExperimentRunner {
78 pub experiments: HashMap<String, ExperimentState>,
79 next_id: u64,
80}
81
82impl ExperimentRunner {
83 pub fn new() -> Self {
85 Self {
86 experiments: HashMap::new(),
87 next_id: 1,
88 }
89 }
90
91 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 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 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 state.config.variants.last()
130 }
131
132 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 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 let p = 2.0 * (1.0 - normal_cdf(t.abs()));
158 p.clamp(0.0, 1.0)
159 }
160
161 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 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 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 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
266fn 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
277fn 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 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; }
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 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}