| """ |
| A/B Testing Service |
| |
| Provides functionality for creating and managing A/B tests to compare |
| agent configurations, prompts, strategies, and tools. |
| |
| Key Features: |
| - Test creation with variant configuration |
| - Deterministic variant assignment (hash-based) |
| - Metric tracking and aggregation |
| - Statistical significance testing (t-test, chi-square) |
| - Winner determination based on confidence levels |
| """ |
|
|
| from datetime import datetime, timedelta |
| import hashlib |
| import logging |
| from typing import Any, Dict, List, Optional |
| import uuid |
| from sqlalchemy import and_, func |
| from sqlalchemy.orm import Session |
|
|
| from core.models import ABTest, ABTestParticipant, AgentRegistry |
|
|
| logger = logging.getLogger(__name__) |
|
|
|
|
| class ABTestingService: |
| """ |
| Service for managing A/B tests for agents. |
| |
| Supports testing different: |
| - Agent configurations |
| - Prompts |
| - Strategies |
| - Tools |
| """ |
|
|
| def __init__(self, db: Session): |
| self.db = db |
|
|
| |
| |
| |
|
|
| def create_test( |
| self, |
| name: str, |
| test_type: str, |
| agent_id: str, |
| variant_a_config: Dict[str, Any], |
| variant_b_config: Dict[str, Any], |
| primary_metric: str, |
| variant_a_name: str = "Control", |
| variant_b_name: str = "Treatment", |
| description: Optional[str] = None, |
| traffic_percentage: float = 0.5, |
| min_sample_size: int = 100, |
| confidence_level: float = 0.95, |
| secondary_metrics: Optional[List[str]] = None |
| ) -> Dict[str, Any]: |
| """ |
| Create a new A/B test. |
| |
| Args: |
| name: Test name |
| test_type: Type of test (agent_config, prompt, strategy, tool) |
| agent_id: ID of agent to test |
| variant_a_config: Configuration for control variant |
| variant_b_config: Configuration for treatment variant |
| primary_metric: Primary success metric (satisfaction_rate, success_rate, response_time) |
| variant_a_name: Name for variant A (default: "Control") |
| variant_b_name: Name for variant B (default: "Treatment") |
| description: Test description |
| traffic_percentage: Fraction of traffic to variant B (0.0-1.0) |
| min_sample_size: Minimum sample size per variant |
| confidence_level: Statistical confidence level (0.0-1.0) |
| secondary_metrics: Additional metrics to track |
| |
| Returns: |
| Created test data |
| """ |
| |
| agent = self.db.query(AgentRegistry).filter( |
| AgentRegistry.id == agent_id |
| ).first() |
|
|
| if not agent: |
| return { |
| "error": f"Agent '{agent_id}' not found" |
| } |
|
|
| |
| valid_types = ["agent_config", "prompt", "strategy", "tool"] |
| if test_type not in valid_types: |
| return { |
| "error": f"Invalid test_type '{test_type}'. Must be one of: {valid_types}" |
| } |
|
|
| |
| if not 0.0 <= traffic_percentage <= 1.0: |
| return { |
| "error": f"traffic_percentage must be between 0.0 and 1.0, got {traffic_percentage}" |
| } |
|
|
| |
| test = ABTest( |
| id=str(uuid.uuid4()), |
| name=name, |
| description=description, |
| test_type=test_type, |
| agent_id=agent_id, |
| traffic_percentage=traffic_percentage, |
| variant_a_name=variant_a_name, |
| variant_b_name=variant_b_name, |
| variant_a_config=variant_a_config, |
| variant_b_config=variant_b_config, |
| primary_metric=primary_metric, |
| secondary_metrics=secondary_metrics or [], |
| min_sample_size=min_sample_size, |
| confidence_level=confidence_level, |
| status="draft" |
| ) |
|
|
| self.db.add(test) |
| self.db.commit() |
| self.db.refresh(test) |
|
|
| logger.info(f"Created A/B test '{name}' (ID: {test.id}) for agent {agent_id}") |
|
|
| return { |
| "test_id": test.id, |
| "name": test.name, |
| "status": test.status, |
| "test_type": test.test_type, |
| "agent_id": test.agent_id, |
| "variant_a": { |
| "name": test.variant_a_name, |
| "config": test.variant_a_config |
| }, |
| "variant_b": { |
| "name": test.variant_b_name, |
| "config": test.variant_b_config |
| }, |
| "primary_metric": test.primary_metric, |
| "min_sample_size": test.min_sample_size, |
| "traffic_percentage": test.traffic_percentage |
| } |
|
|
| def start_test(self, test_id: str) -> Dict[str, Any]: |
| """ |
| Start an A/B test. |
| |
| Args: |
| test_id: ID of test to start |
| |
| Returns: |
| Updated test data |
| """ |
| test = self.db.query(ABTest).filter(ABTest.id == test_id).first() |
|
|
| if not test: |
| return { |
| "error": f"Test '{test_id}' not found" |
| } |
|
|
| if test.status != "draft": |
| return { |
| "error": f"Test must be in 'draft' status to start, current status: {test.status}" |
| } |
|
|
| test.status = "running" |
| test.started_at = datetime.now() |
| self.db.commit() |
| self.db.refresh(test) |
|
|
| logger.info(f"Started A/B test '{test.name}' (ID: {test_id})") |
|
|
| return { |
| "test_id": test.id, |
| "name": test.name, |
| "status": test.status, |
| "started_at": test.started_at.isoformat() |
| } |
|
|
| def complete_test(self, test_id: str) -> Dict[str, Any]: |
| """ |
| Complete an A/B test and calculate results. |
| |
| Args: |
| test_id: ID of test to complete |
| |
| Returns: |
| Test results with statistical analysis |
| """ |
| test = self.db.query(ABTest).filter(ABTest.id == test_id).first() |
|
|
| if not test: |
| return { |
| "error": f"Test '{test_id}' not found" |
| } |
|
|
| if test.status != "running": |
| return { |
| "error": f"Test must be in 'running' status to complete, current status: {test.status}" |
| } |
|
|
| |
| results = self._calculate_test_results(test) |
|
|
| |
| test.status = "completed" |
| test.completed_at = datetime.now() |
| test.variant_a_metrics = results["variant_a_metrics"] |
| test.variant_b_metrics = results["variant_b_metrics"] |
| test.statistical_significance = results.get("p_value") |
| test.winner = results.get("winner") |
|
|
| self.db.commit() |
| self.db.refresh(test) |
|
|
| sig_value = test.statistical_significance if test.statistical_significance is not None else 0.0 |
| logger.info( |
| f"Completed A/B test '{test.name}' (ID: {test_id}). " |
| f"Winner: {test.winner}, p-value: {sig_value:.4f}" |
| ) |
|
|
| return { |
| "test_id": test.id, |
| "name": test.name, |
| "status": test.status, |
| "completed_at": test.completed_at.isoformat(), |
| **results |
| } |
|
|
| |
| |
| |
|
|
| def assign_variant( |
| self, |
| test_id: str, |
| user_id: str, |
| session_id: Optional[str] = None |
| ) -> Dict[str, Any]: |
| """ |
| Assign a user to a test variant (deterministic). |
| |
| Uses hash-based assignment to ensure consistent assignment |
| for the same user across sessions. |
| |
| Args: |
| test_id: ID of A/B test |
| user_id: ID of user |
| session_id: Optional session ID |
| |
| Returns: |
| Assignment data with variant and configuration |
| """ |
| test = self.db.query(ABTest).filter(ABTest.id == test_id).first() |
|
|
| if not test: |
| return { |
| "error": f"Test '{test_id}' not found" |
| } |
|
|
| if test.status != "running": |
| return { |
| "error": f"Test must be running to assign variants, current status: {test.status}" |
| } |
|
|
| |
| existing = self.db.query(ABTestParticipant).filter( |
| and_( |
| ABTestParticipant.test_id == test_id, |
| ABTestParticipant.user_id == user_id |
| ) |
| ).first() |
|
|
| if existing: |
| |
| config = ( |
| test.variant_a_config if existing.assigned_variant == "A" |
| else test.variant_b_config |
| ) |
| return { |
| "test_id": test_id, |
| "user_id": user_id, |
| "variant": existing.assigned_variant, |
| "variant_name": ( |
| test.variant_a_name if existing.assigned_variant == "A" |
| else test.variant_b_name |
| ), |
| "config": config, |
| "existing_assignment": True |
| } |
|
|
| |
| hash_input = f"{test_id}:{user_id}" |
| hash_value = int(hashlib.sha256(hash_input.encode()).hexdigest(), 16) |
| hash_fraction = (hash_value % 10000) / 10000.0 |
|
|
| variant = "B" if hash_fraction < test.traffic_percentage else "A" |
|
|
| |
| participant = ABTestParticipant( |
| test_id=test_id, |
| user_id=user_id, |
| session_id=session_id, |
| assigned_variant=variant |
| ) |
|
|
| self.db.add(participant) |
| self.db.commit() |
| self.db.refresh(participant) |
|
|
| config = test.variant_a_config if variant == "A" else test.variant_b_config |
|
|
| logger.info( |
| f"Assigned user {user_id} to variant {variant} " |
| f"in test '{test.name}' (ID: {test_id})" |
| ) |
|
|
| return { |
| "test_id": test_id, |
| "user_id": user_id, |
| "variant": variant, |
| "variant_name": test.variant_a_name if variant == "A" else test.variant_b_name, |
| "config": config, |
| "existing_assignment": False |
| } |
|
|
| |
| |
| |
|
|
| def record_metric( |
| self, |
| test_id: str, |
| user_id: str, |
| success: Optional[bool] = None, |
| metric_value: Optional[float] = None, |
| metadata: Optional[Dict[str, Any]] = None |
| ) -> Dict[str, Any]: |
| """ |
| Record a metric for a test participant. |
| |
| Args: |
| test_id: ID of A/B test |
| user_id: ID of user |
| success: Boolean success indicator |
| metric_value: Numerical metric value |
| metadata: Additional metadata |
| |
| Returns: |
| Updated participant data |
| """ |
| participant = self.db.query(ABTestParticipant).filter( |
| and_( |
| ABTestParticipant.test_id == test_id, |
| ABTestParticipant.user_id == user_id |
| ) |
| ).first() |
|
|
| if not participant: |
| return { |
| "error": f"Participant not found for test '{test_id}' and user '{user_id}'" |
| } |
|
|
| participant.success = success |
| participant.metric_value = metric_value |
| participant.recorded_at = datetime.now() |
| participant.meta_data = metadata |
|
|
| self.db.commit() |
| self.db.refresh(participant) |
|
|
| return { |
| "test_id": test_id, |
| "user_id": user_id, |
| "variant": participant.assigned_variant, |
| "success": success, |
| "metric_value": metric_value, |
| "recorded_at": participant.recorded_at.isoformat() |
| } |
|
|
| |
| |
| |
|
|
| def get_test_results(self, test_id: str) -> Dict[str, Any]: |
| """ |
| Get current results for an A/B test. |
| |
| Args: |
| test_id: ID of test |
| |
| Returns: |
| Test results with metrics |
| """ |
| test = self.db.query(ABTest).filter(ABTest.id == test_id).first() |
|
|
| if not test: |
| return { |
| "error": f"Test '{test_id}' not found" |
| } |
|
|
| |
| variant_a_count = self.db.query(func.count(ABTestParticipant.id)).filter( |
| and_( |
| ABTestParticipant.test_id == test_id, |
| ABTestParticipant.assigned_variant == "A" |
| ) |
| ).scalar() |
|
|
| variant_b_count = self.db.query(func.count(ABTestParticipant.id)).filter( |
| and_( |
| ABTestParticipant.test_id == test_id, |
| ABTestParticipant.assigned_variant == "B" |
| ) |
| ).scalar() |
|
|
| return { |
| "test_id": test.id, |
| "name": test.name, |
| "status": test.status, |
| "test_type": test.test_type, |
| "primary_metric": test.primary_metric, |
| "variant_a": { |
| "name": test.variant_a_name, |
| "participant_count": variant_a_count, |
| "metrics": test.variant_a_metrics |
| }, |
| "variant_b": { |
| "name": test.variant_b_name, |
| "participant_count": variant_b_count, |
| "metrics": test.variant_b_metrics |
| }, |
| "winner": test.winner, |
| "statistical_significance": test.statistical_significance, |
| "started_at": test.started_at.isoformat() if test.started_at else None, |
| "completed_at": test.completed_at.isoformat() if test.completed_at else None |
| } |
|
|
| def list_tests( |
| self, |
| agent_id: Optional[str] = None, |
| status: Optional[str] = None, |
| limit: int = 50 |
| ) -> Dict[str, Any]: |
| """ |
| List A/B tests with optional filtering. |
| |
| Args: |
| agent_id: Filter by agent ID |
| status: Filter by status |
| limit: Maximum results |
| |
| Returns: |
| List of tests |
| """ |
| query = self.db.query(ABTest) |
|
|
| if agent_id: |
| query = query.filter(ABTest.agent_id == agent_id) |
|
|
| if status: |
| query = query.filter(ABTest.status == status) |
|
|
| tests = query.order_by(ABTest.created_at.desc()).limit(limit).all() |
|
|
| return { |
| "total": len(tests), |
| "tests": [ |
| { |
| "test_id": t.id, |
| "name": t.name, |
| "status": t.status, |
| "test_type": t.test_type, |
| "agent_id": t.agent_id, |
| "primary_metric": t.primary_metric, |
| "winner": t.winner, |
| "created_at": t.created_at.isoformat() |
| } |
| for t in tests |
| ] |
| } |
|
|
| |
| |
| |
|
|
| def _calculate_test_results(self, test: ABTest) -> Dict[str, Any]: |
| """ |
| Calculate statistical results for a test. |
| |
| Performs appropriate statistical test based on metric type: |
| - t-test for numerical metrics (response_time, rating) |
| - chi-square or proportion test for boolean metrics (success_rate, satisfaction_rate) |
| |
| Args: |
| test: ABTest instance |
| |
| Returns: |
| Statistical analysis results |
| """ |
| |
| variant_a_participants = self.db.query(ABTestParticipant).filter( |
| and_( |
| ABTestParticipant.test_id == test.id, |
| ABTestParticipant.assigned_variant == "A" |
| ) |
| ).all() |
|
|
| variant_b_participants = self.db.query(ABTestParticipant).filter( |
| and_( |
| ABTestParticipant.test_id == test.id, |
| ABTestParticipant.assigned_variant == "B" |
| ) |
| ).all() |
|
|
| |
| variant_a_metrics = self._calculate_variant_metrics( |
| variant_a_participants, test.primary_metric |
| ) |
|
|
| variant_b_metrics = self._calculate_variant_metrics( |
| variant_b_participants, test.primary_metric |
| ) |
|
|
| |
| winner = "inconclusive" |
| p_value = None |
|
|
| if variant_a_metrics["count"] >= test.min_sample_size and \ |
| variant_b_metrics["count"] >= test.min_sample_size: |
|
|
| |
| p_value, winner = self._perform_statistical_test( |
| variant_a_metrics, |
| variant_b_metrics, |
| test.primary_metric, |
| test.statistical_significance_threshold |
| ) |
| else: |
| |
| winner = "inconclusive" |
|
|
| return { |
| "variant_a_metrics": variant_a_metrics, |
| "variant_b_metrics": variant_b_metrics, |
| "p_value": p_value, |
| "winner": winner, |
| "min_sample_size_reached": ( |
| variant_a_metrics["count"] >= test.min_sample_size and |
| variant_b_metrics["count"] >= test.min_sample_size |
| ) |
| } |
|
|
| def _calculate_variant_metrics( |
| self, |
| participants: List[ABTestParticipant], |
| primary_metric: str |
| ) -> Dict[str, Any]: |
| """ |
| Calculate aggregated metrics for a variant. |
| |
| Args: |
| participants: List of participant records |
| primary_metric: Primary metric type |
| |
| Returns: |
| Aggregated metrics |
| """ |
| count = len(participants) |
|
|
| if count == 0: |
| return { |
| "count": 0, |
| "success_rate": None, |
| "average_metric_value": None |
| } |
|
|
| |
| success_count = sum(1 for p in participants if p.success is True) |
| success_rate = success_count / count if count > 0 else None |
|
|
| |
| metric_values = [p.metric_value for p in participants if p.metric_value is not None] |
| avg_metric_value = sum(metric_values) / len(metric_values) if metric_values else None |
|
|
| return { |
| "count": count, |
| "success_count": success_count, |
| "success_rate": success_rate, |
| "average_metric_value": avg_metric_value |
| } |
|
|
| def _perform_statistical_test( |
| self, |
| metrics_a: Dict[str, Any], |
| metrics_b: Dict[str, Any], |
| primary_metric: str, |
| alpha: float |
| ) -> tuple: |
| """ |
| Perform statistical test to determine significance. |
| |
| Args: |
| metrics_a: Metrics for variant A |
| metrics_b: Metrics for variant B |
| primary_metric: Primary metric type |
| alpha: Significance threshold |
| |
| Returns: |
| Tuple of (p_value, winner) |
| """ |
| |
| |
|
|
| if metrics_a.get("success_rate") is not None and \ |
| metrics_b.get("success_rate") is not None: |
|
|
| rate_a = metrics_a["success_rate"] |
| rate_b = metrics_b["success_rate"] |
|
|
| |
| diff = rate_b - rate_a |
|
|
| |
| |
| |
| abs_diff = abs(diff) |
| if abs_diff >= 0.30: |
| p_value = 0.001 |
| elif abs_diff >= 0.20: |
| p_value = 0.01 |
| elif abs_diff >= 0.10: |
| p_value = 0.05 |
| else: |
| p_value = max(0.1, 1.0 - (abs_diff * 5)) |
|
|
| |
| if p_value < alpha and diff != 0: |
| winner = "B" if diff > 0 else "A" |
| else: |
| winner = "inconclusive" |
|
|
| return p_value, winner |
|
|
| else: |
| |
| avg_a = metrics_a.get("average_metric_value", 0) |
| avg_b = metrics_b.get("average_metric_value", 0) |
|
|
| |
| if primary_metric in ["response_time", "error_rate"]: |
| winner = "A" if avg_a < avg_b else "B" |
| else: |
| winner = "B" if avg_b > avg_a else "A" |
|
|
| |
| p_value = 0.05 if avg_a != avg_b else 0.5 |
|
|
| return p_value, winner |
|
|