"""Async experiment runner with retry, isolation, and resume support.""" import asyncio from datetime import datetime, timezone from typing import Any, Dict, List, Optional from sqlalchemy.orm import Session from app.dataset.diff_extractor import extract_diff_bundle from app.model_adapters.factory import create_adapter from app.models import ExperimentRun, Result from app.prompts.service import PromptService class ExperimentRunner: def __init__( self, db: Session, concurrency: int = 5, max_retries: int = 3, strategy_id: str = "code_review", ): self.db = db self.concurrency = concurrency self.max_retries = max_retries self.strategy_id = strategy_id self.prompt_service = PromptService(db) def _get_pending_runs(self, experiment_id: Optional[str] = None) -> List[ExperimentRun]: query = self.db.query(ExperimentRun) if experiment_id: query = query.filter_by(experiment_id=experiment_id) return query.filter(ExperimentRun.status.in_(["pending", "failed"])).all() def _reset_stale_running(self) -> None: """Mark runs stuck in running without result back to pending.""" stale = ( self.db.query(ExperimentRun) .filter_by(status="running") .filter(~ExperimentRun.result.has()) .all() ) for run in stale: run.status = "pending" self.db.commit() async def run_experiment( self, experiment_id: Optional[str] = None, progress_callback=None, ) -> Dict[str, Any]: self._reset_stale_running() pending = self._get_pending_runs(experiment_id) semaphore = asyncio.Semaphore(self.concurrency) async def execute(run: ExperimentRun): async with semaphore: return await self._execute_run(run) tasks = [asyncio.create_task(execute(run)) for run in pending] results = [] for coro in asyncio.as_completed(tasks): result = await coro results.append(result) if progress_callback: progress_callback(result) return {"total": len(pending), "completed": len(results)} async def _execute_run(self, run: ExperimentRun) -> Dict[str, Any]: run.status = "running" run.started_at = datetime.now(timezone.utc) self.db.commit() try: sample = run.sample prompt = self.prompt_service.render( self.strategy_id, run.template_version.version_number, {"language": sample.language, "diff": sample.diff}, ) adapter = create_adapter(run.model_id) response = await adapter.chat(prompt, run.sampling_params) run.status = "done" run.completed_at = datetime.now(timezone.utc) result = Result( run_id=run.id, raw_output=response.text, token_usage=response.token_usage, latency_ms=response.latency_ms, ) self.db.add(result) self.db.commit() return { "run_id": str(run.id), "status": "done", "model_id": run.model_id, } except Exception as e: run.retry_count += 1 if run.retry_count > self.max_retries: run.status = "failed" else: run.status = "pending" # will be retried on next resume run.completed_at = datetime.now(timezone.utc) self.db.commit() return { "run_id": str(run.id), "status": run.status, "error": str(e), }