"""API routes for evaluation runs.""" import asyncio from typing import Optional from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel from sqlmodel import Session from agenteval.evaluation.engine import EvalEngine from agenteval.models import EvalRun, RunStatus from agenteval.storage.db import get_session from agenteval.storage.repository import RunRepository, ScenarioRepository, TargetRepository from agenteval.utils.llm import extract_reply_text from agenteval.web.deps import get_db from agenteval.web.websocket import ws_manager router = APIRouter() class StartRunRequest(BaseModel): target_id: str scenario_id: str # ── Task registry for live evaluation runs ───────────────────────────── # Each running evaluation is an asyncio.Task keyed by run_id. The cancel # token is a cooperative ``asyncio.Event`` the engine checks between cases. _tasks: dict[str, asyncio.Task] = {} _cancel_tokens: dict[str, asyncio.Event] = {} async def _run_evaluation(run_id: str, target_id: str, scenario_id: str) -> None: """Background coroutine that drives one evaluation run to completion.""" session = get_session() cancel_token = asyncio.Event() _cancel_tokens[run_id] = cancel_token try: target = TargetRepository(session).get(target_id) scenario = ScenarioRepository(session).get(scenario_id) existing_run = RunRepository(session).get(run_id) if not target or not scenario: return engine = EvalEngine( target=target, scenario=scenario, session=session, cancel_token=cancel_token, ) await engine.run( progress_callback=lambda event, data: ws_manager.emit(run_id, event, data), existing_run=existing_run, ) finally: session.close() _cancel_tokens.pop(run_id, None) _tasks.pop(run_id, None) @router.get("") async def list_runs(session: Session = Depends(get_db)) -> list[dict]: return [r.model_dump() for r in RunRepository(session).list_all()] @router.post("") async def start_run( request: StartRunRequest, session: Session = Depends(get_db), ) -> dict: target = TargetRepository(session).get(request.target_id) scenario = ScenarioRepository(session).get(request.scenario_id) if not target or not scenario: raise HTTPException(status_code=404, detail="target or scenario not found") run = EvalRun(target_id=request.target_id, scenario_id=request.scenario_id) run = RunRepository(session).create(run) task = asyncio.create_task( _run_evaluation(run.id, request.target_id, request.scenario_id), name=f"eval-run-{run.id}", ) _tasks[run.id] = task return run.model_dump() @router.get("/{run_id}") async def get_run(run_id: str, session: Session = Depends(get_db)) -> dict: run = RunRepository(session).get(run_id) if not run: raise HTTPException(status_code=404, detail="run not found") return run.model_dump() @router.post("/{run_id}/cancel") async def cancel_run(run_id: str, session: Session = Depends(get_db)) -> dict: repo = RunRepository(session) run = repo.get(run_id) if not run: raise HTTPException(status_code=404, detail="run not found") if run.status not in (RunStatus.PENDING, RunStatus.RUNNING): raise HTTPException(status_code=400, detail="run is not in a cancellable state") cancel_token = _cancel_tokens.get(run_id) task: Optional[asyncio.Task] = _tasks.get(run_id) if cancel_token is not None: # Cooperative cancel: the engine will catch CancelledError and mark # the run as FAILED with code=cancelled_by_user. cancel_token.set() elif task is not None: # Fallback: hard-cancel the task if no token exists (shouldn't happen). task.cancel() else: # No live task (e.g. process restarted): mark the DB row directly. run.status = RunStatus.FAILED run.summary = { "error": {"code": "cancelled_by_user", "message": "评测已手动停止"}, } repo.update(run) return run.model_dump() @router.get("/{run_id}/logs") async def get_run_logs(run_id: str, session: Session = Depends(get_db)) -> dict: repo = RunRepository(session) run = repo.get(run_id) if not run: raise HTTPException(status_code=404, detail="run not found") turns = repo.get_turns(run_id) results = repo.get_results(run_id) turns_data = [ { "id": t.id, "case_id": t.case_id, "round_index": t.round_index, "latency_ms": t.latency_ms, "sent_text": t.get_sent_message().get("msgBody", {}).get("content", ""), "reply_text": extract_reply_text(t.get_reply()), "sent_at": t.sent_at.isoformat() if t.sent_at else None, "received_at": t.received_at.isoformat() if t.received_at else None, } for t in turns ] results_data = [ { "case_id": r.case_id, "rule_type": r.rule_type, "passed": r.passed, "score": r.score, "reason": r.reason, } for r in results ] scenario_snapshot: dict = {} scenario = ScenarioRepository(session).get(run.scenario_id) if scenario: for case in scenario.cases: scenario_snapshot[case.id] = { "id": case.id, "type": case.type.value if hasattr(case.type, "value") else str(case.type), "messages": list(case.messages), "prompt": case.prompt, "turns": case.turns, "expectations": { "intent": case.expectations.intent, "keywords_include": list(case.expectations.keywords_include), "keywords_exclude": list(case.expectations.keywords_exclude), "response_time_max_ms": case.expectations.response_time_max_ms, "coherence_min_score": case.expectations.coherence_min_score, }, "eval_rules": [{"type": r.type, "params": dict(r.params)} for r in case.eval_rules], } return {"turns": turns_data, "results": results_data, "scenario_snapshot": scenario_snapshot}