Artifex/agents/replay_arena.py
2026-08-15 18:14:21 +07:00

520 lines
31 KiB
Python

from __future__ import annotations
import shutil
import subprocess
import time
from pathlib import Path
from statistics import median
from typing import Any
from django.core.exceptions import ValidationError
from django.db import transaction
from django.utils import timezone
from agents.progeny import ProgenyService
from control_plane.agents.models import (
AgentVersion,
ExperimentComparison,
ExperimentVariant,
ImprovementCandidate,
ProgenyExperiment,
ReplayCase,
ReplayDataset,
ReplayResult,
ReplayRun,
)
from control_plane.events.bus import EventBus
from control_plane.projects.models import CommitRecord, Milestone, Project, ProjectPlan, Task, TaskStatus, Worktree
from control_plane.verification.models import Review, TestRun, Verification, VerificationResult
from graph.bootstrap import champion_task_execution_graph_v1
from graph.langgraph_runtime import LangGraphRuntime
from graph.models import ExecutionGraphDefinition, ExecutionGraphVersion, ExecutionGraphVersionStatus, GraphApproval, GraphApprovalStatus, GraphRun
from graph.native_runtime import GraphExecutionContext
from graph.registry import NodeHandlerRegistry, NodeResult
from graph.task_execution import task_execution_graph_v2_static_analysis
from graph.task_nodes import TaskExecutionServices, task_execution_registry
from model_router.router import ModelRouter
INFRA_FAILURES = {"PROVIDER_FAILURE", "INFRASTRUCTURE_FAILURE", "REPLAY_RUNTIME_FAILURE", "EVALUATOR_FAILURE"}
class ReplayArena:
def __init__(self, router: ModelRouter | None = None, bus: EventBus | None = None, *, test_command: list[str] | None = None) -> None:
self.router = router or ModelRouter({})
self.bus = bus or EventBus()
self.test_command = test_command or ["python", "-m", "pytest"]
def create_dataset(self, name: str, *, description: str = "", version: int = 1, selection_criteria: dict[str, object] | None = None) -> ReplayDataset:
return ReplayDataset.objects.create(name=name, description=description, version=version, selection_criteria=selection_criteria or {})
def add_case_from_task(self, dataset: ReplayDataset, task: Task, *, failure_classification: str = "", selection_metadata: dict[str, object] | None = None) -> ReplayCase:
baseline = self._repository_head(Path(task.project.repository_path)) if task.project.repository_path else ""
return ReplayCase.objects.create(
replay_dataset=dataset,
source_project=task.project,
source_task=task,
source_task_attempt=task.attempts.order_by("-attempt_number").first(),
source_graph_run=task.graph_runs.order_by("-created_at").first(),
task_type=task.task_type,
project_type=task.project.project_type,
original_goal=task.goal,
acceptance_criteria=task.acceptance_criteria,
repository_path=task.project.repository_path,
repository_baseline_ref=baseline,
expected_evaluator_inputs={"acceptance_criteria": task.acceptance_criteria},
failure_classification=failure_classification,
selection_metadata=selection_metadata or {},
)
def create_dataset_from_history(self, name: str, **filters: object) -> ReplayDataset:
dataset = self.create_dataset(name, selection_criteria=filters)
tasks = Task.objects.select_related("project", "milestone").all()
if filters.get("task_type"):
tasks = tasks.filter(task_type=filters["task_type"])
if filters.get("project"):
tasks = tasks.filter(project=filters["project"])
if filters.get("outcome"):
tasks = tasks.filter(status=filters["outcome"])
if filters.get("graph_version"):
tasks = tasks.filter(graph_runs__execution_graph_version=filters["graph_version"])
if filters.get("agent_version"):
tasks = tasks.filter(attempts__coder=filters["agent_version"])
if filters.get("date_start"):
tasks = tasks.filter(created_at__gte=filters["date_start"])
if filters.get("date_end"):
tasks = tasks.filter(created_at__lte=filters["date_end"])
if filters.get("signal_grouping_key"):
tasks = tasks.filter(progenysignal__grouping_key=filters["signal_grouping_key"])
max_cases = int(filters.get("max_cases", 10))
for task in tasks.distinct().order_by("created_at")[:max_cases]:
if task.project.repository_path:
self.add_case_from_task(dataset, task)
return dataset
def freeze_dataset(self, dataset: ReplayDataset) -> ReplayDataset:
if not dataset.cases.exists():
raise ValidationError("ReplayDataset must contain at least one case before freezing.")
dataset.status = "FROZEN"
dataset.frozen_at = timezone.now()
dataset.save(update_fields=["status", "frozen_at", "updated_at"])
return dataset
def create_replacement_dataset_version(self, dataset: ReplayDataset) -> ReplayDataset:
return ReplayDataset.objects.create(
name=dataset.name,
description=dataset.description,
version=dataset.version + 1,
selection_criteria=dataset.selection_criteria,
metadata={"replaces_dataset_id": str(dataset.id)},
)
def create_experiment(self, candidate: ImprovementCandidate, dataset: ReplayDataset, *, success_criteria: dict[str, object] | None = None) -> ProgenyExperiment:
if dataset.status != "FROZEN":
raise ValidationError("Experiments require a frozen replay dataset.")
experiment = ProgenyExperiment.objects.create(
investigation=candidate.investigation,
improvement_candidate=candidate,
target_type=candidate.target_type,
target_identifier=candidate.target_id,
replay_dataset=dataset,
hypothesis=candidate.hypothesis,
success_criteria=success_criteria or {"minimum_replay_cases": 3, "confidence_threshold": 0.7},
status="DRAFT",
metadata={"controls": {"single_variable": True}},
)
return experiment
def add_champion(self, experiment: ProgenyExperiment, *, graph_version: ExecutionGraphVersion | None = None, agent_version: AgentVersion | None = None) -> ExperimentVariant:
return self._add_variant(experiment, "CHAMPION", graph_version=graph_version, agent_version=agent_version)
def add_challenger(self, experiment: ProgenyExperiment, *, graph_version: ExecutionGraphVersion | None = None, agent_version: AgentVersion | None = None) -> ExperimentVariant:
return self._add_variant(experiment, "CHALLENGER", graph_version=graph_version, agent_version=agent_version)
def ensure_static_analysis_graph_challenger(self) -> ExecutionGraphVersion:
spec = task_execution_graph_v2_static_analysis()
definition, _ = ExecutionGraphDefinition.objects.get_or_create(name=spec.name, defaults={"graph_type": spec.graph_type, "description": "Task execution graph"})
version, _ = ExecutionGraphVersion.objects.get_or_create(
graph=definition,
version=spec.version,
defaults={"status": ExecutionGraphVersionStatus.CHALLENGER, "graph_spec": spec.to_dict(), "metadata": {"parent_version": 1, "change_summary": spec.metadata["change_summary"]}},
)
return version
def run_experiment(self, experiment: ProgenyExperiment, *, max_cases: int | None = None) -> ProgenyExperiment:
budget = dict(experiment.metadata.get("budget", {})) if isinstance(experiment.metadata, dict) else {}
case_cap = max_cases or int(budget.get("maximum_replay_cases", experiment.replay_dataset.cases.count()))
model_request_cap = int(budget.get("maximum_model_requests", 10**9))
started = time.monotonic()
runtime_budget = float(budget.get("runtime_budget_seconds", 10**9))
experiment.status = "RUNNING"
experiment.started_at = timezone.now()
experiment.save(update_fields=["status", "started_at", "updated_at"])
for replay_case in experiment.replay_dataset.cases.order_by("created_at")[:case_cap]:
for variant in experiment.variants.order_by("role", "created_at"):
used_requests = sum(int(run.telemetry.get("model_requests_per_task", 0)) for run in experiment.replay_runs.all())
if used_requests >= model_request_cap or time.monotonic() - started > runtime_budget:
experiment.status = "FAILED"
experiment.completed_at = timezone.now()
experiment.metadata = {**experiment.metadata, "incomplete_reason": "budget_exceeded"}
experiment.save(update_fields=["status", "completed_at", "metadata", "updated_at"])
return experiment
self.run_case(experiment, replay_case, variant)
experiment.status = "COMPLETE"
experiment.completed_at = timezone.now()
experiment.save(update_fields=["status", "completed_at", "updated_at"])
self.compare(experiment)
return experiment
def run_case(self, experiment: ProgenyExperiment, replay_case: ReplayCase, variant: ExperimentVariant) -> ReplayRun:
run = ReplayRun.objects.create(experiment=experiment, replay_case=replay_case, variant=variant, started_at=timezone.now(), status="RUNNING")
replay_repo = self._fresh_replay_repository(replay_case, variant)
try:
replay_task = self._clone_task(replay_case, replay_repo, variant)
graph_version = variant.execution_graph_version or champion_task_execution_graph_v1()
graph_run = GraphRun.objects.create(
execution_graph_version=graph_version,
project=replay_task.project,
milestone=replay_task.milestone,
task=replay_task,
current_node=graph_version.graph_spec["entry"],
metadata={"replay_experiment_id": str(experiment.id), "replay_variant_id": str(variant.id), "replay_case_id": str(replay_case.id)},
)
run.graph_run = graph_run
run.replay_task = replay_task
run.save(update_fields=["graph_run", "replay_task", "updated_at"])
overrides = {}
if variant.agent_version_id:
overrides[variant.agent_version.agent.role] = variant.agent_version
services = TaskExecutionServices(self.router, bus=self.bus, test_command=self.test_command, agent_overrides=overrides)
LangGraphRuntime(task_execution_registry(services), bus=self.bus).run_until_terminal_or_paused(graph_run)
replay_task.refresh_from_db()
commit = CommitRecord.objects.filter(task=replay_task).first()
run.graph_run = graph_run
run.replay_task = replay_task
run.commit_candidate_sha = commit.sha if commit else ""
run.status = "COMPLETE" if replay_task.status == TaskStatus.COMPLETE else "FAILED"
run.failure_classification = "VARIANT_FAILURE" if run.status == "FAILED" else ""
run.telemetry = self._run_telemetry(replay_task, graph_run)
run.metadata = {"replay_repository_path": str(replay_repo), "production_safe": True, "commit_label": "REPLAY / EXPERIMENTAL"}
run.completed_at = timezone.now()
run.save(update_fields=["graph_run", "replay_task", "commit_candidate_sha", "status", "failure_classification", "telemetry", "metadata", "completed_at", "updated_at"])
self._persist_result(run, replay_task)
except Exception as exc:
run.status = "FAILED"
run.failure_classification = self._classify_replay_exception(exc)
run.failure_evidence = {"error": str(exc)}
run.completed_at = timezone.now()
run.save(update_fields=["status", "failure_classification", "failure_evidence", "completed_at", "updated_at"])
return run
def compare(self, experiment: ProgenyExperiment) -> ExperimentComparison:
champion = experiment.variants.get(role="CHAMPION")
challenger = experiment.variants.filter(role="CHALLENGER").order_by("created_at").first()
if challenger is None:
raise ValidationError("Experiment requires a challenger variant.")
champion_results = self._eligible_results(experiment, champion)
challenger_results = self._eligible_results(experiment, challenger)
aggregate = {"CHAMPION": self._aggregate(champion_results), "CHALLENGER": self._aggregate(challenger_results)}
aggregate["DELTA"] = self._delta(aggregate["CHAMPION"], aggregate["CHALLENGER"])
paired = self._paired_outcomes(experiment, champion, challenger)
verdict, reasons = self.judge_experiment(experiment, aggregate, paired)
comparison, _ = ExperimentComparison.objects.update_or_create(
experiment=experiment,
defaults={
"champion_variant": champion,
"challenger_variant": challenger,
"aggregate_metrics": aggregate,
"paired_outcomes": paired,
"regression_cases": paired["regression_cases"],
"verdict": verdict,
"reasons": reasons,
},
)
return comparison
def judge_experiment(self, experiment: ProgenyExperiment, aggregate: dict[str, Any], paired: dict[str, Any]) -> tuple[str, list[str]]:
minimum = int(experiment.success_criteria.get("minimum_replay_cases", 3))
reasons: list[str] = []
case_count = paired["case_count"]
if "CHAMPION" in aggregate and "CHALLENGER" in aggregate and (aggregate["CHAMPION"].get("case_count", 0) == 0 or aggregate["CHALLENGER"].get("case_count", 0) == 0):
return "RUN_MORE_REPLAYS", ["Comparable non-infrastructure results are missing for at least one variant."]
if case_count < minimum:
return "RUN_MORE_REPLAYS", [f"Only {case_count} paired replay cases; minimum is {minimum}."]
if paired["CHAMPION_ONLY_PASS"]:
return "REJECT_RECOMMENDED", ["Challenger regressed cases that champion passed."]
quality_delta = aggregate["DELTA"].get("accepted_candidate_rate", 0)
runtime_delta = aggregate["DELTA"].get("median_runtime_seconds", 0)
if quality_delta > 0.05:
reasons.append("Challenger materially improves accepted-candidate rate without critical regressions.")
return "PROMOTE_RECOMMENDED", reasons
if abs(quality_delta) <= 0.01 and runtime_delta < -0.1:
reasons.append("Challenger is quality-equivalent with meaningful runtime improvement.")
return "PROMOTE_RECOMMENDED", reasons
if quality_delta < -0.01:
return "REJECT_RECOMMENDED", ["Challenger quality is worse than champion."]
return "INCONCLUSIVE", ["No material quality or efficiency improvement detected."]
def approve_promotion(self, comparison: ExperimentComparison, *, actor: str = "human") -> ExperimentComparison:
challenger = comparison.challenger_variant
if challenger is None:
raise ValidationError("Comparison has no challenger variant.")
with transaction.atomic():
if challenger.execution_graph_version_id:
graph = challenger.execution_graph_version.graph
ExecutionGraphVersion.objects.filter(graph=graph, status=ExecutionGraphVersionStatus.CHAMPION).update(status=ExecutionGraphVersionStatus.RETIRED)
challenger.execution_graph_version.status = ExecutionGraphVersionStatus.CHAMPION
challenger.execution_graph_version.promoted_at = timezone.now()
challenger.execution_graph_version.save(update_fields=["status", "promoted_at"])
self.bus.publish("EXECUTION_GRAPH_PROMOTED", actor=actor, payload={"experiment_id": str(comparison.experiment_id), "graph_version_id": str(challenger.execution_graph_version_id)})
if challenger.agent_version_id:
agent = challenger.agent_version.agent
if agent.champion_version_id:
agent.champion_version.promotion_status = "CANDIDATE"
agent.champion_version.save(update_fields=["promotion_status", "updated_at"])
challenger.agent_version.promotion_status = "CHAMPION"
challenger.agent_version.save(update_fields=["promotion_status", "updated_at"])
agent.champion_version = challenger.agent_version
agent.save(update_fields=["champion_version", "updated_at"])
self.bus.publish("AGENT_EXPERIMENT_PROMOTED", actor=actor, payload={"experiment_id": str(comparison.experiment_id), "agent_version_id": str(challenger.agent_version_id)})
comparison.approved_at = timezone.now()
comparison.decided_by = actor
comparison.save(update_fields=["approved_at", "decided_by", "updated_at"])
return comparison
def reject_promotion(self, comparison: ExperimentComparison, *, actor: str = "human") -> ExperimentComparison:
comparison.rejected_at = timezone.now()
comparison.decided_by = actor
comparison.save(update_fields=["rejected_at", "decided_by", "updated_at"])
self.bus.publish("EXPERIMENT_PROMOTION_REJECTED", actor=actor, payload={"experiment_id": str(comparison.experiment_id)})
return comparison
def create_experiment_from_candidate(self, candidate: ImprovementCandidate, dataset: ReplayDataset, **kwargs: object) -> ProgenyExperiment:
return self.create_experiment(candidate, dataset, success_criteria=kwargs.get("success_criteria") if isinstance(kwargs.get("success_criteria"), dict) else None)
def _add_variant(self, experiment: ProgenyExperiment, role: str, *, graph_version: ExecutionGraphVersion | None, agent_version: AgentVersion | None) -> ExperimentVariant:
if graph_version is None and agent_version is None:
raise ValidationError("Variant requires an execution graph or agent version target.")
target_type = "EXECUTION_GRAPH" if graph_version else "AGENT"
target = graph_version or agent_version
snapshot = self._variant_snapshot(graph_version=graph_version, agent_version=agent_version)
return ExperimentVariant.objects.create(
experiment=experiment,
role=role,
target_type=target_type,
target_reference=str(target.id),
execution_graph_version=graph_version,
agent_version=agent_version,
configuration_snapshot=snapshot,
metadata={"single_variable_control": True},
)
def _variant_snapshot(self, *, graph_version: ExecutionGraphVersion | None, agent_version: AgentVersion | None) -> dict[str, object]:
if graph_version is not None:
return {"graph": graph_version.graph.name, "version": graph_version.version, "status": graph_version.status, "graph_spec": graph_version.graph_spec, "metadata": graph_version.metadata}
assert agent_version is not None
return {
"agent": agent_version.agent.name,
"role": agent_version.agent.role,
"version": agent_version.version,
"model": agent_version.model,
"system_contract": agent_version.system_contract,
"context_policy": agent_version.context_policy,
"tools": agent_version.tools,
"retry_policy": agent_version.retry_policy,
}
def _repository_head(self, repository_path: Path) -> str:
completed = subprocess.run(["git", "rev-parse", "HEAD"], cwd=repository_path, capture_output=True, text=True, check=True)
return completed.stdout.strip()
def _fresh_replay_repository(self, replay_case: ReplayCase, variant: ExperimentVariant) -> Path:
source = Path(replay_case.repository_path).resolve()
target = source.parent / f"{source.name}-replay-{replay_case.id}-{variant.role.lower()}"
if target.exists():
shutil.rmtree(target)
subprocess.run(["git", "clone", str(source), str(target)], check=True, capture_output=True, text=True)
subprocess.run(["git", "checkout", replay_case.repository_baseline_ref], cwd=target, check=True, capture_output=True, text=True)
return target
def _clone_task(self, replay_case: ReplayCase, replay_repo: Path, variant: ExperimentVariant) -> Task:
source_project = replay_case.source_project
project = Project.objects.create(
name=f"REPLAY {variant.role} {source_project.name if source_project else replay_case.id}",
project_type=replay_case.project_type or (source_project.project_type if source_project else "WEB_APP"),
goal=f"REPLAY / EXPERIMENTAL: {replay_case.original_goal}",
repository_path=str(replay_repo),
)
plan = ProjectPlan.objects.create(project=project, version=1, goal=project.goal)
milestone = Milestone.objects.create(project=project, plan=plan, key="REPLAY", title="Replay", goal="Replay experiment")
task = Task.objects.create(
project=project,
milestone=milestone,
task_type=replay_case.task_type,
status=TaskStatus.RUNNING,
goal=replay_case.original_goal,
acceptance_criteria=replay_case.acceptance_criteria,
max_retries=2,
)
Worktree.objects.create(
task=task,
repository_path=str(replay_repo),
worktree_path=str(replay_repo),
branch_name=f"replay/{variant.role.lower()}/{task.id}",
base_ref=replay_case.repository_baseline_ref,
)
return task
def _run_telemetry(self, task: Task, graph_run: GraphRun) -> dict[str, object]:
telemetry = dict(graph_run.metadata.get("telemetry", {})) if isinstance(graph_run.metadata, dict) else {}
telemetry["runtime_seconds"] = (graph_run.completed_at - graph_run.started_at).total_seconds() if graph_run.started_at and graph_run.completed_at else 0
telemetry["retry_count"] = task.retry_count
telemetry["graph_node_failures"] = graph_run.node_runs.filter(status="FAILED").count()
telemetry["model_requests_per_task"] = telemetry.get("model_requests", 0)
return telemetry
def _persist_result(self, run: ReplayRun, task: Task) -> ReplayResult:
test_run = TestRun.objects.filter(task=task).order_by("-created_at").first()
review = Review.objects.filter(task=task).order_by("-created_at").first()
verification = Verification.objects.filter(task=task).order_by("-created_at").first()
metrics = {
"completion": 1 if task.status == TaskStatus.COMPLETE else 0,
"test_pass": 1 if test_run and test_run.status == "PASS" else 0,
"review_pass": 1 if review and review.status == "PASS" else 0,
"review_rework": 1 if review and review.status == "REWORK_REQUIRED" else 0,
"review_reject": 1 if review and review.status == "REJECTED" else 0,
"judge_pass": 1 if verification and verification.result == VerificationResult.PASS else 0,
"accepted_candidate": 1 if task.status == TaskStatus.COMPLETE else 0,
"retry_count": task.retry_count,
"retry_exhausted": 1 if task.status == TaskStatus.FAILED else 0,
"model_output_invalid": task.progenysignal_set.filter(failure_category="MODEL_OUTPUT_INVALID").count() if hasattr(task, "progenysignal_set") else 0,
"mutation_failures": run.telemetry.get("mutation_failures", 0),
"patch_mismatch": run.telemetry.get("patch_mismatches", 0),
"runtime_seconds": run.telemetry.get("runtime_seconds", 0),
"model_requests": run.telemetry.get("model_requests_per_task", 0),
"mutation_operations": run.telemetry.get("mutation_operations", 0),
}
return ReplayResult.objects.create(
replay_run=run,
completion_status=task.status,
tests_status=test_run.status if test_run else "",
reviewer_status=review.status if review else "",
judge_status=verification.result if verification else "",
accepted_candidate=task.status == TaskStatus.COMPLETE,
metrics=metrics,
safety={"unexpected_file_scope_changes": 0, "policy_violations": 0, "duplicate_side_effect_attempts": 0},
evidence={"commit_candidate_sha": run.commit_candidate_sha, "failure_classification": run.failure_classification},
)
def _eligible_results(self, experiment: ProgenyExperiment, variant: ExperimentVariant) -> list[ReplayResult]:
return list(
ReplayResult.objects.filter(replay_run__experiment=experiment, replay_run__variant=variant)
.exclude(replay_run__failure_classification__in=INFRA_FAILURES)
.select_related("replay_run", "replay_run__replay_case")
)
def _aggregate(self, results: list[ReplayResult]) -> dict[str, float]:
count = len(results)
if count == 0:
return {"case_count": 0}
keys = ["completion", "test_pass", "review_pass", "review_rework", "review_reject", "judge_pass", "accepted_candidate", "retry_exhausted", "model_output_invalid", "mutation_failures", "patch_mismatch", "model_requests", "mutation_operations"]
aggregate = {"case_count": float(count)}
for key in keys:
total = sum(float(result.metrics.get(key, 0)) for result in results)
aggregate[f"{key}_rate" if key in {"completion", "test_pass", "review_pass", "review_rework", "review_reject", "judge_pass", "accepted_candidate", "retry_exhausted", "model_output_invalid"} else f"{key}_per_case"] = total / count
aggregate["median_runtime_seconds"] = median([float(result.metrics.get("runtime_seconds", 0)) for result in results])
return aggregate
def _delta(self, champion: dict[str, float], challenger: dict[str, float]) -> dict[str, float]:
return {key: challenger.get(key, 0) - champion.get(key, 0) for key in set(champion) | set(challenger) if key != "case_count"}
def _paired_outcomes(self, experiment: ProgenyExperiment, champion: ExperimentVariant, challenger: ExperimentVariant) -> dict[str, object]:
outcomes = {"BOTH_PASS": [], "BOTH_FAIL": [], "CHAMPION_ONLY_PASS": [], "CHALLENGER_ONLY_PASS": []}
for replay_case in experiment.replay_dataset.cases.all():
champion_result = ReplayResult.objects.filter(replay_run__experiment=experiment, replay_run__variant=champion, replay_run__replay_case=replay_case).first()
challenger_result = ReplayResult.objects.filter(replay_run__experiment=experiment, replay_run__variant=challenger, replay_run__replay_case=replay_case).first()
if not champion_result or not challenger_result:
continue
if champion_result.replay_run.failure_classification in INFRA_FAILURES or challenger_result.replay_run.failure_classification in INFRA_FAILURES:
continue
champion_pass = champion_result.accepted_candidate
challenger_pass = challenger_result.accepted_candidate
key = "BOTH_PASS" if champion_pass and challenger_pass else "BOTH_FAIL" if not champion_pass and not challenger_pass else "CHAMPION_ONLY_PASS" if champion_pass else "CHALLENGER_ONLY_PASS"
outcomes[key].append(str(replay_case.id))
return {**outcomes, "case_count": sum(len(value) for value in outcomes.values()), "regression_cases": outcomes["CHAMPION_ONLY_PASS"]}
def _classify_replay_exception(self, exc: Exception) -> str:
text = str(exc).lower()
if "provider" in text or "qwen" in text:
return "PROVIDER_FAILURE"
if "git" in text or "worktree" in text or "repository" in text:
return "REPLAY_RUNTIME_FAILURE"
return "INFRASTRUCTURE_FAILURE"
class ReplayArenaNode:
idempotent = True
replay_safe = True
destructive = False
def __init__(self, arena: ReplayArena, node_type: str) -> None:
self.arena = arena
self.node_type = node_type
def experiment(self, context: GraphExecutionContext) -> ProgenyExperiment:
return ProgenyExperiment.objects.get(id=context.graph_run.metadata["experiment_id"])
class ReplayNoopNode(ReplayArenaNode):
def run(self, context: GraphExecutionContext) -> NodeResult:
return NodeResult("COMPLETE", "success")
class ReplayRunChampionNode(ReplayArenaNode):
def run(self, context: GraphExecutionContext) -> NodeResult:
experiment = self.experiment(context)
champion = experiment.variants.get(role="CHAMPION")
for replay_case in experiment.replay_dataset.cases.order_by("created_at"):
if not ReplayRun.objects.filter(experiment=experiment, replay_case=replay_case, variant=champion).exists():
self.arena.run_case(experiment, replay_case, champion)
return NodeResult("COMPLETE", "success")
class ReplayRunChallengerNode(ReplayArenaNode):
def run(self, context: GraphExecutionContext) -> NodeResult:
experiment = self.experiment(context)
challenger = experiment.variants.filter(role="CHALLENGER").order_by("created_at").first()
if challenger is None:
return NodeResult("FAILED", "failure", failure_evidence={"reason": "missing challenger variant"})
for replay_case in experiment.replay_dataset.cases.order_by("created_at"):
if not ReplayRun.objects.filter(experiment=experiment, replay_case=replay_case, variant=challenger).exists():
self.arena.run_case(experiment, replay_case, challenger)
return NodeResult("COMPLETE", "success")
class ReplayCompareNode(ReplayArenaNode):
def run(self, context: GraphExecutionContext) -> NodeResult:
comparison = self.arena.compare(self.experiment(context))
return NodeResult("COMPLETE", "success", {"comparison_id": str(comparison.id), "verdict": comparison.verdict})
class ReplayHumanDecisionNode(ReplayArenaNode):
def run(self, context: GraphExecutionContext) -> NodeResult:
node_run = context.graph_run.node_runs.filter(node_id=context.graph_run.current_node).order_by("-visit_index").first()
if GraphApproval.objects.filter(graph_run=context.graph_run, status=GraphApprovalStatus.APPROVED).exists():
return NodeResult("COMPLETE", "approved")
if GraphApproval.objects.filter(graph_run=context.graph_run, status=GraphApprovalStatus.REJECTED).exists():
return NodeResult("COMPLETE", "rejected")
GraphApproval.objects.get_or_create(graph_run=context.graph_run, node_run=node_run, reason="AWAITING_REPLAY_PROMOTION_DECISION")
return NodeResult("PAUSED", "awaiting", pause_reason="AWAITING_REPLAY_PROMOTION_DECISION")
def replay_experiment_registry(arena: ReplayArena) -> NodeHandlerRegistry:
registry = NodeHandlerRegistry()
for node_type in ["replay_prepare", "replay_select_cases", "replay_validate_variants", "replay_experiment_judge"]:
registry.register(ReplayNoopNode(arena, node_type))
registry.register(ReplayRunChampionNode(arena, "replay_run_champion"))
registry.register(ReplayRunChallengerNode(arena, "replay_run_challenger"))
registry.register(ReplayCompareNode(arena, "replay_compare"))
registry.register(ReplayHumanDecisionNode(arena, "replay_human_decision"))
return registry