"""Start, inspect, and cancel persisted batch text test runs.""" from __future__ import annotations from datetime import UTC, datetime from uuid import uuid4 from db.models import ( Assistant, ModelResource, TestCase, TestRun, TestRunCase, TestSuite, ) from db.session import get_session from fastapi import APIRouter, Depends, HTTPException from services.auth import require_admin from services.config_resolver import resolve_runtime_config from services.test_runs.orchestrator import test_run_orchestrator from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from test_schemas import ( BatchExecutionError, BatchRunCaseOut, BatchRunCreate, BatchRunSnapshotOut, BatchTurnResult, BatchEvaluationResult, TestCaseDefinition, ) router = APIRouter( prefix="/api/test-runs", tags=["tests"], dependencies=[Depends(require_admin)], ) def _new_id(prefix: str) -> str: return f"{prefix}_{uuid4().hex}" async def _run_out( session: AsyncSession, run: TestRun, ) -> BatchRunSnapshotOut: rows = ( await session.execute( select(TestRunCase) .where(TestRunCase.run_id == run.id) .order_by(TestRunCase.position) ) ).scalars().all() cases: list[BatchRunCaseOut] = [] for row in rows: result = dict(row.result or {}) cases.append( BatchRunCaseOut( id=str(row.test_case_id or row.case_snapshot.get("id") or row.id), name=row.test_case_name, status=row.status, turns=[ BatchTurnResult.model_validate(item) for item in result.get("turns") or [] ], overall_criteria=[ BatchEvaluationResult.model_validate(item) for item in result.get("overallCriteria") or result.get("overall_criteria") or [] ], attempt_count=row.attempt_count, max_attempts=row.max_attempts, execution_error=( BatchExecutionError.model_validate(row.execution_error) if row.execution_error else None ), ) ) return BatchRunSnapshotOut( id=run.id, status=run.status, title=run.title, assistant_name=run.assistant_name, config=dict(run.config or {}), cases=cases, started_at=run.started_at or run.created_at, finished_at=run.finished_at, stop_reason=run.stop_reason, ) @router.post("", response_model=BatchRunSnapshotOut) async def create_test_run( body: BatchRunCreate, session: AsyncSession = Depends(get_session), ): assistant = await session.get(Assistant, body.assistant_id) if assistant is None: raise HTTPException(404, "被测助手不存在") if assistant.runtime_mode != "pipeline": raise HTTPException(422, "第一版批量测试只支持 Pipeline 运行模式") if assistant.type not in {"prompt", "workflow"}: raise HTTPException(422, f"第一版批量测试暂不支持 {assistant.type} 类型助手") evaluator_resource = await session.get( ModelResource, body.evaluator_model_resource_id, ) if ( evaluator_resource is None or evaluator_resource.capability != "LLM" or not evaluator_resource.enabled ): raise HTTPException(422, "评估模型不存在、已停用或不是 LLM 资源") try: runtime_config = await resolve_runtime_config(session, assistant.id) except ValueError as exc: raise HTTPException(422, str(exc)) from exc rows = ( await session.execute(select(TestCase).where(TestCase.id.in_(body.case_ids))) ).scalars().all() by_id = {row.id: row for row in rows} missing = [case_id for case_id in body.case_ids if case_id not in by_id] if missing: raise HTTPException(422, "测试用例不存在: " + "、".join(missing)) ordered_cases = [by_id[case_id] for case_id in body.case_ids] tools_by_id = {tool.id: tool for tool in runtime_config.tools} for row in ordered_cases: if row.input_mode != "fixed_script_text": raise HTTPException(422, f"用例“{row.name}”不是可运行的固定文字模式") try: definition = TestCaseDefinition.model_validate(row.definition or {}) except ValueError as exc: raise HTTPException(422, f"用例“{row.name}”校验失败: {exc}") from exc for turn in definition.turns: for behavior in turn.behaviors: if behavior.type != "tool_call": continue tool = tools_by_id.get(behavior.tool_id) if tool is None or tool.function_name != behavior.function_name: raise HTTPException( 422, f"用例“{row.name}”引用的工具 {behavior.function_name} " "不属于当前被测助手", ) suite_ids = list(dict.fromkeys(row.suite_id for row in ordered_cases)) suites = ( await session.execute(select(TestSuite).where(TestSuite.id.in_(suite_ids))) ).scalars().all() suite_by_id = {row.id: row for row in suites} if body.title: title = body.title.strip() elif len(suite_ids) == 1 and suite_ids[0] in suite_by_id: title = f"{suite_by_id[suite_ids[0]].name} · {len(ordered_cases)} 个用例" else: title = f"批量测试 · {len(ordered_cases)} 个用例" now = datetime.now(UTC) run = TestRun( id=_new_id("run"), assistant_id=assistant.id, assistant_name=assistant.name, title=title, status="queued", config={ "suiteCount": len(suite_ids), "evaluatorModelResourceId": evaluator_resource.id, "evaluatorModelResourceName": evaluator_resource.name, "evaluatorModel": str( (evaluator_resource.values or {}).get("modelId") or "" ), **body.config.model_dump(mode="json", by_alias=True), }, started_at=now, ) session.add(run) for position, row in enumerate(ordered_cases): definition = TestCaseDefinition.model_validate(row.definition or {}) snapshot = { "id": row.id, "suiteId": row.suite_id, "name": row.name, "description": row.description, "inputMode": row.input_mode, **definition.model_dump(mode="json", by_alias=True), } session.add( TestRunCase( id=_new_id("rc"), run_id=run.id, test_case_id=row.id, test_case_name=row.name, suite_id=row.suite_id, position=position, case_snapshot=snapshot, status="waiting", attempt_count=0, max_attempts=body.config.error_retry_count + 1, result={"turns": [], "overallCriteria": []}, ) ) await session.commit() await session.refresh(run) test_run_orchestrator.start(run.id) return await _run_out(session, run) @router.get("/{run_id}", response_model=BatchRunSnapshotOut) async def get_test_run( run_id: str, session: AsyncSession = Depends(get_session), ): run = await session.get(TestRun, run_id) if run is None: raise HTTPException(404, "批量测试运行不存在") return await _run_out(session, run) @router.post("/{run_id}/cancel", response_model=BatchRunSnapshotOut) async def cancel_test_run( run_id: str, session: AsyncSession = Depends(get_session), ): run = await session.get(TestRun, run_id) if run is None: raise HTTPException(404, "批量测试运行不存在") if run.status in {"completed", "cancelled"}: return await _run_out(session, run) run.cancel_requested = True run.stop_reason = "manual" await session.commit() await test_run_orchestrator.cancel(run_id) await session.refresh(run) return await _run_out(session, run)