first commit

This commit is contained in:
eeymoo
2026-09-19 12:54:45 +08:00
commit 6fc5b64077
126 changed files with 8601 additions and 0 deletions
+112
View File
@@ -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),
}