Files
2026-09-19 12:54:45 +08:00

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),
}