refactor(tasks): route LLM background tasks through TaskRegistry
架构保养第二轮候选 1:分析 / 周期对比 / judge 复核三条 LLM 任务链 收进各自的模块级 TaskRegistry(强引用防 GC、按 id 幂等、shutdown 统一收敛),删除 judge 的 _BACKGROUND_TASKS 私货,start_* 不再返回 无人消费的 Task。启动清理块补两笔 orphan 清扫:滞留的 generating 分析与对比行标记为 failed,与僵尸运行清扫同构。新增 7 个单测。
This commit is contained in:
parent
df76edcf55
commit
f3a528611e
@ -29,6 +29,7 @@ from agenteval.storage.repository import (
|
|||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
RunRepository,
|
RunRepository,
|
||||||
)
|
)
|
||||||
|
from agenteval.task_registry import TaskRegistry
|
||||||
from agenteval.utils.llm import extract_reply_text, parse_json_from_llm_text
|
from agenteval.utils.llm import extract_reply_text, parse_json_from_llm_text
|
||||||
|
|
||||||
_logger = logging.getLogger("agenteval")
|
_logger = logging.getLogger("agenteval")
|
||||||
@ -330,6 +331,15 @@ async def execute_campaign_analysis(
|
|||||||
session.close()
|
session.close()
|
||||||
|
|
||||||
|
|
||||||
def start_campaign_analysis(campaign_id: str, *, triggered_by: str) -> asyncio.Task:
|
analysis_registry = TaskRegistry()
|
||||||
"""以后台任务启动分析生成(fire-and-forget;状态经 campaign_analyses 表观测)。"""
|
|
||||||
return asyncio.create_task(execute_campaign_analysis(campaign_id, triggered_by=triggered_by))
|
|
||||||
|
def start_campaign_analysis(campaign_id: str, *, triggered_by: str) -> None:
|
||||||
|
"""以后台任务启动分析生成(状态经 campaign_analyses 表观测)。
|
||||||
|
|
||||||
|
registry 持强引用防 GC,shutdown 时统一收敛;同 id 在跑时幂等不重复派生。
|
||||||
|
"""
|
||||||
|
analysis_registry.launch(
|
||||||
|
campaign_id,
|
||||||
|
lambda _cancel: execute_campaign_analysis(campaign_id, triggered_by=triggered_by),
|
||||||
|
)
|
||||||
|
|||||||
@ -6,7 +6,6 @@ diff——全部确定性计算(ADR-0004 口径,经 ``generate_campaign_repo
|
|||||||
机械 diff 之上产出结构化演进叙述(CONTEXT.md「周期对比」)。
|
机械 diff 之上产出结构化演进叙述(CONTEXT.md「周期对比」)。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
@ -23,6 +22,7 @@ from agenteval.storage.repository import (
|
|||||||
CampaignPeriodComparisonRepository,
|
CampaignPeriodComparisonRepository,
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
)
|
)
|
||||||
|
from agenteval.task_registry import TaskRegistry
|
||||||
from agenteval.utils.llm import parse_json_from_llm_text
|
from agenteval.utils.llm import parse_json_from_llm_text
|
||||||
|
|
||||||
_logger = logging.getLogger("agenteval")
|
_logger = logging.getLogger("agenteval")
|
||||||
@ -383,17 +383,24 @@ async def execute_campaign_comparison(
|
|||||||
session.close()
|
session.close()
|
||||||
|
|
||||||
|
|
||||||
|
comparison_registry = TaskRegistry()
|
||||||
|
|
||||||
|
|
||||||
def start_campaign_comparison(
|
def start_campaign_comparison(
|
||||||
campaign_id: str,
|
campaign_id: str,
|
||||||
*,
|
*,
|
||||||
triggered_by: str,
|
triggered_by: str,
|
||||||
baseline_campaign_id: Optional[str] = None,
|
baseline_campaign_id: Optional[str] = None,
|
||||||
) -> asyncio.Task:
|
) -> None:
|
||||||
"""以后台任务启动对比生成(fire-and-forget;状态经 campaign_period_comparisons 表观测)。"""
|
"""以后台任务启动对比生成(状态经 campaign_period_comparisons 表观测)。
|
||||||
return asyncio.create_task(
|
|
||||||
execute_campaign_comparison(
|
registry 持强引用防 GC,shutdown 时统一收敛;同 id 在跑时幂等不重复派生。
|
||||||
|
"""
|
||||||
|
comparison_registry.launch(
|
||||||
|
campaign_id,
|
||||||
|
lambda _cancel: execute_campaign_comparison(
|
||||||
campaign_id,
|
campaign_id,
|
||||||
triggered_by=triggered_by,
|
triggered_by=triggered_by,
|
||||||
baseline_campaign_id=baseline_campaign_id,
|
baseline_campaign_id=baseline_campaign_id,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
|
|||||||
@ -7,7 +7,6 @@
|
|||||||
注入,测试用假客户端覆盖。
|
注入,测试用假客户端覆盖。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
@ -28,6 +27,7 @@ from agenteval.storage.repository import (
|
|||||||
ExplorationMessageRepository,
|
ExplorationMessageRepository,
|
||||||
ExplorationSessionRepository,
|
ExplorationSessionRepository,
|
||||||
)
|
)
|
||||||
|
from agenteval.task_registry import TaskRegistry
|
||||||
from agenteval.utils.llm import parse_json_from_llm_text
|
from agenteval.utils.llm import parse_json_from_llm_text
|
||||||
|
|
||||||
_logger = logging.getLogger("agenteval")
|
_logger = logging.getLogger("agenteval")
|
||||||
@ -180,13 +180,15 @@ async def execute_judge_review(
|
|||||||
session.close()
|
session.close()
|
||||||
|
|
||||||
|
|
||||||
_BACKGROUND_TASKS: set[asyncio.Task] = set()
|
judge_registry = TaskRegistry()
|
||||||
|
|
||||||
|
|
||||||
def start_judge_review(exploration_session_id: str) -> asyncio.Task:
|
def start_judge_review(exploration_session_id: str) -> None:
|
||||||
"""以后台任务启动抽样复核(fire-and-forget;结果经会话 judge_review 观测)。"""
|
"""以后台任务启动抽样复核(结果经会话 judge_review 观测)。
|
||||||
task = asyncio.create_task(execute_judge_review(exploration_session_id))
|
|
||||||
# 强引用防止未持有引用的任务被 GC(asyncio 已知坑)
|
registry 持强引用防 GC,shutdown 时统一收敛;同 id 在跑时幂等不重复派生。
|
||||||
_BACKGROUND_TASKS.add(task)
|
"""
|
||||||
task.add_done_callback(_BACKGROUND_TASKS.discard)
|
judge_registry.launch(
|
||||||
return task
|
exploration_session_id,
|
||||||
|
lambda _cancel: execute_judge_review(exploration_session_id),
|
||||||
|
)
|
||||||
|
|||||||
@ -482,6 +482,22 @@ class CampaignAnalysisRepository:
|
|||||||
self.session.refresh(row)
|
self.session.refresh(row)
|
||||||
return row
|
return row
|
||||||
|
|
||||||
|
def mark_orphans_failed(self) -> int:
|
||||||
|
"""服务启动时清理:把滞留的 generating 分析行标记为 failed。
|
||||||
|
|
||||||
|
分析任务是进程内 asyncio 任务,服务重启后不会恢复;不清理则这些
|
||||||
|
行永远停留在 generating(僵尸状态)。
|
||||||
|
"""
|
||||||
|
rows = self.session.exec(select(CampaignAnalysisDB).where(CampaignAnalysisDB.status == "generating")).all()
|
||||||
|
for row in rows:
|
||||||
|
row.status = "failed"
|
||||||
|
row.error = "服务重启导致分析生成中断"
|
||||||
|
row.updated_at = utc_now()
|
||||||
|
self.session.add(row)
|
||||||
|
if rows:
|
||||||
|
self.session.commit()
|
||||||
|
return len(rows)
|
||||||
|
|
||||||
|
|
||||||
class CampaignPeriodComparisonRepository:
|
class CampaignPeriodComparisonRepository:
|
||||||
"""Repository for period-comparison rows (one per campaign, upserted)."""
|
"""Repository for period-comparison rows (one per campaign, upserted)."""
|
||||||
@ -526,6 +542,24 @@ class CampaignPeriodComparisonRepository:
|
|||||||
self.session.refresh(row)
|
self.session.refresh(row)
|
||||||
return row
|
return row
|
||||||
|
|
||||||
|
def mark_orphans_failed(self) -> int:
|
||||||
|
"""服务启动时清理:把滞留的 generating 周期对比行标记为 failed。
|
||||||
|
|
||||||
|
对比任务是进程内 asyncio 任务,服务重启后不会恢复;不清理则这些
|
||||||
|
行永远停留在 generating(僵尸状态)。基线配对保留,可直接重新触发。
|
||||||
|
"""
|
||||||
|
rows = self.session.exec(
|
||||||
|
select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.status == "generating")
|
||||||
|
).all()
|
||||||
|
for row in rows:
|
||||||
|
row.status = "failed"
|
||||||
|
row.error = "服务重启导致周期对比生成中断"
|
||||||
|
row.updated_at = utc_now()
|
||||||
|
self.session.add(row)
|
||||||
|
if rows:
|
||||||
|
self.session.commit()
|
||||||
|
return len(rows)
|
||||||
|
|
||||||
|
|
||||||
class ResultRepository:
|
class ResultRepository:
|
||||||
"""Repository for evaluation results."""
|
"""Repository for evaluation results."""
|
||||||
|
|||||||
@ -10,7 +10,11 @@ from fastapi.responses import FileResponse, JSONResponse
|
|||||||
|
|
||||||
from agenteval.config import get_settings
|
from agenteval.config import get_settings
|
||||||
from agenteval.storage.db import get_session, init_db
|
from agenteval.storage.db import get_session, init_db
|
||||||
from agenteval.storage.repository import RunRepository
|
from agenteval.storage.repository import (
|
||||||
|
CampaignAnalysisRepository,
|
||||||
|
CampaignPeriodComparisonRepository,
|
||||||
|
RunRepository,
|
||||||
|
)
|
||||||
from agenteval.version import get_build_info, get_version
|
from agenteval.version import get_build_info, get_version
|
||||||
from agenteval.web.deps import require_api_key
|
from agenteval.web.deps import require_api_key
|
||||||
from agenteval.web.routers import (
|
from agenteval.web.routers import (
|
||||||
@ -39,6 +43,10 @@ async def lifespan(_: FastAPI):
|
|||||||
count = RunRepository(session).mark_orphans_failed()
|
count = RunRepository(session).mark_orphans_failed()
|
||||||
if count:
|
if count:
|
||||||
logging.getLogger("agenteval").warning("启动清理:%d 个中断的运行已标记为 failed", count)
|
logging.getLogger("agenteval").warning("启动清理:%d 个中断的运行已标记为 failed", count)
|
||||||
|
llm_orphans = CampaignAnalysisRepository(session).mark_orphans_failed()
|
||||||
|
llm_orphans += CampaignPeriodComparisonRepository(session).mark_orphans_failed()
|
||||||
|
if llm_orphans:
|
||||||
|
logging.getLogger("agenteval").warning("启动清理:%d 条中断的分析/对比已标记为 failed", llm_orphans)
|
||||||
# 据库恢复所有未完成的评估活动,重建其调度循环(不重复派生已到点条目)
|
# 据库恢复所有未完成的评估活动,重建其调度循环(不重复派生已到点条目)
|
||||||
from agenteval.evaluation.campaign_runner import resume_running_campaigns
|
from agenteval.evaluation.campaign_runner import resume_running_campaigns
|
||||||
|
|
||||||
@ -50,13 +58,20 @@ async def lifespan(_: FastAPI):
|
|||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logging.getLogger("agenteval").warning("启动清理失败(忽略): %s", exc)
|
logging.getLogger("agenteval").warning("启动清理失败(忽略): %s", exc)
|
||||||
yield
|
yield
|
||||||
# 优雅停止所有进程内任务:先停活动调度循环,再停在跑的评测运行
|
# 优雅停止所有进程内任务:先停活动调度循环,再停在跑的评测运行,
|
||||||
|
# 最后停三条 LLM 任务链(分析 / 周期对比 / judge 复核)。
|
||||||
try:
|
try:
|
||||||
|
from agenteval.evaluation.analysis import analysis_registry
|
||||||
from agenteval.evaluation.campaign_runner import shutdown_all
|
from agenteval.evaluation.campaign_runner import shutdown_all
|
||||||
|
from agenteval.evaluation.comparison import comparison_registry
|
||||||
|
from agenteval.exploration.judge import judge_registry
|
||||||
from agenteval.web.routers.runs import run_registry
|
from agenteval.web.routers.runs import run_registry
|
||||||
|
|
||||||
await shutdown_all()
|
await shutdown_all()
|
||||||
await run_registry.shutdown_all()
|
await run_registry.shutdown_all()
|
||||||
|
await analysis_registry.shutdown_all()
|
||||||
|
await comparison_registry.shutdown_all()
|
||||||
|
await judge_registry.shutdown_all()
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logging.getLogger("agenteval").warning("活动调度停止失败(忽略): %s", exc)
|
logging.getLogger("agenteval").warning("活动调度停止失败(忽略): %s", exc)
|
||||||
|
|
||||||
|
|||||||
111
tests/unit/test_llm_task_lifecycle.py
Normal file
111
tests/unit/test_llm_task_lifecycle.py
Normal file
@ -0,0 +1,111 @@
|
|||||||
|
"""LLM 后台任务生命周期直测(架构保养第二轮候选 1)。
|
||||||
|
|
||||||
|
分析 / 对比 / judge 复核三条 LLM 任务链收进各自的 TaskRegistry:
|
||||||
|
强引用防 GC、按 id 幂等、shutdown 统一收敛;启动清扫把滞留的
|
||||||
|
generating 行标记为 failed。
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from agenteval.evaluation import analysis, comparison
|
||||||
|
from agenteval.exploration import judge
|
||||||
|
from agenteval.storage.repository import (
|
||||||
|
CampaignAnalysisRepository,
|
||||||
|
CampaignPeriodComparisonRepository,
|
||||||
|
)
|
||||||
|
from sqlmodel import Session, SQLModel, create_engine
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture()
|
||||||
|
def db_session(tmp_path):
|
||||||
|
from agenteval.storage.db import ( # noqa: F401
|
||||||
|
CampaignAnalysisDB,
|
||||||
|
CampaignDB,
|
||||||
|
CampaignPeriodComparisonDB,
|
||||||
|
)
|
||||||
|
|
||||||
|
engine = create_engine(
|
||||||
|
f"sqlite:///{tmp_path / 'llmtasks.db'}",
|
||||||
|
connect_args={"check_same_thread": False},
|
||||||
|
)
|
||||||
|
SQLModel.metadata.create_all(engine)
|
||||||
|
session = Session(engine)
|
||||||
|
try:
|
||||||
|
yield session
|
||||||
|
finally:
|
||||||
|
session.close()
|
||||||
|
engine.dispose()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_start_analysis_registers_in_registry(monkeypatch):
|
||||||
|
monkeypatch.setattr(analysis, "execute_campaign_analysis", lambda *a, **kw: asyncio.sleep(0))
|
||||||
|
analysis.start_campaign_analysis("c-1", triggered_by="manual")
|
||||||
|
assert analysis.analysis_registry.is_running("c-1")
|
||||||
|
await analysis.analysis_registry.shutdown_all()
|
||||||
|
assert not analysis.analysis_registry.is_running("c-1")
|
||||||
|
|
||||||
|
|
||||||
|
async def test_start_comparison_registers_in_registry(monkeypatch):
|
||||||
|
monkeypatch.setattr(comparison, "execute_campaign_comparison", lambda *a, **kw: asyncio.sleep(0))
|
||||||
|
comparison.start_campaign_comparison("c-2", triggered_by="manual")
|
||||||
|
assert comparison.comparison_registry.is_running("c-2")
|
||||||
|
await comparison.comparison_registry.shutdown_all()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_start_judge_registers_in_registry(monkeypatch):
|
||||||
|
monkeypatch.setattr(judge, "execute_judge_review", lambda *a, **kw: asyncio.sleep(0))
|
||||||
|
judge.start_judge_review("s-1")
|
||||||
|
assert judge.judge_registry.is_running("s-1")
|
||||||
|
await judge.judge_registry.shutdown_all()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_launch_is_idempotent_for_live_id(monkeypatch):
|
||||||
|
gate = asyncio.Event()
|
||||||
|
|
||||||
|
async def hang(*args, **kwargs):
|
||||||
|
await gate.wait()
|
||||||
|
|
||||||
|
monkeypatch.setattr(analysis, "execute_campaign_analysis", hang)
|
||||||
|
analysis.start_campaign_analysis("c-dup", triggered_by="manual")
|
||||||
|
analysis.start_campaign_analysis("c-dup", triggered_by="manual")
|
||||||
|
assert len(analysis.analysis_registry._tasks) == 1
|
||||||
|
gate.set()
|
||||||
|
await analysis.analysis_registry.shutdown_all()
|
||||||
|
|
||||||
|
|
||||||
|
async def test_shutdown_all_cancels_hanging_task(monkeypatch):
|
||||||
|
async def hang(*args, **kwargs):
|
||||||
|
await asyncio.Event().wait()
|
||||||
|
|
||||||
|
monkeypatch.setattr(comparison, "execute_campaign_comparison", hang)
|
||||||
|
comparison.start_campaign_comparison("c-hang", triggered_by="manual")
|
||||||
|
assert comparison.comparison_registry.is_running("c-hang")
|
||||||
|
await comparison.comparison_registry.shutdown_all()
|
||||||
|
assert not comparison.comparison_registry.is_running("c-hang")
|
||||||
|
|
||||||
|
|
||||||
|
def test_mark_orphans_failed_flips_generating_analysis(db_session):
|
||||||
|
repo = CampaignAnalysisRepository(db_session)
|
||||||
|
repo.upsert("c-gen", status="generating", triggered_by="auto")
|
||||||
|
repo.upsert("c-done", status="completed", result={"ok": True}, triggered_by="auto")
|
||||||
|
|
||||||
|
count = repo.mark_orphans_failed()
|
||||||
|
|
||||||
|
assert count == 1
|
||||||
|
assert repo.get_by_campaign("c-gen").status == "failed"
|
||||||
|
assert repo.get_by_campaign("c-gen").error
|
||||||
|
assert repo.get_by_campaign("c-done").status == "completed"
|
||||||
|
|
||||||
|
|
||||||
|
def test_mark_orphans_failed_flips_generating_comparison(db_session):
|
||||||
|
repo = CampaignPeriodComparisonRepository(db_session)
|
||||||
|
repo.upsert("c-gen", status="generating", baseline_campaign_id="b-1", triggered_by="auto")
|
||||||
|
repo.upsert("c-done", status="completed", baseline_campaign_id="b-2", result={"ok": True}, triggered_by="auto")
|
||||||
|
|
||||||
|
count = repo.mark_orphans_failed()
|
||||||
|
|
||||||
|
assert count == 1
|
||||||
|
assert repo.get_by_campaign("c-gen").status == "failed"
|
||||||
|
assert repo.get_by_campaign("c-gen").error
|
||||||
|
assert repo.get_by_campaign("c-done").status == "completed"
|
||||||
Loading…
Reference in New Issue
Block a user