"""API routes for evaluation runs.""" import asyncio from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel from sqlmodel import Session from agenteval.evaluation.case_verdict import CaseEvidence, resolve_case_verdicts from agenteval.evaluation.engine import EvalEngine from agenteval.models import EvalRun, RunStatus, RunSummary, RunTrigger from agenteval.storage.db import get_session, iso_utc from agenteval.storage.repository import RunRepository, ScenarioRepository, TargetRepository from agenteval.task_registry import TaskRegistry from agenteval.utils.llm import extract_reply_text from agenteval.utils.webhook import send_run_webhook 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 triggered_by: RunTrigger = RunTrigger.MANUAL # ── 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. run_registry = TaskRegistry() async def _run_evaluation( run_id: str, target_id: str, scenario_id: str, *, cancel_token: asyncio.Event ) -> None: """Background coroutine that drives one evaluation run to completion.""" session = get_session() 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, ) # Fire webhook after run completes (non-blocking, best-effort) completed_run = RunRepository(session).get(run_id) if completed_run: await send_run_webhook( run_id=run_id, status=completed_run.status.value, summary=completed_run.summary.model_dump(mode="json") if completed_run.summary else {}, ) finally: session.close() @router.get("") async def list_runs(session: Session = Depends(get_db)) -> list[dict]: scenario_names = ScenarioRepository(session).name_map() target_names = {t.id: t.name for t in TargetRepository(session).list_all()} return [ { **r.model_dump(), "scenario_name": scenario_names.get(r.scenario_id), "target_name": target_names.get(r.target_id), } 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, scenario_version=scenario.version or 1, triggered_by=request.triggered_by, ) run = RunRepository(session).create(run) run_registry.launch( run.id, lambda cancel_token: _run_evaluation( run.id, request.target_id, request.scenario_id, cancel_token=cancel_token ), ) 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") signalled = run_registry.cancel(run_id) if not signalled: # 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": iso_utc(t.sent_at), "received_at": iso_utc(t.received_at), } 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 ] # Authoritative per-case verdicts: resolve_case_verdicts prefers the engine's # stored case_outcomes and approximates only for legacy runs (single seam). evidence: dict[str, dict] = {} for t in turns: ev = evidence.get(t.case_id) if ev is None: ev = {"has_turns": True, "all_replied": True, "passes": []} evidence[t.case_id] = ev else: ev["has_turns"] = True if t.get_reply() is None: ev["all_replied"] = False for r in results: ev = evidence.setdefault(r.case_id, {"has_turns": False, "all_replied": True, "passes": []}) ev["passes"].append(r.passed) summary = run.summary or RunSummary() errored_case_ids = {e.get("case_id") for e in summary.case_errors} verdicts = resolve_case_verdicts( case_outcomes=summary.case_outcomes, evidence={ cid: CaseEvidence( has_turns=ev["has_turns"], all_replied=ev["all_replied"], result_passes=tuple(ev["passes"]), ) for cid, ev in evidence.items() }, errored_case_ids=errored_case_ids, ) case_verdicts = {cid: {"passed": v.passed, "connectivity": v.connectivity} for cid, v in verdicts.items()} 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), "weight": r.weight} for r in case.eval_rules], "rule_logic": case.rule_logic.value if hasattr(case.rule_logic, "value") else str(case.rule_logic), "rule_pass_threshold": case.rule_pass_threshold, } return { "turns": turns_data, "results": results_data, "case_verdicts": case_verdicts, "scenario_snapshot": scenario_snapshot, }