first commit
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
"""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),
|
||||
}
|
||||
Reference in New Issue
Block a user