"""Evaluation execution engine. Async-first implementation: channels and LLM calls are awaited cooperatively, so multiple cases can run concurrently and a run can be cancelled mid-flight via an ``asyncio.Event`` cancel token. """ import asyncio import uuid from dataclasses import dataclass from typing import Any, Callable, Optional from agenteval.channels.base import EvalChannel from agenteval.channels.factory import ChannelFactory from agenteval.evaluation.rules import RuleResult, get_rule from agenteval.model_gateway import ModelGateway from agenteval.models import ( Case, CaseType, EvalResult, EvalRun, EvalTarget, ModelCapability, ModelPurpose, RuleLogic, RunStatus, Scenario, Turn, ) from agenteval.services.model_configs import ModelConfigService, ModelRuntimeConfig from agenteval.storage.db import get_session, utc_now from agenteval.storage.repository import ResultRepository, RunRepository from agenteval.utils.llm import extract_reply_text, parse_json_from_llm_text # Progress callbacks may be sync or async; the engine awaits the result if # it is a coroutine, otherwise treats it as a plain function. ProgressCallback = Callable[[str, dict[str, Any]], Any] class CancelledError(RuntimeError): """Raised inside the engine when the cancel token fires.""" @dataclass class TimeoutConfig: """Per-operation timeouts (all in seconds).""" poll_reply: float = 30.0 llm_generate: float = 60.0 def _build_send_message(content: str) -> dict[str, Any]: return { "msgType": "text", "msgBody": {"content": content}, } class EvalEngine: """Execute evaluation scenarios against targets. The engine is async so the (slow, network-bound) channel and LLM calls can be awaited cooperatively. Database writes remain synchronous for now (SQLite + StaticPool); they are fast enough not to block the event loop in practice. """ def __init__( self, target: EvalTarget, scenario: Scenario, session=None, run_repo: Optional[RunRepository] = None, result_repo: Optional[ResultRepository] = None, cancel_token: Optional[asyncio.Event] = None, timeout_config: Optional[TimeoutConfig] = None, max_concurrent_cases: int = 1, ): self.target = target self.scenario = scenario self.channel: EvalChannel = ChannelFactory.create(target) self.session = session or get_session() self.run_repo = run_repo or RunRepository(self.session) self.result_repo = result_repo or ResultRepository(self.session) self.cancel_token = cancel_token or asyncio.Event() self.timeout_config = timeout_config or TimeoutConfig() self._case_semaphore = asyncio.Semaphore(max(1, max_concurrent_cases)) # Collects fatal case-level errors (e.g. dynamic message generation # failures) so their cause is persisted into run.summary — not just # emitted transiently over WebSocket. self._case_errors: list[dict[str, str]] = [] self.model_gateway = ModelGateway(timeout=self.timeout_config.llm_generate) self.model_service = ModelConfigService(self.session) self._resolved_models: dict[ModelPurpose, ModelRuntimeConfig] = {} # ── public entry point ──────────────────────────────────────────── async def run( self, progress_callback: Optional[ProgressCallback] = None, existing_run: Optional[EvalRun] = None, ) -> EvalRun: """Run the evaluation and return the completed run record. Cancellation is cooperative: set ``cancel_token`` and the engine will mark the run as FAILED with ``cancelled_by_user`` at the next checkpoint. """ if existing_run: run = existing_run run.status = RunStatus.RUNNING run.started_at = utc_now() run = self.run_repo.update(run) or run else: run = EvalRun( id=str(uuid.uuid4()), target_id=self.target.id or "", scenario_id=self.scenario.id or "", status=RunStatus.RUNNING, started_at=utc_now(), ) run = self.run_repo.create(run) try: total_cases = len(self.scenario.cases) passed_cases = 0 failed_cases = 0 for idx, case in enumerate(self.scenario.cases, start=1): self._check_cancel() await self._emit( progress_callback, "case_start", { "index": idx, "total": total_cases, "case_id": case.id, }, ) async with self._case_semaphore: case_passed, rule_pass, rule_total = await self._run_case( run, case, progress_callback, ) if case_passed: passed_cases += 1 else: failed_cases += 1 await self._emit( progress_callback, "case_end", { "index": idx, "total": total_cases, "case_id": case.id, "passed": case_passed, "rule_pass_count": rule_pass, "rule_total": rule_total, }, ) results = self.run_repo.get_results(run.id) total_rules = len(results) passed_rules = sum(1 for r in results if r.passed) summary = { "total_cases": total_cases, "passed_cases": passed_cases, "failed_cases": failed_cases, "total_rules": total_rules, "passed_rules": passed_rules, "pass_rate": round(passed_rules / total_rules, 4) if total_rules else 0.0, } # Surface fatal case-level errors (e.g. dynamic generation failures) # so the report / DB record shows *why* a run produced no results. if self._case_errors: summary["case_errors"] = self._case_errors if self._resolved_models: summary["model_configs"] = { purpose.value: config.snapshot() for purpose, config in self._resolved_models.items() } run.status = RunStatus.COMPLETED run.completed_at = utc_now() run.summary = summary await self._emit( progress_callback, "run_completed", { "status": "completed", "summary": summary, }, ) except CancelledError: run.status = RunStatus.FAILED run.completed_at = utc_now() run.summary = { "error": {"code": "cancelled_by_user", "message": "评测已手动停止"}, } await self._emit( progress_callback, "run_completed", { "status": "failed", "reason": "cancelled", "error": {"code": "cancelled_by_user", "message": "评测已手动停止"}, }, ) except Exception as exc: run.status = RunStatus.FAILED run.completed_at = utc_now() run.summary = {"error": str(exc)} await self._emit(progress_callback, "error", {"error": str(exc)}) await self._emit( progress_callback, "run_completed", { "status": "failed", "error": str(exc), }, ) raise finally: run = self.run_repo.update(run) or run # Best-effort cleanup of the channel's HTTP client. close = getattr(self.channel, "close", None) if callable(close): try: result = close() if asyncio.iscoroutine(result): await result except Exception: pass try: self.session.close() except Exception: pass return run # ── case / turn execution ───────────────────────────────────────── async def _run_case( self, run: EvalRun, case: Case, progress_callback: Optional[ProgressCallback], ) -> tuple[bool, int, int]: """Run a single case; returns (all_rules_passed, passed_rules, total_rules).""" if case.type == CaseType.DYNAMIC: generated = await self._generate_messages(case, progress_callback) if not generated: await self._emit( progress_callback, "error", { "error": "LLM 未能生成测试消息", "case_id": case.id, }, ) return False, 0, 0 case = case.model_copy(update={"messages": generated}) dialog: list[Turn] = [] for round_index, message in enumerate(case.messages, start=1): self._check_cancel() await self._emit( progress_callback, "turn_start", { "run_id": run.id, "case_id": case.id, "round": round_index, "message": message, }, ) sent_at = utc_now() send_result = await self.channel.send(message) if not send_result.ok: turn = Turn( id=str(uuid.uuid4()), run_id=run.id, case_id=case.id, round_index=round_index, sent_message=_build_send_message(message), sent_at=sent_at, ) self.result_repo.save_turn(turn) await self._save_rule_results(run, case, turn, [], progress_callback) await self._emit( progress_callback, "turn_error", { "case_id": case.id, "round": round_index, "error": send_result.error, }, ) return False, 0, 0 try: reply = await self.channel.poll_reply( send_result.question_msg_id or "", timeout=self.timeout_config.poll_reply, ) except Exception as poll_exc: received_at = utc_now() turn = Turn( id=str(uuid.uuid4()), run_id=run.id, case_id=case.id, round_index=round_index, sent_message=_build_send_message(message), sent_at=sent_at, question_msg_id=send_result.question_msg_id, received_at=received_at, ) self.result_repo.save_turn(turn) await self._emit( progress_callback, "turn_error", { "case_id": case.id, "round": round_index, "error": f"poll_reply 异常: {poll_exc}", }, ) return False, 0, 0 received_at = utc_now() latency_ms = None if sent_at and received_at: latency_ms = int((received_at - sent_at).total_seconds() * 1000) turn = Turn( id=str(uuid.uuid4()), run_id=run.id, case_id=case.id, round_index=round_index, sent_message=_build_send_message(message), sent_at=sent_at, question_msg_id=send_result.question_msg_id, reply=reply.raw_message if reply else None, received_at=received_at, latency_ms=latency_ms, ) self.result_repo.save_turn(turn) dialog.append(turn) await self._emit( progress_callback, "turn_end", { "run_id": run.id, "case_id": case.id, "round": round_index, "latency_ms": latency_ms, "reply_text": extract_reply_text(reply.raw_message if reply else None), }, ) if not dialog: return False, 0, 0 return await self._save_rule_results(run, case, dialog[-1], dialog, progress_callback) async def _save_rule_results( self, run: EvalRun, case: Case, turn: Turn, dialog: list[Turn], progress_callback: Optional[ProgressCallback], ) -> tuple[bool, int, int]: """Apply rules and save results; returns (case_passed, passed_count, total_count). Combination logic (case.rule_logic): ALL — all rules must pass (default) ANY — at least one rule must pass WEIGHTED — weighted average score >= case.rule_pass_threshold """ from agenteval.models import EvalRuleConfig rules_config: list[EvalRuleConfig] = list(case.eval_rules) # If no explicit rules, derive implicit rules from expectations. if not rules_config: if case.expectations.response_time_max_ms: rules_config.append( EvalRuleConfig( type="response_time", params={ "max_ms": case.expectations.response_time_max_ms, }, ) ) if case.expectations.keywords_include or case.expectations.keywords_exclude: rules_config.append( EvalRuleConfig( type="keyword_match", params={ "keywords": case.expectations.keywords_include, "exclude_keywords": case.expectations.keywords_exclude, }, ) ) if not rules_config: # No rules defined and no expectations → case passes with no checks return True, 0, 0 passed_count = 0 total_count = 0 weighted_score = 0.0 total_weight = 0.0 for rule_config in rules_config: purpose = { "llm_score": ModelPurpose.JUDGE, "semantic_similarity": ModelPurpose.EMBEDDING, "safety": ModelPurpose.MODERATION, }.get(rule_config.type) try: model_config = self._resolve_model(purpose) if purpose else None rule = get_rule( rule_config.type, rule_config.params, model_config=model_config, gateway=self.model_gateway if model_config else None, ) result = await rule.evaluate(case, dialog) except Exception as exc: result = RuleResult(passed=False, reason=f"模型配置解析失败: {exc}") eval_result = EvalResult( id=str(uuid.uuid4()), run_id=run.id, case_id=case.id, turn_id=turn.id or "", rule_type=rule_config.type, passed=result.passed, score=result.score, reason=result.reason, ) self.result_repo.save_result(eval_result) total_count += 1 if result.passed: passed_count += 1 # Weighted scoring: use rule score (default 1.0 if passed, 0.0 if failed) score_val = result.score if result.score is not None else (1.0 if result.passed else 0.0) weight = rule_config.weight weighted_score += score_val * weight total_weight += weight await self._emit( progress_callback, "rule_result", { "run_id": run.id, "case_id": case.id, "rule_type": rule_config.type, "passed": result.passed, "score": result.score, "reason": result.reason, "weight": weight, }, ) # Determine case pass/fail based on rule_logic logic = case.rule_logic if logic == RuleLogic.ALL: case_passed = passed_count == total_count elif logic == RuleLogic.ANY: case_passed = passed_count > 0 elif logic == RuleLogic.WEIGHTED: avg = weighted_score / total_weight if total_weight > 0 else 0.0 case_passed = avg >= case.rule_pass_threshold else: case_passed = passed_count == total_count return case_passed, passed_count, total_count async def _generate_messages( self, case: Case, progress_callback: Optional[ProgressCallback], ) -> list[str]: """Use LLM to generate test messages for dynamic cases.""" async def _fail(msg: str) -> list[str]: # Persist the reason into run-level case_errors (surfaced in summary), # not just a transient WebSocket emit that's lost after the run. self._case_errors.append({"case_id": case.id, "stage": "generate_messages", "error": msg}) await self._emit(progress_callback, "error", {"error": msg, "case_id": case.id}) return [] turns = case.turns or 3 prompt = case.prompt or "请生成一些测试问题" system_prompt = ( f"你需要扮演一个真实的用户/患者,根据以下要求生成 {turns} 条独立的测试问题。\n\n" f"要求:{prompt}\n\n" "输出格式要求:只输出一个 JSON 数组,包含 " + str(turns) + " 个字符串,每个字符串是一条消息。" "不要输出任何解释、markdown 或其他内容。" ) messages_payload = [ {"role": "system", "content": system_prompt}, {"role": "user", "content": f"请生成 {turns} 条测试消息"}, ] try: model_config = self._resolve_model(ModelPurpose.GENERATOR) if model_config: content = await self.model_gateway.chat(model_config, messages_payload, temperature=0.7) else: content = await self._generate_messages_legacy(messages_payload) try: parsed = parse_json_from_llm_text(content) except (ValueError, Exception) as parse_exc: return await _fail(f"LLM 返回无法解析为数组: {parse_exc}") if not isinstance(parsed, list): return await _fail("LLM 返回的不是数组") messages = [str(m) for m in parsed if isinstance(m, str) and m.strip()] if not messages: return await _fail("LLM 返回的消息为空") await self._emit( progress_callback, "messages_generated", { "case_id": case.id, "messages": messages, }, ) return messages except Exception as exc: return await _fail(f"LLM 生成消息失败: {exc}") async def _generate_messages_legacy(self, messages: list[dict[str, str]]) -> str: """Temporary fallback for scenarios not yet migrated to model bindings.""" import httpx from agenteval.utils.llm import extract_content_from_llm_response llm_config = self.scenario.llm_config if not llm_config or not llm_config.get("api_url"): raise ValueError("动态用例未绑定生成模型,且兼容 llm_config 缺少 api_url") headers = {"Content-Type": "application/json"} if llm_config.get("api_key"): headers["Authorization"] = f"Bearer {llm_config['api_key']}" payload = { "model": llm_config.get("model", "doubao-seed-2.0-lite"), "messages": messages, "temperature": 0.7, } async with httpx.AsyncClient(timeout=self.timeout_config.llm_generate) as client: response = await client.post(llm_config["api_url"], headers=headers, json=payload) response.raise_for_status() content = extract_content_from_llm_response(response.json()) if not content: raise ValueError("LLM 返回内容为空或无法解析") return content def _resolve_model(self, purpose: ModelPurpose | None) -> ModelRuntimeConfig | None: if purpose is None: return None if purpose in self._resolved_models: return self._resolved_models[purpose] config_id = self.scenario.model_bindings.get(purpose) if not config_id: return None expected = { ModelPurpose.GENERATOR: ModelCapability.CHAT, ModelPurpose.JUDGE: ModelCapability.CHAT, ModelPurpose.EMBEDDING: ModelCapability.EMBEDDING, ModelPurpose.MODERATION: ModelCapability.MODERATION, }[purpose] runtime = self.model_service.resolve(config_id, expected) self._resolved_models[purpose] = runtime return runtime # ── helpers ──────────────────────────────────────────────────────── def _check_cancel(self) -> None: if self.cancel_token.is_set(): raise CancelledError("run cancelled") async def _emit( self, callback: Optional[ProgressCallback], event: str, data: dict[str, Any], ) -> None: if not callback: return try: result = callback(event, data) if asyncio.iscoroutine(result): await result except Exception: pass