tokio_prompt_orchestrator/
ab_testing.rs1#![allow(dead_code)]
7
8use std::collections::HashMap;
9use std::collections::hash_map::DefaultHasher;
10use std::hash::{Hash, Hasher};
11use std::sync::atomic::{AtomicU64, Ordering};
12use std::sync::{Arc, RwLock};
13
14#[derive(Clone, Debug)]
20pub struct Variant {
21 pub id: String,
23 pub model: String,
25 pub prompt_modifier: String,
27 pub weight: f64,
29}
30
31#[derive(Debug)]
37pub struct Assignment {
38 pub variant_id: String,
40 pub session_id: String,
42 pub assigned_at: std::time::Instant,
44}
45
46pub struct VariantMetrics {
52 pub variant_id: String,
54 pub requests: AtomicU64,
56 pub total_latency_ms: AtomicU64,
58 pub total_tokens: AtomicU64,
60 pub successes: AtomicU64,
62}
63
64impl VariantMetrics {
65 pub fn new(variant_id: impl Into<String>) -> Self {
67 Self {
68 variant_id: variant_id.into(),
69 requests: AtomicU64::new(0),
70 total_latency_ms: AtomicU64::new(0),
71 total_tokens: AtomicU64::new(0),
72 successes: AtomicU64::new(0),
73 }
74 }
75
76 pub fn avg_latency(&self) -> f64 {
78 let reqs = self.requests.load(Ordering::Relaxed);
79 if reqs == 0 {
80 return 0.0;
81 }
82 self.total_latency_ms.load(Ordering::Relaxed) as f64 / reqs as f64
83 }
84
85 pub fn success_rate(&self) -> f64 {
87 let reqs = self.requests.load(Ordering::Relaxed);
88 if reqs == 0 {
89 return 0.0;
90 }
91 self.successes.load(Ordering::Relaxed) as f64 / reqs as f64
92 }
93
94 pub fn avg_tokens(&self) -> f64 {
96 let reqs = self.requests.load(Ordering::Relaxed);
97 if reqs == 0 {
98 return 0.0;
99 }
100 self.total_tokens.load(Ordering::Relaxed) as f64 / reqs as f64
101 }
102}
103
104pub struct AbTest {
110 pub id: String,
112 pub variants: Vec<Variant>,
114 pub assignments: RwLock<HashMap<String, String>>,
116 pub metrics: HashMap<String, Arc<VariantMetrics>>,
118 pub start_time: std::time::Instant,
120 pub traffic_pct: f64,
122}
123
124impl AbTest {
125 pub fn new(id: impl Into<String>, variants: Vec<Variant>, traffic_pct: f64) -> Self {
131 let mut metrics = HashMap::new();
132 for v in &variants {
133 metrics.insert(v.id.clone(), Arc::new(VariantMetrics::new(v.id.clone())));
134 }
135 Self {
136 id: id.into(),
137 variants,
138 assignments: RwLock::new(HashMap::new()),
139 metrics,
140 start_time: std::time::Instant::now(),
141 traffic_pct: traffic_pct.clamp(0.0, 1.0),
142 }
143 }
144
145 pub fn assign(&self, session_id: &str) -> Option<&Variant> {
150 if self.variants.is_empty() {
151 return None;
152 }
153
154 let hash = hash_str(session_id);
156 if (hash % 100) as f64 >= self.traffic_pct * 100.0 {
157 return None;
158 }
159
160 {
162 let assignments = self.assignments.read().ok()?;
163 if let Some(vid) = assignments.get(session_id) {
164 return self.variants.iter().find(|v| &v.id == vid);
165 }
166 }
167
168 let total_weight: f64 = self.variants.iter().map(|v| v.weight.max(0.0)).sum();
170 if total_weight == 0.0 {
171 return None;
172 }
173
174 let selection_hash = hash_str(&format!("{}{}", self.id, session_id));
177 let cursor = (selection_hash as f64 / u64::MAX as f64) * total_weight;
178
179 let mut cumulative = 0.0;
180 let mut chosen: Option<&Variant> = None;
181 for v in &self.variants {
182 cumulative += v.weight.max(0.0);
183 if cursor <= cumulative {
184 chosen = Some(v);
185 break;
186 }
187 }
188 if chosen.is_none() {
190 chosen = self.variants.last();
191 }
192
193 if let Some(v) = chosen {
194 if let Ok(mut map) = self.assignments.write() {
195 map.insert(session_id.to_string(), v.id.clone());
196 }
197 }
198
199 chosen
200 }
201
202 pub fn record_result(&self, session_id: &str, latency_ms: u64, tokens: u64, success: bool) {
204 let variant_id = {
205 match self.assignments.read() {
206 Ok(map) => map.get(session_id).cloned(),
207 Err(_) => return,
208 }
209 };
210 if let Some(vid) = variant_id {
211 if let Some(m) = self.metrics.get(&vid) {
212 m.requests.fetch_add(1, Ordering::Relaxed);
213 m.total_latency_ms.fetch_add(latency_ms, Ordering::Relaxed);
214 m.total_tokens.fetch_add(tokens, Ordering::Relaxed);
215 if success {
216 m.successes.fetch_add(1, Ordering::Relaxed);
217 }
218 }
219 }
220 }
221
222 pub fn winner(&self) -> Option<&str> {
225 if self.variants.len() < 2 {
226 return None;
227 }
228 let mut rates: Vec<(&str, f64)> = self
230 .variants
231 .iter()
232 .map(|v| {
233 let rate = self
234 .metrics
235 .get(&v.id)
236 .map(|m| m.success_rate())
237 .unwrap_or(0.0);
238 (v.id.as_str(), rate)
239 })
240 .collect();
241
242 rates.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
243
244 let best = rates[0];
245 let second = rates[1];
246 if best.1 - second.1 > 0.05 {
247 Some(best.0)
248 } else {
249 None
250 }
251 }
252
253 pub fn summary(&self) -> String {
255 let mut out = format!("## A/B Test: {}\n\n", self.id);
256 out.push_str("| Variant | Requests | Avg Latency (ms) | Avg Tokens | Success Rate |\n");
257 out.push_str("|---------|----------|-----------------|------------|-------------|\n");
258 for v in &self.variants {
259 if let Some(m) = self.metrics.get(&v.id) {
260 out.push_str(&format!(
261 "| {} | {} | {:.1} | {:.1} | {:.1}% |\n",
262 v.id,
263 m.requests.load(Ordering::Relaxed),
264 m.avg_latency(),
265 m.avg_tokens(),
266 m.success_rate() * 100.0,
267 ));
268 }
269 }
270 if let Some(w) = self.winner() {
271 out.push_str(&format!("\n**Winner: {}**\n", w));
272 } else {
273 out.push_str("\n*No winner determined yet.*\n");
274 }
275 out
276 }
277}
278
279pub struct AbTestRegistry {
285 tests: RwLock<HashMap<String, AbTest>>,
286}
287
288impl AbTestRegistry {
289 pub fn new() -> Self {
291 Self {
292 tests: RwLock::new(HashMap::new()),
293 }
294 }
295
296 pub fn create_test(&self, test: AbTest) {
298 if let Ok(mut map) = self.tests.write() {
299 map.insert(test.id.clone(), test);
300 }
301 }
302
303 pub fn get_test_summary(&self, test_id: &str) -> Option<String> {
308 let map = self.tests.read().ok()?;
309 map.get(test_id).map(|t| t.summary())
310 }
311
312 pub fn conclude_test(&self, test_id: &str) -> Option<String> {
314 let summary = self.get_test_summary(test_id);
315 if let Ok(mut map) = self.tests.write() {
316 map.remove(test_id);
317 }
318 summary
319 }
320
321 pub fn active_tests(&self) -> Vec<String> {
323 self.tests
324 .read()
325 .map(|m| m.keys().cloned().collect())
326 .unwrap_or_default()
327 }
328}
329
330impl Default for AbTestRegistry {
331 fn default() -> Self {
332 Self::new()
333 }
334}
335
336fn hash_str(s: &str) -> u64 {
341 let mut h = DefaultHasher::new();
342 s.hash(&mut h);
343 h.finish()
344}
345
346#[cfg(test)]
351mod tests {
352 use super::*;
353
354 fn make_test(traffic: f64) -> AbTest {
355 AbTest::new(
356 "test-1",
357 vec![
358 Variant { id: "A".into(), model: "gpt-4".into(), prompt_modifier: "".into(), weight: 1.0 },
359 Variant { id: "B".into(), model: "gpt-3.5".into(), prompt_modifier: "[fast] ".into(), weight: 1.0 },
360 ],
361 traffic,
362 )
363 }
364
365 #[test]
366 fn assignment_is_deterministic() {
367 let t = make_test(1.0);
368 let first = t.assign("session-abc").map(|v| v.id.clone());
369 let second = t.assign("session-abc").map(|v| v.id.clone());
370 assert_eq!(first, second);
371 }
372
373 #[test]
374 fn no_traffic_returns_none() {
375 let t = make_test(0.0);
376 let mut excluded = 0usize;
378 for i in 0..100 {
379 if t.assign(&format!("s-{i}")).is_none() {
380 excluded += 1;
381 }
382 }
383 assert!(excluded > 90, "expected most sessions excluded, got {} included", 100 - excluded);
384 }
385
386 #[test]
387 fn winner_requires_5pct_margin() {
388 let t = make_test(1.0);
389 for i in 0..50 {
391 let s = format!("s-{i}");
392 t.assign(&s);
393 t.record_result(&s, 100, 50, true);
394 }
395 let _ = t.winner();
398 }
399
400 #[test]
401 fn summary_contains_header() {
402 let t = make_test(1.0);
403 let s = t.summary();
404 assert!(s.contains("Variant"), "summary missing header");
405 }
406
407 #[test]
408 fn registry_lifecycle() {
409 let reg = AbTestRegistry::new();
410 reg.create_test(make_test(1.0));
411 assert!(reg.active_tests().contains(&"test-1".to_string()));
412 let summary = reg.conclude_test("test-1");
413 assert!(summary.is_some());
414 assert!(reg.active_tests().is_empty());
415 }
416}