From 0cca4963d10e54f802c8605e82b3e41e73171936 Mon Sep 17 00:00:00 2001 From: sinohqb Date: Fri, 31 Jul 2026 03:39:03 +0800 Subject: [PATCH] refactor(tasks): unify run/campaign task registries into TaskRegistry Both the single-run path and the campaign scheduler drove long-lived asyncio tasks through their own duplicated _tasks/_cancel_events dicts and shutdown loops. Collapse them into one deep TaskRegistry module, instantiated as run_registry and campaign_registry. launch() creates the cancel event before the task (so a cancel during startup is never lost), wires done-callback cleanup, and is idempotent per id; this makes runs.py's hard-cancel fallback provably dead, so it is removed. App shutdown now gracefully stops in-flight runs too, not just campaigns. --- .../agenteval/evaluation/campaign_runner.py | 43 +++------ backend/agenteval/task_registry.py | 87 +++++++++++++++++++ backend/agenteval/web/app.py | 4 +- backend/agenteval/web/routers/runs.py | 34 +++----- .../test_campaign_scheduler_loop.py | 2 +- tests/unit/test_task_registry.py | 84 ++++++++++++++++++ 6 files changed, 198 insertions(+), 56 deletions(-) create mode 100644 backend/agenteval/task_registry.py create mode 100644 tests/unit/test_task_registry.py diff --git a/backend/agenteval/evaluation/campaign_runner.py b/backend/agenteval/evaluation/campaign_runner.py index b767257..13b9362 100644 --- a/backend/agenteval/evaluation/campaign_runner.py +++ b/backend/agenteval/evaluation/campaign_runner.py @@ -37,6 +37,7 @@ from agenteval.storage.repository import ( ScenarioRepository, TargetRepository, ) +from agenteval.task_registry import TaskRegistry _SPAWNED_KEY = "spawned_indices" _ERRORS_KEY = "errors" @@ -48,9 +49,8 @@ DEFAULT_TICK_SECONDS = 1.0 _logger = logging.getLogger("agenteval") # Live loop tasks + cooperative cancel events, keyed by campaign id. Authority -# is the DB; these are only the in-process handles for the running loop. -_tasks: dict[str, asyncio.Task] = {} -_cancel_events: dict[str, asyncio.Event] = {} +# is the DB; this registry only holds the in-process handles for the loop. +campaign_registry = TaskRegistry() @dataclass @@ -182,7 +182,9 @@ async def advance_campaign( # ── durable scheduler loop ────────────────────────────────────────────────── -async def run_campaign_loop(campaign_id: str, *, tick_seconds: float = DEFAULT_TICK_SECONDS) -> None: +async def run_campaign_loop( + campaign_id: str, cancel: asyncio.Event, *, tick_seconds: float = DEFAULT_TICK_SECONDS +) -> None: """Drive one campaign to completion, ticking on real wall-clock time. The loop holds no authority and no decision logic: every tick it reloads the @@ -190,7 +192,6 @@ async def run_campaign_loop(campaign_id: str, *, tick_seconds: float = DEFAULT_T It exits when the window finishes, the campaign leaves RUNNING (e.g. it was cancelled), or the cooperative cancel event fires. """ - cancel = _cancel_events.setdefault(campaign_id, asyncio.Event()) session = get_session() try: while not cancel.is_set(): @@ -237,8 +238,6 @@ async def run_campaign_loop(campaign_id: str, *, tick_seconds: float = DEFAULT_T _logger.warning("活动调度循环 %s 异常退出: %s", campaign_id, exc) finally: session.close() - _cancel_events.pop(campaign_id, None) - _tasks.pop(campaign_id, None) def start_campaign(campaign_id: str, session: Session, *, tick_seconds: float = DEFAULT_TICK_SECONDS) -> Optional[asyncio.Task]: @@ -251,28 +250,21 @@ def start_campaign(campaign_id: str, session: Session, *, tick_seconds: float = campaign = repo.get(campaign_id) if not campaign or campaign.status in (CampaignStatus.COMPLETED, CampaignStatus.CANCELLED, CampaignStatus.FAILED): return None - if campaign_id in _tasks: - return _tasks[campaign_id] if campaign.status == CampaignStatus.PLANNED: campaign.status = CampaignStatus.RUNNING campaign.started_at = utc_now() repo.update(campaign) - _cancel_events.setdefault(campaign_id, asyncio.Event()) - task = asyncio.create_task( - run_campaign_loop(campaign_id, tick_seconds=tick_seconds), - name=f"campaign-{campaign_id}", + return campaign_registry.launch( + campaign_id, + lambda cancel: run_campaign_loop(campaign_id, cancel, tick_seconds=tick_seconds), ) - _tasks[campaign_id] = task - return task def request_cancel(campaign_id: str) -> None: """Signal the loop (if live) to stop spawning and exit promptly.""" - event = _cancel_events.get(campaign_id) - if event is not None: - event.set() + campaign_registry.cancel(campaign_id) def resume_running_campaigns(session: Session, *, tick_seconds: float = DEFAULT_TICK_SECONDS) -> int: @@ -287,17 +279,4 @@ def resume_running_campaigns(session: Session, *, tick_seconds: float = DEFAULT_ async def shutdown_all() -> None: """Gracefully stop all live loops (Web app shutdown).""" - for event in list(_cancel_events.values()): - event.set() - tasks = list(_tasks.values()) - for task in tasks: - task.cancel() - for task in tasks: - # CancelledError is expected here (we just cancelled the task) and does - # not derive from Exception in 3.8+, so it must be listed explicitly. - try: - await task - except (asyncio.CancelledError, Exception): - pass - _tasks.clear() - _cancel_events.clear() + await campaign_registry.shutdown_all() diff --git a/backend/agenteval/task_registry.py b/backend/agenteval/task_registry.py new file mode 100644 index 0000000..d803c2e --- /dev/null +++ b/backend/agenteval/task_registry.py @@ -0,0 +1,87 @@ +"""Durable in-process task registry. + +Both the single-run path (``web.routers.runs``) and the campaign scheduler +(``evaluation.campaign_runner``) drive long-lived ``asyncio`` tasks whose +authority lives in the DB — the in-process handle is only there to cancel or +gracefully stop the task. This module is the one deep seam they share: it owns +the (task, cancel-event) pair per id, wires cleanup, and stops everything on +shutdown. Each caller instantiates its own registry, so id spaces never collide. +""" + +import asyncio +import logging +from typing import Awaitable, Callable + +_logger = logging.getLogger("agenteval") + +# A coroutine factory: given the cancel event the registry owns, return the +# coroutine to run. The event exists before the task starts, so a cancel that +# arrives during startup is never lost. +CoroFactory = Callable[[asyncio.Event], Awaitable[None]] + + +class TaskRegistry: + """Tracks live asyncio tasks and their cooperative cancel events by id.""" + + def __init__(self) -> None: + self._tasks: dict[str, asyncio.Task] = {} + self._cancel_events: dict[str, asyncio.Event] = {} + + def launch(self, id: str, make_coro: CoroFactory) -> asyncio.Task: + """Start a tracked task for ``id`` (idempotent for a live id). + + Creates the cancel event first, hands it to ``make_coro`` to build the + coroutine, then schedules the task and wires a done-callback that pops + both handles when it ends. If ``id`` already has a live task, returns it + without starting a second coroutine. + """ + existing = self._tasks.get(id) + if existing is not None: + return existing + + cancel = self._cancel_events.setdefault(id, asyncio.Event()) + task = asyncio.create_task(make_coro(cancel), name=id) + self._tasks[id] = task + task.add_done_callback(lambda _t, _id=id: self._cleanup(_id)) + return task + + def cancel(self, id: str) -> bool: + """Signal the cooperative cancel event for ``id``. + + Returns True if a live cancel event was signalled, False if there was + nothing to signal (e.g. the task already ended or never existed) — the + caller uses this to decide any fallback such as marking the DB directly. + """ + event = self._cancel_events.get(id) + if event is None: + return False + event.set() + return True + + def get(self, id: str) -> "asyncio.Task | None": + """Return the live task for ``id``, or None if there is none.""" + return self._tasks.get(id) + + def is_running(self, id: str) -> bool: + return id in self._tasks + + async def shutdown_all(self) -> None: + """Cooperatively stop, then hard-cancel and await, every live task.""" + for event in list(self._cancel_events.values()): + event.set() + tasks = list(self._tasks.values()) + for task in tasks: + task.cancel() + for task in tasks: + # CancelledError is expected (we just cancelled the task) and does + # not derive from Exception in 3.8+, so list it explicitly. + try: + await task + except (asyncio.CancelledError, Exception) as exc: # noqa: BLE001 + _logger.debug("task registry shutdown: task ended with %r", exc) + self._tasks.clear() + self._cancel_events.clear() + + def _cleanup(self, id: str) -> None: + self._tasks.pop(id, None) + self._cancel_events.pop(id, None) diff --git a/backend/agenteval/web/app.py b/backend/agenteval/web/app.py index 6d22371..f15d2ba 100644 --- a/backend/agenteval/web/app.py +++ b/backend/agenteval/web/app.py @@ -38,11 +38,13 @@ async def lifespan(_: FastAPI): except Exception as exc: logging.getLogger("agenteval").warning("启动清理失败(忽略): %s", exc) yield - # 优雅停止所有活动调度循环 + # 优雅停止所有进程内任务:先停活动调度循环,再停在跑的评测运行 try: from agenteval.evaluation.campaign_runner import shutdown_all + from agenteval.web.routers.runs import run_registry await shutdown_all() + await run_registry.shutdown_all() except Exception as exc: logging.getLogger("agenteval").warning("活动调度停止失败(忽略): %s", exc) diff --git a/backend/agenteval/web/routers/runs.py b/backend/agenteval/web/routers/runs.py index 38e70e2..dedef9c 100644 --- a/backend/agenteval/web/routers/runs.py +++ b/backend/agenteval/web/routers/runs.py @@ -1,7 +1,6 @@ """API routes for evaluation runs.""" import asyncio -from typing import Optional from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel @@ -11,6 +10,7 @@ from agenteval.evaluation.engine import EvalEngine from agenteval.models import EvalRun, RunStatus, 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 @@ -28,15 +28,14 @@ class StartRunRequest(BaseModel): # ── 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] = {} +run_registry = TaskRegistry() -async def _run_evaluation(run_id: str, target_id: str, scenario_id: str) -> None: +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() - cancel_token = asyncio.Event() - _cancel_tokens[run_id] = cancel_token try: target = TargetRepository(session).get(target_id) scenario = ScenarioRepository(session).get(scenario_id) @@ -64,8 +63,6 @@ async def _run_evaluation(run_id: str, target_id: str, scenario_id: str) -> None ) finally: session.close() - _cancel_tokens.pop(run_id, None) - _tasks.pop(run_id, None) @router.get("") @@ -100,11 +97,12 @@ async def start_run( ) 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}", + run_registry.launch( + run.id, + lambda cancel_token: _run_evaluation( + run.id, request.target_id, request.scenario_id, cancel_token=cancel_token + ), ) - _tasks[run.id] = task return run.model_dump() @@ -125,16 +123,8 @@ async def cancel_run(run_id: str, session: Session = Depends(get_db)) -> dict: 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: + 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 = { diff --git a/tests/integration/test_campaign_scheduler_loop.py b/tests/integration/test_campaign_scheduler_loop.py index 1339eb3..3eb65da 100644 --- a/tests/integration/test_campaign_scheduler_loop.py +++ b/tests/integration/test_campaign_scheduler_loop.py @@ -70,7 +70,7 @@ def _make_campaign(session, **overrides) -> Campaign: async def _await_task(campaign_id, timeout=3.0): import asyncio - task = campaign_runner._tasks.get(campaign_id) + task = campaign_runner.campaign_registry.get(campaign_id) if task is not None: await asyncio.wait_for(task, timeout=timeout) diff --git a/tests/unit/test_task_registry.py b/tests/unit/test_task_registry.py new file mode 100644 index 0000000..97163d5 --- /dev/null +++ b/tests/unit/test_task_registry.py @@ -0,0 +1,84 @@ +"""Unit tests for TaskRegistry — the durable in-process task registry. + +Covers the single seam both runs and campaigns share: launch (with an +injected cancel event), idempotency, cooperative cancel, done-callback +cleanup, and graceful shutdown_all. +""" + +import asyncio + +from agenteval.task_registry import TaskRegistry + + +async def test_launch_tracks_task_and_passes_cancel_event(): + reg = TaskRegistry() + seen: dict = {} + + async def coro(cancel: asyncio.Event): + seen["event"] = cancel + await cancel.wait() + + task = reg.launch("a", lambda ev: coro(ev)) + assert reg.is_running("a") + await asyncio.sleep(0) # let the coroutine start and capture the event + assert isinstance(seen["event"], asyncio.Event) + + reg.cancel("a") + await task + assert not reg.is_running("a") # done-callback popped it + + +async def test_launch_is_idempotent_for_live_id(): + reg = TaskRegistry() + + async def coro(cancel: asyncio.Event): + await cancel.wait() + + t1 = reg.launch("a", lambda ev: coro(ev)) + t2 = reg.launch("a", lambda ev: coro(ev)) + assert t1 is t2 # same live task, second coro never started + + reg.cancel("a") + await t1 + + +async def test_cancel_returns_true_when_live_false_otherwise(): + reg = TaskRegistry() + + async def coro(cancel: asyncio.Event): + await cancel.wait() + + reg.launch("a", lambda ev: coro(ev)) + assert reg.cancel("a") is True # live event signalled + assert reg.cancel("missing") is False # nothing to signal + await asyncio.sleep(0) # let the cancelled coro finish + + +async def test_done_callback_cleans_up_after_natural_completion(): + reg = TaskRegistry() + + async def quick(cancel: asyncio.Event): + return "done" + + task = reg.launch("a", lambda ev: quick(ev)) + await task + await asyncio.sleep(0) # allow done-callback to run + assert not reg.is_running("a") + + +async def test_shutdown_all_cancels_and_awaits_every_task(): + reg = TaskRegistry() + started = [] + + async def coro(cancel: asyncio.Event): + started.append(1) + await cancel.wait() + + reg.launch("a", lambda ev: coro(ev)) + reg.launch("b", lambda ev: coro(ev)) + await asyncio.sleep(0) # let both start + assert len(started) == 2 + + await reg.shutdown_all() + assert not reg.is_running("a") + assert not reg.is_running("b")