113 lines
3.7 KiB
Python
113 lines
3.7 KiB
Python
"""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),
|
|
}
|