Compare commits
1 Commits
fix/101-cr
...
fix/issue-
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
cf23e93787 |
@@ -1,152 +0,0 @@
|
||||
"""
|
||||
A/B Test Framework for Crisis Detection in the-door.
|
||||
|
||||
Allows running two crisis detection variants side-by-side with
|
||||
logged outcomes for comparison. No PII stored — only variant labels,
|
||||
levels, and timing.
|
||||
|
||||
Usage:
|
||||
from crisis.ab_testing import ABTestCrisisDetector
|
||||
|
||||
detector = ABTestCrisisDetector(variant_a=detect_v1, variant_b=detect_v2)
|
||||
result, variant = detector.detect("I feel hopeless")
|
||||
# result: CrisisDetectionResult
|
||||
# variant: "A" or "B"
|
||||
|
||||
# Get comparison metrics
|
||||
stats = detector.get_stats()
|
||||
# {"A": {"count": 100, "avg_latency_ms": 2.3, ...}, "B": {...}}
|
||||
"""
|
||||
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
|
||||
from .detect import CrisisDetectionResult
|
||||
|
||||
|
||||
# ── Feature flag ───────────────────────────────────────────────
|
||||
|
||||
def _get_variant_override() -> Optional[str]:
|
||||
"""Check for environment variable override (testing/debugging)."""
|
||||
val = os.environ.get("CRISIS_AB_VARIANT", "").upper()
|
||||
if val in ("A", "B"):
|
||||
return val
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class VariantRecord:
|
||||
"""Single detection event record — no PII, only metadata."""
|
||||
variant: str
|
||||
level: str
|
||||
latency_ms: float
|
||||
indicator_count: int
|
||||
|
||||
|
||||
class ABTestCrisisDetector:
|
||||
"""
|
||||
A/B test wrapper for crisis detection.
|
||||
|
||||
Routes calls to variant A or B based on configurable split,
|
||||
logs outcomes for comparison, and provides aggregate stats.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
variant_a: Callable[[str], CrisisDetectionResult],
|
||||
variant_b: Callable[[str], CrisisDetectionResult],
|
||||
split: float = 0.5,
|
||||
variant_a_name: str = "A",
|
||||
variant_b_name: str = "B",
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
variant_a: First detection function
|
||||
variant_b: Second detection function
|
||||
split: Probability of selecting variant A (0.0 to 1.0)
|
||||
variant_a_name: Label for variant A in reports
|
||||
variant_b_name: Label for variant B in reports
|
||||
"""
|
||||
self.variant_a = variant_a
|
||||
self.variant_b = variant_b
|
||||
self.split = split
|
||||
self.variant_a_name = variant_a_name
|
||||
self.variant_b_name = variant_b_name
|
||||
self.records: List[VariantRecord] = []
|
||||
|
||||
def _select_variant(self) -> str:
|
||||
"""Select variant based on split and optional env override."""
|
||||
override = _get_variant_override()
|
||||
if override:
|
||||
return override
|
||||
return "A" if random.random() < self.split else "B"
|
||||
|
||||
def detect(self, text: str) -> Tuple[CrisisDetectionResult, str]:
|
||||
"""
|
||||
Run detection on the selected variant and log the result.
|
||||
|
||||
Returns:
|
||||
(CrisisDetectionResult, variant_label)
|
||||
"""
|
||||
variant = self._select_variant()
|
||||
|
||||
if variant == "A":
|
||||
fn = self.variant_a
|
||||
else:
|
||||
fn = self.variant_b
|
||||
|
||||
start = time.perf_counter()
|
||||
result = fn(text)
|
||||
latency_ms = (time.perf_counter() - start) * 1000
|
||||
|
||||
# Log record (no PII — only level, timing, count)
|
||||
record = VariantRecord(
|
||||
variant=variant,
|
||||
level=result.level,
|
||||
latency_ms=latency_ms,
|
||||
indicator_count=len(result.indicators),
|
||||
)
|
||||
self.records.append(record)
|
||||
|
||||
return result, variant
|
||||
|
||||
def get_stats(self) -> Dict[str, dict]:
|
||||
"""
|
||||
Get per-variant comparison statistics.
|
||||
|
||||
Returns dict with variant labels as keys:
|
||||
{
|
||||
"A": {"count": 100, "avg_latency_ms": 2.3, "levels": {...}},
|
||||
"B": {"count": 95, "avg_latency_ms": 3.1, "levels": {...}}
|
||||
"""
|
||||
stats = {}
|
||||
for label in ("A", "B"):
|
||||
recs = [r for r in self.records if r.variant == label]
|
||||
if not recs:
|
||||
stats[label] = {"count": 0}
|
||||
continue
|
||||
|
||||
latencies = [r.latency_ms for r in recs]
|
||||
levels = {}
|
||||
for r in recs:
|
||||
levels[r.level] = levels.get(r.level, 0) + 1
|
||||
|
||||
stats[label] = {
|
||||
"count": len(recs),
|
||||
"avg_latency_ms": round(sum(latencies) / len(latencies), 2),
|
||||
"max_latency_ms": round(max(latencies), 2),
|
||||
"min_latency_ms": round(min(latencies), 2),
|
||||
"levels": levels,
|
||||
"avg_indicators": round(
|
||||
sum(r.indicator_count for r in recs) / len(recs), 2
|
||||
),
|
||||
}
|
||||
|
||||
return stats
|
||||
|
||||
def reset(self) -> None:
|
||||
"""Clear all records. For testing."""
|
||||
self.records.clear()
|
||||
43
index.html
43
index.html
@@ -475,6 +475,26 @@ html, body {
|
||||
margin-bottom: 24px;
|
||||
}
|
||||
|
||||
.modal-status {
|
||||
min-height: 22px;
|
||||
margin: 0 0 16px;
|
||||
font-size: 0.9rem;
|
||||
line-height: 1.45;
|
||||
color: #8b949e;
|
||||
}
|
||||
|
||||
.modal-status.is-visible {
|
||||
display: block;
|
||||
}
|
||||
|
||||
.modal-status.success {
|
||||
color: #3fb950;
|
||||
}
|
||||
|
||||
.modal-status.error {
|
||||
color: #ff7b72;
|
||||
}
|
||||
|
||||
.form-group {
|
||||
margin-bottom: 16px;
|
||||
}
|
||||
@@ -737,6 +757,7 @@ html, body {
|
||||
<textarea id="sp-environment" placeholder="e.g., Giving my car keys to a friend, locking away meds..."></textarea>
|
||||
</div>
|
||||
</div>
|
||||
<div id="safety-plan-status" class="modal-status" role="status" aria-live="polite" aria-atomic="true"></div>
|
||||
<div class="modal-footer">
|
||||
<button class="btn btn-secondary" id="cancel-safety-plan">Cancel</button>
|
||||
<button class="btn btn-primary" id="save-safety-plan">Save Plan</button>
|
||||
@@ -818,6 +839,7 @@ Sovereignty and service always.`;
|
||||
var closeSafetyPlan = document.getElementById('close-safety-plan');
|
||||
var cancelSafetyPlan = document.getElementById('cancel-safety-plan');
|
||||
var saveSafetyPlan = document.getElementById('save-safety-plan');
|
||||
var safetyPlanStatus = document.getElementById('safety-plan-status');
|
||||
var clearChatBtn = document.getElementById('clear-chat-btn');
|
||||
|
||||
// ===== STATE =====
|
||||
@@ -1183,12 +1205,24 @@ Sovereignty and service always.`;
|
||||
} catch (e) {}
|
||||
}
|
||||
|
||||
function setSafetyPlanStatus(message, type) {
|
||||
safetyPlanStatus.textContent = message;
|
||||
safetyPlanStatus.className = 'modal-status is-visible ' + (type || '');
|
||||
}
|
||||
|
||||
function clearSafetyPlanStatus() {
|
||||
safetyPlanStatus.textContent = '';
|
||||
safetyPlanStatus.className = 'modal-status';
|
||||
}
|
||||
|
||||
closeSafetyPlan.addEventListener('click', function() {
|
||||
clearSafetyPlanStatus();
|
||||
safetyPlanModal.classList.remove('active');
|
||||
_restoreSafetyPlanFocus();
|
||||
});
|
||||
|
||||
cancelSafetyPlan.addEventListener('click', function() {
|
||||
clearSafetyPlanStatus();
|
||||
safetyPlanModal.classList.remove('active');
|
||||
_restoreSafetyPlanFocus();
|
||||
});
|
||||
@@ -1203,11 +1237,9 @@ Sovereignty and service always.`;
|
||||
};
|
||||
try {
|
||||
localStorage.setItem('timmy_safety_plan', JSON.stringify(plan));
|
||||
safetyPlanModal.classList.remove('active');
|
||||
_restoreSafetyPlanFocus();
|
||||
alert('Safety plan saved locally.');
|
||||
setSafetyPlanStatus('Safety plan saved locally.', 'success');
|
||||
} catch (e) {
|
||||
alert('Error saving plan.');
|
||||
setSafetyPlanStatus('Error saving plan.', 'error');
|
||||
}
|
||||
});
|
||||
|
||||
@@ -1285,6 +1317,7 @@ Sovereignty and service always.`;
|
||||
|
||||
// Wire open buttons to activate focus trap
|
||||
safetyPlanBtn.addEventListener('click', function() {
|
||||
clearSafetyPlanStatus();
|
||||
loadSafetyPlan();
|
||||
safetyPlanModal.classList.add('active');
|
||||
_activateSafetyPlanFocusTrap(safetyPlanBtn);
|
||||
@@ -1293,6 +1326,8 @@ Sovereignty and service always.`;
|
||||
// Crisis panel safety plan button (if crisis panel is visible)
|
||||
if (crisisSafetyPlanBtn) {
|
||||
crisisSafetyPlanBtn.addEventListener('click', function() {
|
||||
clearSafetyPlanStatus();
|
||||
clearSafetyPlanStatus();
|
||||
loadSafetyPlan();
|
||||
safetyPlanModal.classList.add('active');
|
||||
_activateSafetyPlanFocusTrap(crisisSafetyPlanBtn);
|
||||
|
||||
@@ -1,129 +0,0 @@
|
||||
"""
|
||||
Tests for crisis/ab_testing.py — A/B test framework for crisis detection.
|
||||
|
||||
Verifies variant selection, logging, stats aggregation, and env override.
|
||||
"""
|
||||
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from crisis.ab_testing import ABTestCrisisDetector
|
||||
from crisis.detect import CrisisDetectionResult, detect_crisis
|
||||
|
||||
|
||||
def _make_variant(level: str):
|
||||
"""Create a mock detection function that returns a fixed level."""
|
||||
def fn(text: str) -> CrisisDetectionResult:
|
||||
return CrisisDetectionResult(level=level, indicators=[f"mock_{level}"])
|
||||
return fn
|
||||
|
||||
|
||||
class TestABTestCrisisDetector:
|
||||
"""A/B test framework unit tests."""
|
||||
|
||||
def setup_method(self):
|
||||
"""Ensure no env override."""
|
||||
os.environ.pop("CRISIS_AB_VARIANT", None)
|
||||
|
||||
def test_returns_result_and_variant(self):
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
result, variant = detector.detect("test message")
|
||||
assert isinstance(result, CrisisDetectionResult)
|
||||
assert variant in ("A", "B")
|
||||
|
||||
def test_records_are_logged(self):
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
# Force variant A
|
||||
with patch.object(detector, "_select_variant", return_value="A"):
|
||||
detector.detect("test")
|
||||
assert len(detector.records) == 1
|
||||
assert detector.records[0].variant == "A"
|
||||
assert detector.records[0].level == "LOW"
|
||||
|
||||
def test_stats_empty(self):
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
stats = detector.get_stats()
|
||||
assert stats["A"]["count"] == 0
|
||||
assert stats["B"]["count"] == 0
|
||||
|
||||
def test_stats_with_data(self):
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
# Force 5 A and 3 B
|
||||
with patch.object(detector, "_select_variant", side_effect=["A"] * 5 + ["B"] * 3):
|
||||
for _ in range(8):
|
||||
detector.detect("test")
|
||||
|
||||
stats = detector.get_stats()
|
||||
assert stats["A"]["count"] == 5
|
||||
assert stats["B"]["count"] == 3
|
||||
assert "avg_latency_ms" in stats["A"]
|
||||
assert stats["A"]["levels"]["LOW"] == 5
|
||||
assert stats["B"]["levels"]["HIGH"] == 3
|
||||
|
||||
def test_env_override_a(self):
|
||||
os.environ["CRISIS_AB_VARIANT"] = "A"
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
for _ in range(10):
|
||||
result, variant = detector.detect("test")
|
||||
assert variant == "A"
|
||||
assert result.level == "LOW"
|
||||
|
||||
def test_env_override_b(self):
|
||||
os.environ["CRISIS_AB_VARIANT"] = "b"
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
for _ in range(10):
|
||||
result, variant = detector.detect("test")
|
||||
assert variant == "B"
|
||||
assert result.level == "HIGH"
|
||||
|
||||
def test_reset_clears_records(self):
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
)
|
||||
detector.detect("test")
|
||||
detector.detect("test")
|
||||
assert len(detector.records) == 2
|
||||
detector.reset()
|
||||
assert len(detector.records) == 0
|
||||
|
||||
def test_split_respected(self):
|
||||
"""With split=1.0, always get variant A."""
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=_make_variant("LOW"),
|
||||
variant_b=_make_variant("HIGH"),
|
||||
split=1.0,
|
||||
)
|
||||
for _ in range(10):
|
||||
_, variant = detector.detect("test")
|
||||
assert variant == "A"
|
||||
|
||||
def test_with_real_detector(self):
|
||||
"""Integration test using actual detect_crisis as both variants."""
|
||||
detector = ABTestCrisisDetector(
|
||||
variant_a=detect_crisis,
|
||||
variant_b=detect_crisis,
|
||||
)
|
||||
result, variant = detector.detect("I want to kill myself")
|
||||
assert result.level == "CRITICAL"
|
||||
assert variant in ("A", "B")
|
||||
52
tests/test_safety_plan_save_feedback.py
Normal file
52
tests/test_safety_plan_save_feedback.py
Normal file
@@ -0,0 +1,52 @@
|
||||
import pathlib
|
||||
import re
|
||||
import unittest
|
||||
|
||||
ROOT = pathlib.Path(__file__).resolve().parents[1]
|
||||
INDEX_HTML = ROOT / 'index.html'
|
||||
|
||||
|
||||
class TestSafetyPlanSaveFeedback(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.html = INDEX_HTML.read_text()
|
||||
|
||||
def test_modal_has_inline_status_live_region(self):
|
||||
self.assertRegex(
|
||||
self.html,
|
||||
r'<div[^>]+id="safety-plan-status"[^>]+role="status"[^>]+aria-live="polite"[^>]*>',
|
||||
'Expected an inline polite live region for safety plan save feedback.',
|
||||
)
|
||||
|
||||
def test_save_feedback_does_not_use_blocking_alerts(self):
|
||||
self.assertNotIn(
|
||||
"alert('Safety plan saved locally.')",
|
||||
self.html,
|
||||
'Expected success feedback to stop using blocking alert().',
|
||||
)
|
||||
self.assertNotIn(
|
||||
"alert('Error saving plan.')",
|
||||
self.html,
|
||||
'Expected error feedback to stop using blocking alert().',
|
||||
)
|
||||
|
||||
def test_save_logic_updates_inline_status_for_success_and_error(self):
|
||||
self.assertRegex(
|
||||
self.html,
|
||||
r'function\s+setSafetyPlanStatus\s*\(',
|
||||
'Expected a helper to update inline save feedback.',
|
||||
)
|
||||
self.assertRegex(
|
||||
self.html,
|
||||
r"setSafetyPlanStatus\('Safety plan saved locally\.'\s*,\s*'success'\)",
|
||||
'Expected success path to update inline status.',
|
||||
)
|
||||
self.assertRegex(
|
||||
self.html,
|
||||
r"setSafetyPlanStatus\('Error saving plan\.'\s*,\s*'error'\)",
|
||||
'Expected error path to update inline status.',
|
||||
)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user