Skip to main content

tokio_prompt_orchestrator/
ab_testing.rs

1//! # A/B Testing Framework
2//!
3//! A/B testing for model/prompt comparisons with deterministic assignment,
4//! atomic metrics collection, and statistical winner determination.
5
6#![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// ---------------------------------------------------------------------------
15// Variant
16// ---------------------------------------------------------------------------
17
18/// A single experiment variant with routing metadata.
19#[derive(Clone, Debug)]
20pub struct Variant {
21    /// Unique identifier for this variant.
22    pub id: String,
23    /// Model name to use for this variant.
24    pub model: String,
25    /// String prepended/appended to the prompt for this variant.
26    pub prompt_modifier: String,
27    /// Relative traffic weight (will be normalised against sum of all weights).
28    pub weight: f64,
29}
30
31// ---------------------------------------------------------------------------
32// Assignment
33// ---------------------------------------------------------------------------
34
35/// A recorded session→variant assignment.
36#[derive(Debug)]
37pub struct Assignment {
38    /// The variant this session was assigned to.
39    pub variant_id: String,
40    /// The session that was assigned.
41    pub session_id: String,
42    /// Wall-clock instant of assignment.
43    pub assigned_at: std::time::Instant,
44}
45
46// ---------------------------------------------------------------------------
47// VariantMetrics
48// ---------------------------------------------------------------------------
49
50/// Atomic per-variant statistics accumulated during the test.
51pub struct VariantMetrics {
52    /// The variant these metrics belong to.
53    pub variant_id: String,
54    /// Total requests routed to this variant.
55    pub requests: AtomicU64,
56    /// Sum of observed latencies in milliseconds.
57    pub total_latency_ms: AtomicU64,
58    /// Sum of tokens consumed.
59    pub total_tokens: AtomicU64,
60    /// Number of successful completions.
61    pub successes: AtomicU64,
62}
63
64impl VariantMetrics {
65    /// Create zeroed metrics for a variant.
66    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    /// Average latency in milliseconds, or 0.0 if no requests yet.
77    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    /// Fraction of requests that succeeded (0.0–1.0).
86    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    /// Average tokens per request.
95    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
104// ---------------------------------------------------------------------------
105// AbTest
106// ---------------------------------------------------------------------------
107
108/// An active A/B test comparing two or more variants.
109pub struct AbTest {
110    /// Unique test identifier.
111    pub id: String,
112    /// All variants in this test.
113    pub variants: Vec<Variant>,
114    /// Map of session_id → variant_id for all assigned sessions.
115    pub assignments: RwLock<HashMap<String, String>>,
116    /// Per-variant metrics (Arc so VariantMetrics with atomics can live behind shared refs).
117    pub metrics: HashMap<String, Arc<VariantMetrics>>,
118    /// When the test was started.
119    pub start_time: std::time::Instant,
120    /// Fraction of traffic routed into this test (0.0–1.0).
121    pub traffic_pct: f64,
122}
123
124impl AbTest {
125    /// Create a new A/B test.
126    ///
127    /// * `id` — unique test name
128    /// * `variants` — must be non-empty; weights are normalised internally
129    /// * `traffic_pct` — fraction of sessions to route into this test (0.0–1.0)
130    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    /// Deterministically assign a session to a variant using consistent hashing.
146    ///
147    /// Returns `None` if traffic sampling excludes this session or there are no
148    /// variants.
149    pub fn assign(&self, session_id: &str) -> Option<&Variant> {
150        if self.variants.is_empty() {
151            return None;
152        }
153
154        // Determine whether this session is in the traffic bucket.
155        let hash = hash_str(session_id);
156        if (hash % 100) as f64 >= self.traffic_pct * 100.0 {
157            return None;
158        }
159
160        // Check existing assignment first.
161        {
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        // Assign via weighted CDF.
169        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        // Use a second hash (salted) for variant selection to decouple from
175        // the traffic-inclusion hash above.
176        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        // Fallback: last variant.
189        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    /// Record the outcome of a request for the session's assigned variant.
203    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    /// Return the variant_id of the winner if its success_rate exceeds all
223    /// others by more than 5 percentage points.
224    pub fn winner(&self) -> Option<&str> {
225        if self.variants.len() < 2 {
226            return None;
227        }
228        // Collect (variant_id, success_rate) pairs.
229        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    /// Render a markdown table summarising per-variant statistics.
254    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
279// ---------------------------------------------------------------------------
280// AbTestRegistry
281// ---------------------------------------------------------------------------
282
283/// Central registry that manages the lifecycle of multiple A/B tests.
284pub struct AbTestRegistry {
285    tests: RwLock<HashMap<String, AbTest>>,
286}
287
288impl AbTestRegistry {
289    /// Create an empty registry.
290    pub fn new() -> Self {
291        Self {
292            tests: RwLock::new(HashMap::new()),
293        }
294    }
295
296    /// Register a new test.
297    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    /// Retrieve an immutable reference to a test by id.
304    ///
305    /// Returns a cloned summary string because holding the read-guard across
306    /// arbitrary caller code would require unsafe lifetime tricks.
307    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    /// Remove a concluded test, returning its final summary.
313    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    /// Return the ids of all currently active tests.
322    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
336// ---------------------------------------------------------------------------
337// Private helpers
338// ---------------------------------------------------------------------------
339
340fn hash_str(s: &str) -> u64 {
341    let mut h = DefaultHasher::new();
342    s.hash(&mut h);
343    h.finish()
344}
345
346// ---------------------------------------------------------------------------
347// Tests
348// ---------------------------------------------------------------------------
349
350#[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        // At 0% traffic, all sessions should be excluded.
377        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        // Force assign sessions and record results.
390        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        // Without clear margin winner should be None (both variants get 100%).
396        // Just assert it returns Some or None without panicking.
397        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}