File size: 5,240 Bytes
6993919 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 | """A/B Testing Framework — split traffic, measure accuracy, auto-promote winners."""
import hashlib
from datetime import UTC, datetime
from fastapi import APIRouter, Query
router = APIRouter(prefix="/api/v1/ab-testing", tags=["ab-testing"])
_experiments: dict[str, dict] = {}
_results: dict[str, list] = {}
@router.post("/experiment")
async def create_experiment(
name: str,
variants: str = Query(..., description="Comma-separated variant names, e.g. 'prompt_v1,prompt_v2'"),
traffic_split: str = Query("50,50", description="Traffic split percentages"),
metric: str = Query("accuracy", description="Metric to evaluate"),
):
"""Create a new A/B test experiment."""
variant_list = [v.strip() for v in variants.split(",")]
split_list = [int(s.strip()) for s in traffic_split.split(",")]
if len(variant_list) != len(split_list):
return {"error": "Variants and splits must have same length"}
if sum(split_list) != 100:
return {"error": "Traffic splits must sum to 100"}
exp_id = hashlib.md5(name.encode()).hexdigest()[:8]
_experiments[exp_id] = {
"id": exp_id,
"name": name,
"variants": variant_list,
"traffic_split": split_list,
"metric": metric,
"created_at": datetime.now(UTC).isoformat(),
"status": "active",
}
_results[exp_id] = []
return _experiments[exp_id]
@router.get("/assign/{experiment_id}")
async def assign_variant(experiment_id: str, user_id: str = "anonymous"):
"""Assign a user to a variant deterministically (same user always gets same variant)."""
if experiment_id not in _experiments:
return {"variant": "control", "note": "Experiment not found"}
exp = _experiments[experiment_id]
# Deterministic assignment based on user_id hash
h = int(hashlib.md5(f"{experiment_id}:{user_id}".encode()).hexdigest(), 16)
bucket = h % 100
cumulative = 0
for variant, split in zip(exp["variants"], exp["traffic_split"], strict=False):
cumulative += split
if bucket < cumulative:
return {"experiment": experiment_id, "variant": variant, "user": user_id}
return {"experiment": experiment_id, "variant": exp["variants"][-1], "user": user_id}
@router.post("/record/{experiment_id}")
async def record_result(experiment_id: str, variant: str, correct: bool = False):
"""Record a result for an A/B test variant."""
if experiment_id not in _experiments:
return {"error": "Experiment not found"}
_results[experiment_id].append(
{
"variant": variant,
"correct": correct,
"timestamp": datetime.now(UTC).isoformat(),
}
)
# Check if we have enough data to declare a winner (30+ samples per variant)
results = _results[experiment_id]
variant_counts = {}
variant_correct = {}
for r in results:
v = r["variant"]
variant_counts[v] = variant_counts.get(v, 0) + 1
if r["correct"]:
variant_correct[v] = variant_correct.get(v, 0) + 1
# Check if all variants have enough samples
min_samples = 30
all_ready = all(count >= min_samples for count in variant_counts.values())
if all_ready:
scores = {v: variant_correct.get(v, 0) / variant_counts[v] for v in variant_counts}
best = max(scores, key=scores.get)
worst = min(scores, key=scores.get)
improvement = scores[best] - scores[worst]
result = {
"experiment": experiment_id,
"status": "complete" if improvement > 0.05 else "inconclusive",
"winner": best if improvement > 0.05 else None,
"scores": {v: round(s, 3) for v, s in scores.items()},
"samples": variant_counts,
"confidence": round(improvement * 100, 1),
}
if improvement > 0.05:
_experiments[experiment_id]["status"] = "complete"
_experiments[experiment_id]["winner"] = best
return result
return {
"experiment": experiment_id,
"status": "collecting",
"samples": variant_counts,
"min_needed": min_samples,
}
@router.get("/experiments")
async def list_experiments():
return {
"experiments": list(_experiments.values()),
"active": sum(1 for e in _experiments.values() if e["status"] == "active"),
}
@router.get("/experiment/{experiment_id}")
async def get_experiment(experiment_id: str):
if experiment_id not in _experiments:
return {"error": "Experiment not found"}
exp = _experiments[experiment_id]
results = _results.get(experiment_id, [])
return {
**exp,
"total_trials": len(results),
"results": _summarize_results(results, exp["variants"]),
}
def _summarize_results(results: list, variants: list) -> dict:
counts = dict.fromkeys(variants, 0)
correct = dict.fromkeys(variants, 0)
for r in results:
v = r["variant"]
if v in counts:
counts[v] += 1
if r["correct"]:
correct[v] += 1
return {
v: {"trials": counts[v], "correct": correct[v], "accuracy": round(correct[v] / max(counts[v], 1), 3)}
for v in variants
}
|