refactor(evaluation/storage): 结算统一与 repository 拆分(Phase 2 + 3)
合并两个不可分割的深化: Phase 2 — 智能作业结算统一(ADR-0012) - intelligence_jobs.execute(job_kind, campaign_id, ...) 作为结算的 唯一实现:建行 → 认领 → 校验 → generating → 落账,一处编排、 一处截断(500 字符)。两个 executor 退化为 ensure_queued / validate / work_fn 三个小 adapter。 - analysis.validate_analysis_request() 共享校验入口(活动终态 → 模型),路由捕获映射 400、executor 捕获落 failed 行,与 validate_comparison_request 先例同构。 - campaign_runner._auto_start_analysis 的跳过守卫收敛至 auto_intelligence_eligible 单一判断点。 - comparison.py 删除零调用的 build_comparison_payload; load_comparison_view 投影归位至 campaign_read_model。 - 新增 characterization 测试(认领竞争、重复触发、截断、恢复上限)。 Phase 3 — storage/repository.py 拆分 - AsyncJobRepository 及两个子类迁至 storage/async_job_repository.py(Phase 2 的 intelligence_jobs 与 comparison 必须 import 自该路径,故与 Phase 2 同 commit)。 - ExplorationSession / ExplorationMessage 迁至 storage/exploration_repository.py;repository.py 由 1180 行降至 约 814 行,grep 确认无残留符号。 - exploration 子模块与路由 import 全部更新;测试 import 跟随。 刻意不做:CAS 共享原语、app.py 五 registry 关停顺序归一 (ADR-0006 精神,等真实需求出现再议)。
This commit is contained in:
parent
71543f042a
commit
7eae6de52d
@ -13,7 +13,7 @@ from typing import Any, Awaitable, Callable, Optional
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.model_gateway import ModelGateway
|
||||
from agenteval.models import Campaign, ModelCapability, RunStatus
|
||||
from agenteval.models import Campaign, CampaignStatus, ModelCapability, RunStatus
|
||||
from agenteval.services.model_configs import (
|
||||
ModelConfigError,
|
||||
ModelConfigService,
|
||||
@ -52,6 +52,20 @@ def resolve_analysis_model(campaign: Campaign, session: Session) -> Optional[Mod
|
||||
return None
|
||||
|
||||
|
||||
def validate_analysis_request(session: Session, campaign: Campaign) -> ModelRuntimeConfig:
|
||||
"""共享校验入口:活动终态 → 模型。违规抛 AnalysisError。
|
||||
|
||||
router 触发端点捕获映射 400;执行器捕获落 failed 行。校验顺序权威,
|
||||
两处不再漂移(validate_comparison_request 先例)。
|
||||
"""
|
||||
if campaign.status in (CampaignStatus.PLANNED, CampaignStatus.RUNNING):
|
||||
raise AnalysisError("活动完成后才能生成智能分析")
|
||||
runtime = resolve_analysis_model(campaign, session)
|
||||
if runtime is None:
|
||||
raise AnalysisError("未配置分析模型:请在模型配置中心将某个 chat 配置设为「分析默认」,或为该活动指定分析模型")
|
||||
return runtime
|
||||
|
||||
|
||||
def collect_failure_samples(
|
||||
campaign_id: str,
|
||||
session: Session,
|
||||
@ -76,12 +90,14 @@ def collect_failure_samples(
|
||||
turn = turns.get(result.turn_id)
|
||||
user = extract_reply_text(turn.get_sent_message().get("msgBody")) if turn else ""
|
||||
reply = extract_reply_text(turn.get_reply().get("msgBody")) if turn and turn.get_reply() else ""
|
||||
bucket.append({
|
||||
bucket.append(
|
||||
{
|
||||
"run_id": run.id or "",
|
||||
"user": user[:text_limit],
|
||||
"reply": reply[:text_limit],
|
||||
"reason": (result.reason or "")[:text_limit],
|
||||
})
|
||||
}
|
||||
)
|
||||
return {sid: items for sid, items in samples.items() if items}
|
||||
|
||||
|
||||
@ -123,10 +139,12 @@ async def _analyze_scenario(
|
||||
ensure_ascii=False,
|
||||
)
|
||||
parsed = _parse_stage(
|
||||
await chat_client([
|
||||
await chat_client(
|
||||
[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
]),
|
||||
]
|
||||
),
|
||||
f"场景「{entry.get('scenario_name', entry['scenario_id'])}」阶段一",
|
||||
)
|
||||
narrative = parsed.get("narrative")
|
||||
@ -169,10 +187,12 @@ async def _synthesize(
|
||||
payload["探索发现"] = exploration_summary
|
||||
user_prompt = json.dumps(payload, ensure_ascii=False)
|
||||
parsed = _parse_stage(
|
||||
await chat_client([
|
||||
await chat_client(
|
||||
[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
]),
|
||||
]
|
||||
),
|
||||
"阶段二综合研判",
|
||||
)
|
||||
overall = parsed.get("overall")
|
||||
@ -200,10 +220,9 @@ async def analyze_campaign(
|
||||
if not capability:
|
||||
raise AnalysisError("活动没有可分析的场景数据")
|
||||
|
||||
stage1 = await asyncio.gather(*[
|
||||
_analyze_scenario(entry, failure_samples.get(entry["scenario_id"], []), chat_client)
|
||||
for entry in capability
|
||||
])
|
||||
stage1 = await asyncio.gather(
|
||||
*[_analyze_scenario(entry, failure_samples.get(entry["scenario_id"], []), chat_client) for entry in capability]
|
||||
)
|
||||
stage2 = await _synthesize(campaign, report, list(stage1), chat_client, exploration_summary=exploration_summary)
|
||||
|
||||
valid_scenario_ids = {entry["scenario_id"] for entry in capability}
|
||||
@ -212,13 +231,15 @@ async def analyze_campaign(
|
||||
if not isinstance(p, dict):
|
||||
continue
|
||||
severity = p.get("severity")
|
||||
problems.append({
|
||||
problems.append(
|
||||
{
|
||||
"severity": severity if severity in _VALID_SEVERITIES else "medium",
|
||||
"title": str(p.get("title", "")),
|
||||
"description": str(p.get("description", "")),
|
||||
"scenario_ids": [s for s in p.get("scenario_ids") or [] if s in valid_scenario_ids],
|
||||
"evidence_run_ids": [r for r in p.get("evidence_run_ids") or [] if r in valid_run_ids],
|
||||
})
|
||||
}
|
||||
)
|
||||
suggestions = [
|
||||
{"priority": int(s.get("priority", i + 1)), "text": str(s.get("text", ""))}
|
||||
for i, s in enumerate(stage2.get("suggestions") or [])
|
||||
@ -227,9 +248,7 @@ async def analyze_campaign(
|
||||
return {
|
||||
"overall": stage2["overall"],
|
||||
"problems": problems,
|
||||
"scenario_narratives": [
|
||||
{"scenario_id": s["scenario_id"], "narrative": s["narrative"]} for s in stage1
|
||||
],
|
||||
"scenario_narratives": [{"scenario_id": s["scenario_id"], "narrative": s["narrative"]} for s in stage1],
|
||||
"suggestions": suggestions,
|
||||
}
|
||||
|
||||
|
||||
@ -5,16 +5,19 @@ from typing import Any, Optional
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.campaign_scheduler import clock_offset, elapsed_seconds
|
||||
from agenteval.evaluation.comparison import load_comparison_view
|
||||
from agenteval.evaluation.comparison import compute_metric_diff, resolve_auto_baseline
|
||||
from agenteval.evaluation.report import (
|
||||
build_campaign_timeline,
|
||||
generate_campaign_report,
|
||||
load_campaign_report,
|
||||
summarize_campaign_progress,
|
||||
)
|
||||
from agenteval.exploration.summary import summarize_campaign_exploration
|
||||
from agenteval.models import Campaign
|
||||
from agenteval.storage.async_job_repository import CampaignAnalysisRepository, CampaignPeriodComparisonRepository
|
||||
from agenteval.storage.db import iso_utc, utc_now
|
||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignRepository,
|
||||
RunRepository,
|
||||
ScenarioRepository,
|
||||
@ -22,6 +25,77 @@ from agenteval.storage.repository import (
|
||||
)
|
||||
|
||||
|
||||
def load_comparison_view(session: Session, campaign: Campaign) -> dict[str, Any]:
|
||||
"""周期对比读模型单一出口:返回 GET /comparison 完整响应形状。
|
||||
|
||||
无行时 status=none + auto_baseline;有行时含 comparison dict(含 model_name
|
||||
标签)+ 对生效基线的 metric_diff。markdown 导出从同一 view 投影。
|
||||
"""
|
||||
payload = _build_auto_baseline_payload(campaign, session)
|
||||
row = CampaignPeriodComparisonRepository(session).get_by_campaign(campaign.id)
|
||||
if row is None:
|
||||
return {"status": "none", "comparison": None, **payload}
|
||||
|
||||
effective_baseline = CampaignRepository(session).get(row.baseline_campaign_id)
|
||||
metric_diff = (
|
||||
compute_metric_diff(
|
||||
load_campaign_report(session, effective_baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
if effective_baseline is not None
|
||||
else None
|
||||
)
|
||||
model_cfg = ModelConfigRepository(session).get(row.model_config_id) if row.model_config_id else None
|
||||
model_label = (
|
||||
f"{model_cfg.name}({model_cfg.model_name})"
|
||||
if model_cfg and model_cfg.model_name
|
||||
else (model_cfg.name if model_cfg else None)
|
||||
)
|
||||
comparison = {
|
||||
"baseline_campaign_id": row.baseline_campaign_id,
|
||||
"baseline": (
|
||||
{
|
||||
"id": effective_baseline.id,
|
||||
"name": effective_baseline.name,
|
||||
"completed_at": iso_utc(effective_baseline.completed_at),
|
||||
}
|
||||
if effective_baseline
|
||||
else None
|
||||
),
|
||||
"result": row.get_result(),
|
||||
"error": row.error,
|
||||
"model_config_id": row.model_config_id,
|
||||
"model_name": model_label,
|
||||
"triggered_by": row.triggered_by,
|
||||
"updated_at": iso_utc(row.updated_at),
|
||||
}
|
||||
return {
|
||||
"status": row.status,
|
||||
"comparison": comparison,
|
||||
"auto_baseline": payload["auto_baseline"],
|
||||
"metric_diff": metric_diff,
|
||||
}
|
||||
|
||||
|
||||
def _build_auto_baseline_payload(campaign: Campaign, session: Session) -> dict[str, Any]:
|
||||
"""自动基线信息 + 机械 diff(无基线时两者均为 null)。"""
|
||||
baseline = resolve_auto_baseline(campaign, session)
|
||||
if baseline is None:
|
||||
return {"auto_baseline": None, "metric_diff": None}
|
||||
diff = compute_metric_diff(
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
return {
|
||||
"auto_baseline": {
|
||||
"id": baseline.id,
|
||||
"name": baseline.name,
|
||||
"completed_at": iso_utc(baseline.completed_at),
|
||||
},
|
||||
"metric_diff": diff,
|
||||
}
|
||||
|
||||
|
||||
class CampaignReadModel:
|
||||
"""One interface for Campaign list, detail, report and export projections."""
|
||||
|
||||
|
||||
@ -21,7 +21,6 @@ from typing import Awaitable, Callable, Optional
|
||||
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.analysis import resolve_analysis_model
|
||||
from agenteval.evaluation.campaign_lifecycle import complete_campaign
|
||||
from agenteval.evaluation.campaign_lifecycle import start_campaign as start_campaign_lifecycle
|
||||
from agenteval.evaluation.campaign_scheduler import (
|
||||
@ -33,7 +32,7 @@ from agenteval.evaluation.campaign_scheduler import (
|
||||
resolve_finalize,
|
||||
)
|
||||
from agenteval.evaluation.engine import EvalEngine
|
||||
from agenteval.evaluation.intelligence_jobs import enqueue_campaign_analysis
|
||||
from agenteval.evaluation.intelligence_jobs import auto_intelligence_eligible, enqueue_campaign_analysis
|
||||
from agenteval.models import (
|
||||
Campaign,
|
||||
CampaignStatus,
|
||||
@ -239,11 +238,9 @@ def _auto_start_analysis(campaign: Campaign, session: Session) -> None:
|
||||
failure all skip silently — the analysis is an enhancement and must never
|
||||
block or break campaign completion.
|
||||
"""
|
||||
if campaign.time_scale != 1 or not campaign.id:
|
||||
if not auto_intelligence_eligible(campaign, session):
|
||||
return
|
||||
try:
|
||||
if resolve_analysis_model(campaign, session) is None:
|
||||
return
|
||||
enqueue_campaign_analysis(campaign.id, triggered_by="auto")
|
||||
except Exception as exc:
|
||||
_logger.warning("活动 %s 自动分析触发失败(已跳过): %s", campaign.id, exc)
|
||||
|
||||
@ -13,15 +13,10 @@ from typing import Any, Optional
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.analysis import ChatClient, resolve_analysis_model
|
||||
from agenteval.evaluation.report import load_campaign_report
|
||||
from agenteval.models import Campaign, CampaignStatus
|
||||
from agenteval.storage.db import iso_utc, utc_now
|
||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignPeriodComparisonRepository,
|
||||
CampaignRepository,
|
||||
)
|
||||
from agenteval.storage.async_job_repository import CampaignAnalysisRepository
|
||||
from agenteval.storage.db import utc_now
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
from agenteval.utils.llm import parse_json_from_llm_text
|
||||
|
||||
_SAME_MOMENT_EPS = 1e-3
|
||||
@ -177,9 +172,7 @@ def validate_comparison_request(
|
||||
else:
|
||||
baseline = resolve_auto_baseline(campaign, session)
|
||||
if baseline is None:
|
||||
raise ComparisonError(
|
||||
"未找到自动基线:历史活动中没有同计划指纹且已完成分析的活动,可手动选择基线活动"
|
||||
)
|
||||
raise ComparisonError("未找到自动基线:历史活动中没有同计划指纹且已完成分析的活动,可手动选择基线活动")
|
||||
|
||||
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
||||
if baseline_analysis is None or baseline_analysis.status != "completed":
|
||||
@ -188,98 +181,6 @@ def validate_comparison_request(
|
||||
return baseline
|
||||
|
||||
|
||||
def load_comparison_view(session: Session, campaign: Campaign) -> dict[str, Any]:
|
||||
"""周期对比读模型单一出口:返回 GET /comparison 完整响应形状。
|
||||
|
||||
无行时 status=none + auto_baseline;有行时含 comparison dict(含 model_name
|
||||
标签)+ 对生效基线的 metric_diff。markdown 导出从同一 view 投影。
|
||||
"""
|
||||
payload = _build_auto_baseline_payload(campaign, session)
|
||||
row = CampaignPeriodComparisonRepository(session).get_by_campaign(campaign.id)
|
||||
if row is None:
|
||||
return {"status": "none", "comparison": None, **payload}
|
||||
|
||||
effective_baseline = CampaignRepository(session).get(row.baseline_campaign_id)
|
||||
metric_diff = (
|
||||
compute_metric_diff(
|
||||
load_campaign_report(session, effective_baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
if effective_baseline is not None
|
||||
else None
|
||||
)
|
||||
model_cfg = (
|
||||
ModelConfigRepository(session).get(row.model_config_id) if row.model_config_id else None
|
||||
)
|
||||
model_label = (
|
||||
f"{model_cfg.name}({model_cfg.model_name})"
|
||||
if model_cfg and model_cfg.model_name
|
||||
else (model_cfg.name if model_cfg else None)
|
||||
)
|
||||
comparison = {
|
||||
"baseline_campaign_id": row.baseline_campaign_id,
|
||||
"baseline": (
|
||||
{
|
||||
"id": effective_baseline.id,
|
||||
"name": effective_baseline.name,
|
||||
"completed_at": iso_utc(effective_baseline.completed_at),
|
||||
}
|
||||
if effective_baseline
|
||||
else None
|
||||
),
|
||||
"result": row.get_result(),
|
||||
"error": row.error,
|
||||
"model_config_id": row.model_config_id,
|
||||
"model_name": model_label,
|
||||
"triggered_by": row.triggered_by,
|
||||
"updated_at": iso_utc(row.updated_at),
|
||||
}
|
||||
return {
|
||||
"status": row.status,
|
||||
"comparison": comparison,
|
||||
"auto_baseline": payload["auto_baseline"],
|
||||
"metric_diff": metric_diff,
|
||||
}
|
||||
|
||||
|
||||
def _build_auto_baseline_payload(campaign: Campaign, session: Session) -> dict[str, Any]:
|
||||
"""自动基线信息 + 机械 diff(无基线时两者均为 null)。"""
|
||||
baseline = resolve_auto_baseline(campaign, session)
|
||||
if baseline is None:
|
||||
return {"auto_baseline": None, "metric_diff": None}
|
||||
diff = compute_metric_diff(
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
return {
|
||||
"auto_baseline": {
|
||||
"id": baseline.id,
|
||||
"name": baseline.name,
|
||||
"completed_at": iso_utc(baseline.completed_at),
|
||||
},
|
||||
"metric_diff": diff,
|
||||
}
|
||||
|
||||
|
||||
def build_comparison_payload(campaign: Campaign, session: Session) -> dict[str, Any]:
|
||||
"""GET 返回体:自动基线信息 + 机械 diff(无基线时两者均为 null)。"""
|
||||
baseline = resolve_auto_baseline(campaign, session)
|
||||
if baseline is None:
|
||||
return {"auto_baseline": None, "metric_diff": None}
|
||||
diff = compute_metric_diff(
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
return {
|
||||
"auto_baseline": {
|
||||
"id": baseline.id,
|
||||
"name": baseline.name,
|
||||
"completed_at": iso_utc(baseline.completed_at),
|
||||
},
|
||||
"metric_diff": diff,
|
||||
}
|
||||
|
||||
|
||||
# ── 叙述半边:单次 LLM 调用编排 ─────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
@ -1,26 +1,40 @@
|
||||
"""Durable runtime for Campaign intelligence jobs.
|
||||
|
||||
智能分析与周期对比是两个领域工作 adapter;本 module 统一掌握它们的
|
||||
持久排队、进程内幂等启动、重启恢复和关闭顺序。数据库行是耐久权威,
|
||||
TaskRegistry 只保存当前进程中的任务句柄。
|
||||
智能分析与周期对比是两个领域工作 adapter;本 module 的 ``execute``
|
||||
统一掌握它们的结算契约:建行/认领/校验/generating/failed/completed。
|
||||
数据库行是耐久权威,TaskRegistry 只保存当前进程中的任务句柄。
|
||||
"""
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, Optional
|
||||
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.storage.async_job_repository import CampaignAnalysisRepository, CampaignPeriodComparisonRepository
|
||||
from agenteval.storage.db import get_session
|
||||
from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignPeriodComparisonRepository,
|
||||
)
|
||||
from agenteval.task_registry import TaskRegistry
|
||||
|
||||
_registry = TaskRegistry()
|
||||
_logger = logging.getLogger("agenteval")
|
||||
MAX_QUEUED_RECOVERY_ATTEMPTS = 3
|
||||
_ERROR_TRUNCATE_LEN = 500
|
||||
|
||||
# adapter 契约:
|
||||
# ensure_queued(session) -> row | None(建行或确认已有行;None = 静默放弃)
|
||||
# validate(session) -> 落账附加字段(认领后校验,失败抛 JobValidationError)
|
||||
# work_fn(session) -> (result, 落账附加字段)
|
||||
EnsureQueuedFn = Callable[[Session], Any]
|
||||
ValidateFn = Callable[[Session], Optional[dict]]
|
||||
WorkFn = Callable[[Session], Awaitable[tuple[dict, dict]]]
|
||||
|
||||
|
||||
class JobValidationError(RuntimeError):
|
||||
"""认领后校验失败:携带已确定的落账字段(如对比的生效基线)进 failed 行。"""
|
||||
|
||||
def __init__(self, message: str, **meta: Any):
|
||||
super().__init__(message)
|
||||
self.meta = meta
|
||||
|
||||
|
||||
def _job_key(kind: str, campaign_id: str) -> str:
|
||||
@ -50,6 +64,81 @@ def _launch_comparison(
|
||||
_registry.launch(_job_key("comparison", campaign_id), run)
|
||||
|
||||
|
||||
async def execute(
|
||||
job_kind: str,
|
||||
campaign_id: str,
|
||||
*,
|
||||
triggered_by: str,
|
||||
repo_cls,
|
||||
ensure_queued: EnsureQueuedFn,
|
||||
validate: Optional[ValidateFn] = None,
|
||||
work_fn: WorkFn,
|
||||
session_factory: Optional[Callable[[], Session]] = None,
|
||||
) -> None:
|
||||
"""结算一个耐久智能作业:建行 → 认领 → 校验 → generating → 落账。
|
||||
|
||||
结算契约(两个 adapter 共用,只此一处):
|
||||
|
||||
- ``ensure_queued`` 幂等建行(或确认已有行);返回 None 表示静默
|
||||
放弃(如活动不存在),不写任何行。
|
||||
- ``claim_queued`` 是耐久幂等权威:抢不到(重复触发 / 已在执行)
|
||||
静默返回,绝不重复跑领域工作。
|
||||
- ``validate`` 在认领后校验并产出附加字段(模型、基线),随
|
||||
generating/结算落账;抛 ``JobValidationError`` → 落 failed 行
|
||||
(异常携带的已确定字段一并落账)。
|
||||
- ``work_fn`` 抛任何异常 → failed 行,error 截断到 500 字符(一处)。
|
||||
"""
|
||||
session = (session_factory or get_session)()
|
||||
repo = repo_cls(session)
|
||||
try:
|
||||
row = ensure_queued(session)
|
||||
if row is None:
|
||||
return
|
||||
if not repo.claim_queued(campaign_id).claimed:
|
||||
return
|
||||
effective_trigger = row.triggered_by or triggered_by
|
||||
|
||||
meta: dict[str, Any] = {}
|
||||
if validate is not None:
|
||||
try:
|
||||
meta = validate(session) or {}
|
||||
except JobValidationError as exc:
|
||||
repo.upsert(
|
||||
campaign_id,
|
||||
status="failed",
|
||||
triggered_by=effective_trigger,
|
||||
error=str(exc),
|
||||
**exc.meta,
|
||||
)
|
||||
return
|
||||
|
||||
repo.upsert(campaign_id, status="generating", triggered_by=effective_trigger, **meta)
|
||||
work_meta: dict[str, Any] = {}
|
||||
try:
|
||||
result, work_meta = await work_fn(session)
|
||||
except Exception as exc:
|
||||
_logger.warning("%s 作业 %s 失败: %s", job_kind, campaign_id, exc)
|
||||
repo.upsert(
|
||||
campaign_id,
|
||||
status="failed",
|
||||
error=str(exc)[:_ERROR_TRUNCATE_LEN],
|
||||
triggered_by=effective_trigger,
|
||||
**meta,
|
||||
**work_meta,
|
||||
)
|
||||
return
|
||||
repo.upsert(
|
||||
campaign_id,
|
||||
status="completed",
|
||||
result=result,
|
||||
triggered_by=effective_trigger,
|
||||
**meta,
|
||||
**work_meta,
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
async def execute_campaign_analysis_job(
|
||||
campaign_id: str,
|
||||
*,
|
||||
@ -57,53 +146,38 @@ async def execute_campaign_analysis_job(
|
||||
chat_client: Any = None,
|
||||
session_factory: Optional[Callable[[], Session]] = None,
|
||||
) -> None:
|
||||
"""Claim and settle one intelligent-analysis job."""
|
||||
"""Claim and settle one intelligent-analysis job (work adapter)."""
|
||||
|
||||
def ensure_queued(session: Session):
|
||||
analyses = CampaignAnalysisRepository(session)
|
||||
row = analyses.get_by_campaign(campaign_id)
|
||||
return row or analyses.enqueue(campaign_id, triggered_by=triggered_by)
|
||||
|
||||
def validate(session: Session) -> dict:
|
||||
from agenteval.evaluation.analysis import AnalysisError, validate_analysis_request
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if campaign is None:
|
||||
raise JobValidationError("campaign not found")
|
||||
try:
|
||||
runtime = validate_analysis_request(session, campaign)
|
||||
except AnalysisError as exc:
|
||||
raise JobValidationError(str(exc)) from exc
|
||||
return {"model_config_id": runtime.id}
|
||||
|
||||
async def work(session: Session) -> tuple[dict, dict]:
|
||||
from agenteval.evaluation.analysis import (
|
||||
analyze_campaign,
|
||||
collect_failure_samples,
|
||||
gateway_chat_client,
|
||||
resolve_analysis_model,
|
||||
)
|
||||
from agenteval.evaluation.comparison import resolve_auto_baseline
|
||||
from agenteval.evaluation.report import load_campaign_view
|
||||
from agenteval.storage.repository import CampaignRepository, RunRepository
|
||||
|
||||
session = (session_factory or get_session)()
|
||||
try:
|
||||
analyses = CampaignAnalysisRepository(session)
|
||||
row = analyses.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = analyses.enqueue(campaign_id, triggered_by=triggered_by)
|
||||
if not analyses.claim_queued(campaign_id).claimed:
|
||||
return
|
||||
effective_trigger = row.triggered_by or triggered_by
|
||||
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if campaign is None:
|
||||
analyses.upsert(
|
||||
campaign_id,
|
||||
status="failed",
|
||||
triggered_by=effective_trigger,
|
||||
error="campaign not found",
|
||||
)
|
||||
return
|
||||
runtime = resolve_analysis_model(campaign, session)
|
||||
if runtime is None:
|
||||
analyses.upsert(
|
||||
campaign_id,
|
||||
status="failed",
|
||||
triggered_by=effective_trigger,
|
||||
error="未配置分析模型:请在模型配置中心将某个 chat 配置设为「分析默认」",
|
||||
)
|
||||
return
|
||||
analyses.upsert(
|
||||
campaign_id,
|
||||
status="generating",
|
||||
model_config_id=runtime.id,
|
||||
triggered_by=effective_trigger,
|
||||
)
|
||||
try:
|
||||
client = chat_client or gateway_chat_client(runtime)
|
||||
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||||
view = load_campaign_view(session, campaign)
|
||||
result = await analyze_campaign(
|
||||
@ -111,29 +185,51 @@ async def execute_campaign_analysis_job(
|
||||
report=view["report"],
|
||||
failure_samples=collect_failure_samples(campaign_id, session),
|
||||
valid_run_ids={run.id for run in runs if run.id},
|
||||
chat_client=client,
|
||||
chat_client=chat_client or gateway_chat_client(runtime),
|
||||
exploration_summary=view["exploration"],
|
||||
)
|
||||
except Exception as exc:
|
||||
_logger.warning("活动 %s 智能分析失败: %s", campaign_id, exc)
|
||||
analyses.upsert(
|
||||
campaign_id,
|
||||
status="failed",
|
||||
model_config_id=runtime.id,
|
||||
error=str(exc)[:500],
|
||||
triggered_by=effective_trigger,
|
||||
)
|
||||
return
|
||||
analyses.upsert(
|
||||
campaign_id,
|
||||
status="completed",
|
||||
result=result,
|
||||
model_config_id=runtime.id,
|
||||
triggered_by=effective_trigger,
|
||||
)
|
||||
return result, {}
|
||||
|
||||
await execute(
|
||||
"analysis",
|
||||
campaign_id,
|
||||
triggered_by=triggered_by,
|
||||
repo_cls=CampaignAnalysisRepository,
|
||||
ensure_queued=ensure_queued,
|
||||
validate=validate,
|
||||
work_fn=work,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
await _maybe_enqueue_auto_comparison(campaign_id, session_factory=session_factory)
|
||||
|
||||
|
||||
def auto_intelligence_eligible(campaign, session: Session) -> bool:
|
||||
"""自动触发智能分析/周期对比的唯一跳过守卫:正式线 + 分析模型可解析。
|
||||
|
||||
活动完成时的自动分析与分析完成后的自动对比共用这一判断点;
|
||||
加速线与未配置模型一律静默跳过(增强能力,不阻断主流程)。
|
||||
"""
|
||||
from agenteval.evaluation.analysis import resolve_analysis_model
|
||||
|
||||
return bool(campaign.id) and campaign.time_scale == 1 and resolve_analysis_model(campaign, session) is not None
|
||||
|
||||
|
||||
async def _maybe_enqueue_auto_comparison(
|
||||
campaign_id: str,
|
||||
*,
|
||||
session_factory: Optional[Callable[[], Session]] = None,
|
||||
) -> None:
|
||||
"""分析完成后:存在自动基线 → 入队周期对比。任何一步不满足或出错都静默跳过。"""
|
||||
from agenteval.evaluation.comparison import resolve_auto_baseline
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
|
||||
session = (session_factory or get_session)()
|
||||
try:
|
||||
if campaign.time_scale != 1 or resolve_analysis_model(campaign, session) is None:
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if campaign is None or not auto_intelligence_eligible(campaign, session):
|
||||
return
|
||||
row = CampaignAnalysisRepository(session).get_by_campaign(campaign_id)
|
||||
if row is None or row.status != "completed":
|
||||
return
|
||||
baseline = resolve_auto_baseline(campaign, session)
|
||||
if baseline is None:
|
||||
@ -145,7 +241,7 @@ async def execute_campaign_analysis_job(
|
||||
session=session,
|
||||
)
|
||||
except Exception as exc:
|
||||
_logger.warning("活动 %s 自动周期对比跳过: %s", campaign.id, exc)
|
||||
_logger.warning("活动 %s 自动周期对比跳过: %s", campaign_id, exc)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
@ -158,26 +254,20 @@ async def execute_campaign_comparison_job(
|
||||
chat_client: Any = None,
|
||||
session_factory: Optional[Callable[[], Session]] = None,
|
||||
) -> None:
|
||||
"""Claim and settle one period-comparison job."""
|
||||
from agenteval.evaluation.analysis import gateway_chat_client, resolve_analysis_model
|
||||
from agenteval.evaluation.comparison import (
|
||||
ComparisonError,
|
||||
compute_metric_diff,
|
||||
narrate_period_comparison,
|
||||
validate_comparison_request,
|
||||
)
|
||||
from agenteval.evaluation.report import load_campaign_report
|
||||
"""Claim and settle one period-comparison job (work adapter)."""
|
||||
from agenteval.evaluation.comparison import ComparisonError
|
||||
|
||||
def ensure_queued(session: Session):
|
||||
from agenteval.evaluation.comparison import validate_comparison_request
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
|
||||
session = (session_factory or get_session)()
|
||||
try:
|
||||
comparisons = CampaignPeriodComparisonRepository(session)
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if campaign is None:
|
||||
return
|
||||
|
||||
return None
|
||||
row = comparisons.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
if row is not None:
|
||||
return row
|
||||
try:
|
||||
initial_baseline = validate_comparison_request(
|
||||
session,
|
||||
@ -191,17 +281,24 @@ async def execute_campaign_comparison_job(
|
||||
triggered_by=triggered_by,
|
||||
error=str(exc),
|
||||
)
|
||||
return
|
||||
row = comparisons.enqueue(
|
||||
return None
|
||||
return comparisons.enqueue(
|
||||
campaign_id,
|
||||
baseline_campaign_id=initial_baseline.id,
|
||||
triggered_by=triggered_by,
|
||||
)
|
||||
if not comparisons.claim_queued(campaign_id).claimed:
|
||||
return
|
||||
|
||||
effective_trigger = row.triggered_by or triggered_by
|
||||
effective_baseline_id = baseline_campaign_id or row.baseline_campaign_id
|
||||
def validate(session: Session) -> dict:
|
||||
from agenteval.evaluation.analysis import resolve_analysis_model
|
||||
from agenteval.evaluation.comparison import (
|
||||
ComparisonError,
|
||||
validate_comparison_request,
|
||||
)
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
row = CampaignPeriodComparisonRepository(session).get_by_campaign(campaign_id)
|
||||
effective_baseline_id = baseline_campaign_id or (row.baseline_campaign_id if row else None)
|
||||
try:
|
||||
baseline = validate_comparison_request(
|
||||
session,
|
||||
@ -209,26 +306,23 @@ async def execute_campaign_comparison_job(
|
||||
explicit_baseline_id=effective_baseline_id,
|
||||
)
|
||||
except ComparisonError as exc:
|
||||
comparisons.upsert(
|
||||
campaign_id,
|
||||
status="failed",
|
||||
baseline_campaign_id=effective_baseline_id,
|
||||
triggered_by=effective_trigger,
|
||||
error=str(exc),
|
||||
)
|
||||
return
|
||||
raise JobValidationError(str(exc), baseline_campaign_id=effective_baseline_id) from exc
|
||||
runtime = resolve_analysis_model(campaign, session)
|
||||
return {"baseline_campaign_id": baseline.id, "model_config_id": runtime.id}
|
||||
|
||||
async def work(session: Session) -> tuple[dict, dict]:
|
||||
from agenteval.evaluation.analysis import gateway_chat_client, resolve_analysis_model
|
||||
from agenteval.evaluation.comparison import compute_metric_diff, narrate_period_comparison
|
||||
from agenteval.evaluation.report import load_campaign_report
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
row = CampaignPeriodComparisonRepository(session).get_by_campaign(campaign_id)
|
||||
effective_baseline_id = baseline_campaign_id or (row.baseline_campaign_id if row else None)
|
||||
baseline = CampaignRepository(session).get(effective_baseline_id)
|
||||
runtime = resolve_analysis_model(campaign, session)
|
||||
baseline_analysis = CampaignAnalysisRepository(session).get_by_campaign(baseline.id)
|
||||
current_analysis = CampaignAnalysisRepository(session).get_by_campaign(campaign_id)
|
||||
comparisons.upsert(
|
||||
campaign_id,
|
||||
status="generating",
|
||||
baseline_campaign_id=baseline.id,
|
||||
model_config_id=runtime.id,
|
||||
triggered_by=effective_trigger,
|
||||
)
|
||||
try:
|
||||
diff = compute_metric_diff(
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
@ -240,27 +334,18 @@ async def execute_campaign_comparison_job(
|
||||
valid_scenario_ids={item["scenario_id"] for item in diff["scenarios"]},
|
||||
chat_client=chat_client or gateway_chat_client(runtime),
|
||||
)
|
||||
except Exception as exc:
|
||||
_logger.warning("活动 %s 周期对比失败: %s", campaign_id, exc)
|
||||
comparisons.upsert(
|
||||
return result, {}
|
||||
|
||||
await execute(
|
||||
"comparison",
|
||||
campaign_id,
|
||||
status="failed",
|
||||
baseline_campaign_id=baseline.id,
|
||||
model_config_id=runtime.id,
|
||||
error=str(exc)[:500],
|
||||
triggered_by=effective_trigger,
|
||||
triggered_by=triggered_by,
|
||||
repo_cls=CampaignPeriodComparisonRepository,
|
||||
ensure_queued=ensure_queued,
|
||||
validate=validate,
|
||||
work_fn=work,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
return
|
||||
comparisons.upsert(
|
||||
campaign_id,
|
||||
status="completed",
|
||||
baseline_campaign_id=baseline.id,
|
||||
result=result,
|
||||
model_config_id=runtime.id,
|
||||
triggered_by=effective_trigger,
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def enqueue_campaign_analysis(
|
||||
|
||||
@ -22,11 +22,8 @@ from agenteval.exploration.models import ExplorationMessage, ExplorationSession
|
||||
from agenteval.models import Campaign
|
||||
from agenteval.services.model_configs import ModelRuntimeConfig
|
||||
from agenteval.storage.db import get_session, iso_utc, utc_now
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationMessageRepository,
|
||||
ExplorationSessionRepository,
|
||||
)
|
||||
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
from agenteval.task_registry import TaskRegistry
|
||||
from agenteval.utils.llm import parse_json_from_llm_text
|
||||
|
||||
@ -72,9 +69,7 @@ def sample_round_indexes(rounds: list[int], limit: int = MAX_JUDGE_SAMPLES) -> l
|
||||
return [rounds[round(i * span / (limit - 1))] for i in range(limit)]
|
||||
|
||||
|
||||
def build_judge_messages(
|
||||
session_obj: ExplorationSession, samples: list[ExplorationMessage]
|
||||
) -> list[dict[str, str]]:
|
||||
def build_judge_messages(session_obj: ExplorationSession, samples: list[ExplorationMessage]) -> list[dict[str, str]]:
|
||||
transcript = []
|
||||
for message in samples:
|
||||
speaker = "虚拟用户" if message.role == "user" else "被评对象"
|
||||
|
||||
@ -30,10 +30,9 @@ from agenteval.exploration.models import (
|
||||
)
|
||||
from agenteval.models import CampaignStatus
|
||||
from agenteval.storage.db import as_utc, utc_now
|
||||
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationMessageRepository,
|
||||
ExplorationSessionRepository,
|
||||
TargetRepository,
|
||||
)
|
||||
|
||||
@ -144,9 +143,7 @@ async def conduct_turn(db_session: Session, *, session_id: str, content: str) ->
|
||||
|
||||
received_at = utc_now()
|
||||
latency_ms = (
|
||||
outcome.latency_ms
|
||||
if outcome.latency_ms is not None
|
||||
else int((received_at - sent_at).total_seconds() * 1000)
|
||||
outcome.latency_ms if outcome.latency_ms is not None else int((received_at - sent_at).total_seconds() * 1000)
|
||||
)
|
||||
reply_text = outcome.reply_text or ""
|
||||
message_repo.save_message(
|
||||
|
||||
@ -14,9 +14,9 @@ from agenteval.evaluation.report import generate_campaign_report
|
||||
from agenteval.exploration.models import resolve_budget
|
||||
from agenteval.models import Campaign, CampaignStatus, EvalRun
|
||||
from agenteval.storage.db import as_utc, iso_utc, utc_now
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationSessionRepository,
|
||||
RunRepository,
|
||||
ScenarioRepository,
|
||||
TargetRepository,
|
||||
|
||||
@ -52,7 +52,7 @@ def summarize_exploration(sessions: list[ExplorationSession]) -> Optional[dict[s
|
||||
|
||||
def summarize_campaign_exploration(db_session: Session, campaign_id: str) -> Optional[dict[str, Any]]:
|
||||
"""取数 + 聚合一步完成:报告 / 分析 / 导出三个出口共用的探索摘要取法。"""
|
||||
from agenteval.storage.repository import ExplorationSessionRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
|
||||
return summarize_exploration(ExplorationSessionRepository(db_session).list_by_campaign(campaign_id))
|
||||
|
||||
|
||||
247
backend/agenteval/storage/async_job_repository.py
Normal file
247
backend/agenteval/storage/async_job_repository.py
Normal file
@ -0,0 +1,247 @@
|
||||
"""Repositories for durable async job rows (活动分析 / 周期对比).
|
||||
|
||||
DB 行是任务的持久权威:queued/generating/completed/failed 状态机与
|
||||
claim 竞争都住在这里;进程内任务句柄另由 TaskRegistry 管理。
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Generic, Optional, TypeVar
|
||||
|
||||
from sqlalchemy import update as sql_update
|
||||
from sqlmodel import Session, select
|
||||
|
||||
from agenteval.storage.db import CampaignAnalysisDB, CampaignPeriodComparisonDB, get_session, utc_now
|
||||
|
||||
DB = TypeVar("DB") # persisted table row
|
||||
|
||||
|
||||
class AsyncJobClaimStatus(str, Enum):
|
||||
CLAIMED = "claimed"
|
||||
NOT_FOUND = "not_found"
|
||||
ALREADY_CLAIMED = "already_claimed"
|
||||
NOT_QUEUED = "not_queued"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AsyncJobClaimResult:
|
||||
status: AsyncJobClaimStatus
|
||||
|
||||
@property
|
||||
def claimed(self) -> bool:
|
||||
return self.status is AsyncJobClaimStatus.CLAIMED
|
||||
|
||||
|
||||
class AsyncJobRepository(Generic[DB]):
|
||||
"""Base class for async job repositories (analysis, comparison).
|
||||
|
||||
Provides common pattern: get_by_campaign, mark_orphans_failed.
|
||||
Subclasses implement upsert with their specific fields.
|
||||
"""
|
||||
|
||||
_table: type[DB]
|
||||
|
||||
def __init__(self, session: Optional[Session] = None):
|
||||
self.session = session or get_session()
|
||||
|
||||
def get_by_campaign(self, campaign_id: str) -> Optional[DB]:
|
||||
statement = select(self._table).where(self._table.campaign_id == campaign_id) # type: ignore[attr-defined]
|
||||
return self.session.exec(statement).first()
|
||||
|
||||
def list_queued(self) -> list[DB]:
|
||||
statement = select(self._table).where(self._table.status == "queued") # type: ignore[attr-defined]
|
||||
return list(self.session.exec(statement).all())
|
||||
|
||||
def claim_queued(self, campaign_id: str) -> AsyncJobClaimResult:
|
||||
"""Atomically move one queued job to generating.
|
||||
|
||||
The status predicate is the durable idempotency authority. Competing
|
||||
workers may observe the same queued row, but only one can claim it.
|
||||
"""
|
||||
statement = (
|
||||
sql_update(self._table)
|
||||
.where(
|
||||
self._table.campaign_id == campaign_id, # type: ignore[attr-defined]
|
||||
self._table.status == "queued", # type: ignore[attr-defined]
|
||||
)
|
||||
.values(status="generating", error=None, updated_at=utc_now())
|
||||
)
|
||||
result = self.session.exec(statement)
|
||||
self.session.commit()
|
||||
self.session.expire_all()
|
||||
if result.rowcount == 1:
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.CLAIMED)
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.NOT_FOUND)
|
||||
if row.status == "generating": # type: ignore[attr-defined]
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.ALREADY_CLAIMED)
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.NOT_QUEUED)
|
||||
|
||||
def prepare_queued_recovery(self, max_attempts: int, exhausted_error: str) -> list[DB]:
|
||||
"""Increment queued recovery attempts and fail exhausted jobs."""
|
||||
recoverable: list[DB] = []
|
||||
rows = self.list_queued()
|
||||
for row in rows:
|
||||
if row.recovery_attempts >= max_attempts: # type: ignore[attr-defined]
|
||||
row.status = "failed" # type: ignore[attr-defined]
|
||||
row.error = exhausted_error # type: ignore[attr-defined]
|
||||
else:
|
||||
row.recovery_attempts += 1 # type: ignore[attr-defined]
|
||||
recoverable.append(row)
|
||||
row.updated_at = utc_now() # type: ignore[attr-defined]
|
||||
self.session.add(row)
|
||||
if rows:
|
||||
self.session.commit()
|
||||
self.session.expire_all()
|
||||
return recoverable
|
||||
|
||||
def mark_orphans_failed(self, error_message: str) -> int:
|
||||
"""服务启动时清理:把滞留的 generating 行标记为 failed。
|
||||
|
||||
异步任务是进程内 asyncio 任务,服务重启后不会恢复;不清理则这些
|
||||
行永远停留在 generating(僵尸状态)。
|
||||
"""
|
||||
rows = self.session.exec(
|
||||
select(self._table).where(self._table.status == "generating") # type: ignore[attr-defined]
|
||||
).all()
|
||||
for row in rows:
|
||||
row.status = "failed" # type: ignore[attr-defined]
|
||||
row.error = error_message # type: ignore[attr-defined]
|
||||
row.updated_at = utc_now() # type: ignore[attr-defined]
|
||||
self.session.add(row)
|
||||
if rows:
|
||||
self.session.commit()
|
||||
return len(rows)
|
||||
|
||||
|
||||
class CampaignAnalysisRepository(AsyncJobRepository[CampaignAnalysisDB]):
|
||||
"""Repository for campaign analysis rows (one per campaign, upserted)."""
|
||||
|
||||
_table = CampaignAnalysisDB
|
||||
|
||||
def enqueue(self, campaign_id: str, *, triggered_by: str = "manual") -> CampaignAnalysisDB:
|
||||
"""Persist an analysis job before launching its process-local task."""
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignAnalysisDB(campaign_id=campaign_id, status="queued", triggered_by=triggered_by)
|
||||
elif row.status in {"queued", "generating"}:
|
||||
return row
|
||||
else:
|
||||
row.status = "queued"
|
||||
row.result = None
|
||||
row.model_config_id = None
|
||||
row.error = None
|
||||
row.triggered_by = triggered_by
|
||||
row.recovery_attempts = 0
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def upsert(
|
||||
self,
|
||||
campaign_id: str,
|
||||
*,
|
||||
status: str,
|
||||
result: Optional[dict] = None,
|
||||
model_config_id: Optional[str] = None,
|
||||
error: Optional[str] = None,
|
||||
triggered_by: str = "manual",
|
||||
) -> CampaignAnalysisDB:
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignAnalysisDB(campaign_id=campaign_id)
|
||||
row.status = status
|
||||
if result is not None:
|
||||
row.set_result(result)
|
||||
else:
|
||||
row.result = None
|
||||
row.model_config_id = model_config_id
|
||||
row.error = error
|
||||
row.triggered_by = triggered_by
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def mark_orphans_failed(self) -> int:
|
||||
"""服务启动时清理:把滞留的 generating 分析行标记为 failed。"""
|
||||
return super().mark_orphans_failed("服务重启导致分析生成中断")
|
||||
|
||||
|
||||
class CampaignPeriodComparisonRepository(AsyncJobRepository[CampaignPeriodComparisonDB]):
|
||||
"""Repository for period-comparison rows (one per campaign, upserted)."""
|
||||
|
||||
_table = CampaignPeriodComparisonDB
|
||||
|
||||
def enqueue(
|
||||
self,
|
||||
campaign_id: str,
|
||||
*,
|
||||
baseline_campaign_id: str,
|
||||
triggered_by: str = "manual",
|
||||
) -> CampaignPeriodComparisonDB:
|
||||
"""Persist a comparison job before launching its process-local task."""
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignPeriodComparisonDB(
|
||||
campaign_id=campaign_id,
|
||||
baseline_campaign_id=baseline_campaign_id,
|
||||
status="queued",
|
||||
triggered_by=triggered_by,
|
||||
)
|
||||
elif row.status in {"queued", "generating"}:
|
||||
return row
|
||||
else:
|
||||
row.baseline_campaign_id = baseline_campaign_id
|
||||
row.status = "queued"
|
||||
row.result = None
|
||||
row.model_config_id = None
|
||||
row.error = None
|
||||
row.triggered_by = triggered_by
|
||||
row.recovery_attempts = 0
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def upsert(
|
||||
self,
|
||||
campaign_id: str,
|
||||
*,
|
||||
status: str,
|
||||
baseline_campaign_id: Optional[str] = None,
|
||||
result: Optional[dict] = None,
|
||||
model_config_id: Optional[str] = None,
|
||||
error: Optional[str] = None,
|
||||
triggered_by: str = "manual",
|
||||
) -> CampaignPeriodComparisonDB:
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignPeriodComparisonDB(
|
||||
campaign_id=campaign_id,
|
||||
baseline_campaign_id=baseline_campaign_id or "",
|
||||
)
|
||||
row.status = status
|
||||
if baseline_campaign_id is not None:
|
||||
row.baseline_campaign_id = baseline_campaign_id
|
||||
if result is not None:
|
||||
row.set_result(result)
|
||||
else:
|
||||
row.result = None
|
||||
row.model_config_id = model_config_id
|
||||
row.error = error
|
||||
row.triggered_by = triggered_by
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def mark_orphans_failed(self) -> int:
|
||||
"""服务启动时清理:把滞留的 generating 周期对比行标记为 failed。"""
|
||||
return super().mark_orphans_failed("服务重启导致周期对比生成中断")
|
||||
135
backend/agenteval/storage/exploration_repository.py
Normal file
135
backend/agenteval/storage/exploration_repository.py
Normal file
@ -0,0 +1,135 @@
|
||||
"""Repositories for virtual-user exploration sessions (探索式评测)."""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from sqlmodel import Session, select
|
||||
|
||||
from agenteval.exploration.models import ExplorationMessage, ExplorationSession, ExplorationSessionStatus
|
||||
from agenteval.storage.db import ExplorationMessageDB, ExplorationSessionDB, get_session, utc_now
|
||||
from agenteval.storage.repository import BaseRepository
|
||||
|
||||
|
||||
class ExplorationSessionRepository(BaseRepository[ExplorationSession, ExplorationSessionDB]):
|
||||
"""Repository for virtual-user exploration sessions (探索会话)."""
|
||||
|
||||
_table = ExplorationSessionDB
|
||||
_order_by = "created_at"
|
||||
|
||||
def _copy_mutable(self, db: ExplorationSessionDB, session_obj: ExplorationSession) -> None:
|
||||
db.goal = session_obj.goal
|
||||
db.status = session_obj.status.value
|
||||
db.triggered_by = session_obj.triggered_by.value
|
||||
db.turn_count = session_obj.turn_count
|
||||
db.error = session_obj.error
|
||||
db.closed_at = session_obj.closed_at
|
||||
db.set_persona(session_obj.persona)
|
||||
if session_obj.seed_ref is not None:
|
||||
db.set_seed_ref(session_obj.seed_ref)
|
||||
if session_obj.experience is not None:
|
||||
db.set_experience(session_obj.experience)
|
||||
if session_obj.judge_review is not None:
|
||||
db.set_judge_review(session_obj.judge_review)
|
||||
|
||||
def _to_db(self, session_obj: ExplorationSession) -> ExplorationSessionDB:
|
||||
db = ExplorationSessionDB(
|
||||
id=session_obj.id,
|
||||
campaign_id=session_obj.campaign_id,
|
||||
target_id=session_obj.target_id,
|
||||
created_at=session_obj.created_at,
|
||||
)
|
||||
self._copy_mutable(db, session_obj)
|
||||
return db
|
||||
|
||||
def _from_db(self, db: ExplorationSessionDB) -> ExplorationSession:
|
||||
return ExplorationSession(
|
||||
id=db.id,
|
||||
campaign_id=db.campaign_id,
|
||||
target_id=db.target_id,
|
||||
persona=db.get_persona(),
|
||||
goal=db.goal,
|
||||
seed_ref=db.get_seed_ref(),
|
||||
status=db.status,
|
||||
triggered_by=db.triggered_by,
|
||||
experience=db.get_experience(),
|
||||
judge_review=db.get_judge_review(),
|
||||
turn_count=db.turn_count,
|
||||
error=db.error,
|
||||
created_at=db.created_at,
|
||||
closed_at=db.closed_at,
|
||||
)
|
||||
|
||||
def update(self, session_obj: ExplorationSession) -> Optional[ExplorationSession]:
|
||||
existing = self.session.get(ExplorationSessionDB, session_obj.id)
|
||||
if not existing:
|
||||
return None
|
||||
self._copy_mutable(existing, session_obj)
|
||||
self.session.add(existing)
|
||||
self.session.commit()
|
||||
self.session.refresh(existing)
|
||||
return self._from_db(existing)
|
||||
|
||||
def list_by_campaign(self, campaign_id: str) -> list[ExplorationSession]:
|
||||
statement = (
|
||||
select(ExplorationSessionDB)
|
||||
.where(ExplorationSessionDB.campaign_id == campaign_id)
|
||||
.order_by(ExplorationSessionDB.created_at)
|
||||
)
|
||||
return [self._from_db(r) for r in self.session.exec(statement).all()]
|
||||
|
||||
def expire_running_sessions(self, campaign_id: str) -> int:
|
||||
"""Expire every still-running exploration session of a finalized campaign.
|
||||
|
||||
Returns the number of sessions expired. Completed/failed sessions keep
|
||||
their evidence untouched.
|
||||
"""
|
||||
expired = 0
|
||||
for session_obj in self.list_by_campaign(campaign_id):
|
||||
if session_obj.status != ExplorationSessionStatus.RUNNING:
|
||||
continue
|
||||
session_obj.status = ExplorationSessionStatus.EXPIRED
|
||||
session_obj.closed_at = utc_now()
|
||||
self.update(session_obj)
|
||||
expired += 1
|
||||
return expired
|
||||
|
||||
|
||||
class ExplorationMessageRepository:
|
||||
"""Append-only repository for exploration session chat rows."""
|
||||
|
||||
def __init__(self, session: Optional[Session] = None):
|
||||
self.session = session or get_session()
|
||||
|
||||
def save_message(self, message: ExplorationMessage) -> ExplorationMessage:
|
||||
db = ExplorationMessageDB(
|
||||
id=message.id,
|
||||
session_id=message.session_id,
|
||||
round_index=message.round_index,
|
||||
role=message.role,
|
||||
content=message.content,
|
||||
latency_ms=message.latency_ms,
|
||||
created_at=message.created_at,
|
||||
)
|
||||
self.session.add(db)
|
||||
self.session.commit()
|
||||
self.session.refresh(db)
|
||||
message.id = db.id
|
||||
return message
|
||||
|
||||
def list_by_session(self, session_id: str) -> list[ExplorationMessage]:
|
||||
statement = (
|
||||
select(ExplorationMessageDB)
|
||||
.where(ExplorationMessageDB.session_id == session_id)
|
||||
.order_by(ExplorationMessageDB.created_at)
|
||||
)
|
||||
return [
|
||||
ExplorationMessage(
|
||||
id=r.id,
|
||||
session_id=r.session_id,
|
||||
round_index=r.round_index,
|
||||
role=r.role,
|
||||
content=r.content,
|
||||
latency_ms=r.latency_ms,
|
||||
created_at=r.created_at,
|
||||
)
|
||||
for r in self.session.exec(statement).all()
|
||||
]
|
||||
@ -11,7 +11,7 @@ from sqlalchemy import update as sql_update
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlmodel import Session, select
|
||||
|
||||
from agenteval.exploration.models import ExplorationMessage, ExplorationSession, ExplorationSessionStatus
|
||||
from agenteval.exploration.models import ExplorationSessionStatus
|
||||
from agenteval.models import (
|
||||
Campaign,
|
||||
CampaignStatus,
|
||||
@ -26,13 +26,10 @@ from agenteval.models import (
|
||||
)
|
||||
from agenteval.services.model_configs import ModelConfigService
|
||||
from agenteval.storage.db import (
|
||||
CampaignAnalysisDB,
|
||||
CampaignDB,
|
||||
CampaignPeriodComparisonDB,
|
||||
EvalResultDB,
|
||||
EvalRunDB,
|
||||
EvalTargetDB,
|
||||
ExplorationMessageDB,
|
||||
ExplorationSessionDB,
|
||||
ScenarioDB,
|
||||
TurnDB,
|
||||
@ -46,22 +43,6 @@ M = TypeVar("M") # domain model
|
||||
DB = TypeVar("DB") # persisted table row
|
||||
|
||||
|
||||
class AsyncJobClaimStatus(str, Enum):
|
||||
CLAIMED = "claimed"
|
||||
NOT_FOUND = "not_found"
|
||||
ALREADY_CLAIMED = "already_claimed"
|
||||
NOT_QUEUED = "not_queued"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AsyncJobClaimResult:
|
||||
status: AsyncJobClaimStatus
|
||||
|
||||
@property
|
||||
def claimed(self) -> bool:
|
||||
return self.status is AsyncJobClaimStatus.CLAIMED
|
||||
|
||||
|
||||
class BaseRepository(Generic[M, DB]):
|
||||
"""Shared CRUD skeleton for id-keyed entity repositories.
|
||||
|
||||
@ -393,9 +374,7 @@ class RunRepository(BaseRepository[EvalRun, EvalRunDB]):
|
||||
return CampaignRunClaimResult(CampaignRunClaimStatus.EXISTING, self._from_db(existing))
|
||||
return CampaignRunClaimResult(CampaignRunClaimStatus.CONFLICT)
|
||||
|
||||
def _get_campaign_identity(
|
||||
self, campaign_id: str, plan_index: int, occurrence_index: int
|
||||
) -> Optional[EvalRunDB]:
|
||||
def _get_campaign_identity(self, campaign_id: str, plan_index: int, occurrence_index: int) -> Optional[EvalRunDB]:
|
||||
statement = select(EvalRunDB).where(
|
||||
EvalRunDB.campaign_id == campaign_id,
|
||||
EvalRunDB.campaign_plan_index == plan_index,
|
||||
@ -477,11 +456,7 @@ class RunRepository(BaseRepository[EvalRun, EvalRunDB]):
|
||||
return len(failed)
|
||||
|
||||
def _is_recoverable_campaign_pending(self, db: EvalRunDB) -> bool:
|
||||
if (
|
||||
db.campaign_id is None
|
||||
or db.campaign_plan_index is None
|
||||
or db.campaign_occurrence_index is None
|
||||
):
|
||||
if db.campaign_id is None or db.campaign_plan_index is None or db.campaign_occurrence_index is None:
|
||||
return False
|
||||
campaign = self.session.get(CampaignDB, db.campaign_id)
|
||||
return campaign is not None and campaign.status == CampaignStatus.RUNNING.value
|
||||
@ -779,221 +754,6 @@ class CampaignRepository(BaseRepository[Campaign, CampaignDB]):
|
||||
self.session.refresh(db)
|
||||
|
||||
|
||||
class AsyncJobRepository(Generic[DB]):
|
||||
"""Base class for async job repositories (analysis, comparison).
|
||||
|
||||
Provides common pattern: get_by_campaign, mark_orphans_failed.
|
||||
Subclasses implement upsert with their specific fields.
|
||||
"""
|
||||
|
||||
_table: type[DB]
|
||||
|
||||
def __init__(self, session: Optional[Session] = None):
|
||||
self.session = session or get_session()
|
||||
|
||||
def get_by_campaign(self, campaign_id: str) -> Optional[DB]:
|
||||
statement = select(self._table).where(self._table.campaign_id == campaign_id) # type: ignore[attr-defined]
|
||||
return self.session.exec(statement).first()
|
||||
|
||||
def list_queued(self) -> list[DB]:
|
||||
statement = select(self._table).where(self._table.status == "queued") # type: ignore[attr-defined]
|
||||
return list(self.session.exec(statement).all())
|
||||
|
||||
def claim_queued(self, campaign_id: str) -> "AsyncJobClaimResult":
|
||||
"""Atomically move one queued job to generating.
|
||||
|
||||
The status predicate is the durable idempotency authority. Competing
|
||||
workers may observe the same queued row, but only one can claim it.
|
||||
"""
|
||||
statement = (
|
||||
sql_update(self._table)
|
||||
.where(
|
||||
self._table.campaign_id == campaign_id, # type: ignore[attr-defined]
|
||||
self._table.status == "queued", # type: ignore[attr-defined]
|
||||
)
|
||||
.values(status="generating", error=None, updated_at=utc_now())
|
||||
)
|
||||
result = self.session.exec(statement)
|
||||
self.session.commit()
|
||||
self.session.expire_all()
|
||||
if result.rowcount == 1:
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.CLAIMED)
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.NOT_FOUND)
|
||||
if row.status == "generating": # type: ignore[attr-defined]
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.ALREADY_CLAIMED)
|
||||
return AsyncJobClaimResult(AsyncJobClaimStatus.NOT_QUEUED)
|
||||
|
||||
def prepare_queued_recovery(self, max_attempts: int, exhausted_error: str) -> list[DB]:
|
||||
"""Increment queued recovery attempts and fail exhausted jobs."""
|
||||
recoverable: list[DB] = []
|
||||
rows = self.list_queued()
|
||||
for row in rows:
|
||||
if row.recovery_attempts >= max_attempts: # type: ignore[attr-defined]
|
||||
row.status = "failed" # type: ignore[attr-defined]
|
||||
row.error = exhausted_error # type: ignore[attr-defined]
|
||||
else:
|
||||
row.recovery_attempts += 1 # type: ignore[attr-defined]
|
||||
recoverable.append(row)
|
||||
row.updated_at = utc_now() # type: ignore[attr-defined]
|
||||
self.session.add(row)
|
||||
if rows:
|
||||
self.session.commit()
|
||||
self.session.expire_all()
|
||||
return recoverable
|
||||
|
||||
def mark_orphans_failed(self, error_message: str) -> int:
|
||||
"""服务启动时清理:把滞留的 generating 行标记为 failed。
|
||||
|
||||
异步任务是进程内 asyncio 任务,服务重启后不会恢复;不清理则这些
|
||||
行永远停留在 generating(僵尸状态)。
|
||||
"""
|
||||
rows = self.session.exec(
|
||||
select(self._table).where(self._table.status == "generating") # type: ignore[attr-defined]
|
||||
).all()
|
||||
for row in rows:
|
||||
row.status = "failed" # type: ignore[attr-defined]
|
||||
row.error = error_message # type: ignore[attr-defined]
|
||||
row.updated_at = utc_now() # type: ignore[attr-defined]
|
||||
self.session.add(row)
|
||||
if rows:
|
||||
self.session.commit()
|
||||
return len(rows)
|
||||
|
||||
|
||||
class CampaignAnalysisRepository(AsyncJobRepository[CampaignAnalysisDB]):
|
||||
"""Repository for campaign analysis rows (one per campaign, upserted)."""
|
||||
|
||||
_table = CampaignAnalysisDB
|
||||
|
||||
def enqueue(self, campaign_id: str, *, triggered_by: str = "manual") -> CampaignAnalysisDB:
|
||||
"""Persist an analysis job before launching its process-local task."""
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignAnalysisDB(campaign_id=campaign_id, status="queued", triggered_by=triggered_by)
|
||||
elif row.status in {"queued", "generating"}:
|
||||
return row
|
||||
else:
|
||||
row.status = "queued"
|
||||
row.result = None
|
||||
row.model_config_id = None
|
||||
row.error = None
|
||||
row.triggered_by = triggered_by
|
||||
row.recovery_attempts = 0
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def upsert(
|
||||
self,
|
||||
campaign_id: str,
|
||||
*,
|
||||
status: str,
|
||||
result: Optional[dict] = None,
|
||||
model_config_id: Optional[str] = None,
|
||||
error: Optional[str] = None,
|
||||
triggered_by: str = "manual",
|
||||
) -> CampaignAnalysisDB:
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignAnalysisDB(campaign_id=campaign_id)
|
||||
row.status = status
|
||||
if result is not None:
|
||||
row.set_result(result)
|
||||
else:
|
||||
row.result = None
|
||||
row.model_config_id = model_config_id
|
||||
row.error = error
|
||||
row.triggered_by = triggered_by
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def mark_orphans_failed(self) -> int:
|
||||
"""服务启动时清理:把滞留的 generating 分析行标记为 failed。"""
|
||||
return super().mark_orphans_failed("服务重启导致分析生成中断")
|
||||
|
||||
|
||||
class CampaignPeriodComparisonRepository(AsyncJobRepository[CampaignPeriodComparisonDB]):
|
||||
"""Repository for period-comparison rows (one per campaign, upserted)."""
|
||||
|
||||
_table = CampaignPeriodComparisonDB
|
||||
|
||||
def enqueue(
|
||||
self,
|
||||
campaign_id: str,
|
||||
*,
|
||||
baseline_campaign_id: str,
|
||||
triggered_by: str = "manual",
|
||||
) -> CampaignPeriodComparisonDB:
|
||||
"""Persist a comparison job before launching its process-local task."""
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignPeriodComparisonDB(
|
||||
campaign_id=campaign_id,
|
||||
baseline_campaign_id=baseline_campaign_id,
|
||||
status="queued",
|
||||
triggered_by=triggered_by,
|
||||
)
|
||||
elif row.status in {"queued", "generating"}:
|
||||
return row
|
||||
else:
|
||||
row.baseline_campaign_id = baseline_campaign_id
|
||||
row.status = "queued"
|
||||
row.result = None
|
||||
row.model_config_id = None
|
||||
row.error = None
|
||||
row.triggered_by = triggered_by
|
||||
row.recovery_attempts = 0
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def upsert(
|
||||
self,
|
||||
campaign_id: str,
|
||||
*,
|
||||
status: str,
|
||||
baseline_campaign_id: Optional[str] = None,
|
||||
result: Optional[dict] = None,
|
||||
model_config_id: Optional[str] = None,
|
||||
error: Optional[str] = None,
|
||||
triggered_by: str = "manual",
|
||||
) -> CampaignPeriodComparisonDB:
|
||||
row = self.get_by_campaign(campaign_id)
|
||||
if row is None:
|
||||
row = CampaignPeriodComparisonDB(
|
||||
campaign_id=campaign_id,
|
||||
baseline_campaign_id=baseline_campaign_id or "",
|
||||
)
|
||||
row.status = status
|
||||
if baseline_campaign_id is not None:
|
||||
row.baseline_campaign_id = baseline_campaign_id
|
||||
if result is not None:
|
||||
row.set_result(result)
|
||||
else:
|
||||
row.result = None
|
||||
row.model_config_id = model_config_id
|
||||
row.error = error
|
||||
row.triggered_by = triggered_by
|
||||
row.updated_at = utc_now()
|
||||
self.session.add(row)
|
||||
self.session.commit()
|
||||
self.session.refresh(row)
|
||||
return row
|
||||
|
||||
def mark_orphans_failed(self) -> int:
|
||||
"""服务启动时清理:把滞留的 generating 周期对比行标记为 failed。"""
|
||||
return super().mark_orphans_failed("服务重启导致周期对比生成中断")
|
||||
|
||||
|
||||
class ResultRepository:
|
||||
"""Repository for evaluation results."""
|
||||
|
||||
@ -1052,129 +812,3 @@ class ResultRepository:
|
||||
self.session.commit()
|
||||
self.session.refresh(db)
|
||||
return _result_from_db(db)
|
||||
|
||||
|
||||
class ExplorationSessionRepository(BaseRepository[ExplorationSession, ExplorationSessionDB]):
|
||||
"""Repository for virtual-user exploration sessions (探索会话)."""
|
||||
|
||||
_table = ExplorationSessionDB
|
||||
_order_by = "created_at"
|
||||
|
||||
def _copy_mutable(self, db: ExplorationSessionDB, session_obj: ExplorationSession) -> None:
|
||||
db.goal = session_obj.goal
|
||||
db.status = session_obj.status.value
|
||||
db.triggered_by = session_obj.triggered_by.value
|
||||
db.turn_count = session_obj.turn_count
|
||||
db.error = session_obj.error
|
||||
db.closed_at = session_obj.closed_at
|
||||
db.set_persona(session_obj.persona)
|
||||
if session_obj.seed_ref is not None:
|
||||
db.set_seed_ref(session_obj.seed_ref)
|
||||
if session_obj.experience is not None:
|
||||
db.set_experience(session_obj.experience)
|
||||
if session_obj.judge_review is not None:
|
||||
db.set_judge_review(session_obj.judge_review)
|
||||
|
||||
def _to_db(self, session_obj: ExplorationSession) -> ExplorationSessionDB:
|
||||
db = ExplorationSessionDB(
|
||||
id=session_obj.id,
|
||||
campaign_id=session_obj.campaign_id,
|
||||
target_id=session_obj.target_id,
|
||||
created_at=session_obj.created_at,
|
||||
)
|
||||
self._copy_mutable(db, session_obj)
|
||||
return db
|
||||
|
||||
def _from_db(self, db: ExplorationSessionDB) -> ExplorationSession:
|
||||
return ExplorationSession(
|
||||
id=db.id,
|
||||
campaign_id=db.campaign_id,
|
||||
target_id=db.target_id,
|
||||
persona=db.get_persona(),
|
||||
goal=db.goal,
|
||||
seed_ref=db.get_seed_ref(),
|
||||
status=db.status,
|
||||
triggered_by=db.triggered_by,
|
||||
experience=db.get_experience(),
|
||||
judge_review=db.get_judge_review(),
|
||||
turn_count=db.turn_count,
|
||||
error=db.error,
|
||||
created_at=db.created_at,
|
||||
closed_at=db.closed_at,
|
||||
)
|
||||
|
||||
def update(self, session_obj: ExplorationSession) -> Optional[ExplorationSession]:
|
||||
existing = self.session.get(ExplorationSessionDB, session_obj.id)
|
||||
if not existing:
|
||||
return None
|
||||
self._copy_mutable(existing, session_obj)
|
||||
self.session.add(existing)
|
||||
self.session.commit()
|
||||
self.session.refresh(existing)
|
||||
return self._from_db(existing)
|
||||
|
||||
def list_by_campaign(self, campaign_id: str) -> list[ExplorationSession]:
|
||||
statement = (
|
||||
select(ExplorationSessionDB)
|
||||
.where(ExplorationSessionDB.campaign_id == campaign_id)
|
||||
.order_by(ExplorationSessionDB.created_at)
|
||||
)
|
||||
return [self._from_db(r) for r in self.session.exec(statement).all()]
|
||||
|
||||
def expire_running_sessions(self, campaign_id: str) -> int:
|
||||
"""Expire every still-running exploration session of a finalized campaign.
|
||||
|
||||
Returns the number of sessions expired. Completed/failed sessions keep
|
||||
their evidence untouched.
|
||||
"""
|
||||
expired = 0
|
||||
for session_obj in self.list_by_campaign(campaign_id):
|
||||
if session_obj.status != ExplorationSessionStatus.RUNNING:
|
||||
continue
|
||||
session_obj.status = ExplorationSessionStatus.EXPIRED
|
||||
session_obj.closed_at = utc_now()
|
||||
self.update(session_obj)
|
||||
expired += 1
|
||||
return expired
|
||||
|
||||
|
||||
class ExplorationMessageRepository:
|
||||
"""Append-only repository for exploration session chat rows."""
|
||||
|
||||
def __init__(self, session: Optional[Session] = None):
|
||||
self.session = session or get_session()
|
||||
|
||||
def save_message(self, message: ExplorationMessage) -> ExplorationMessage:
|
||||
db = ExplorationMessageDB(
|
||||
id=message.id,
|
||||
session_id=message.session_id,
|
||||
round_index=message.round_index,
|
||||
role=message.role,
|
||||
content=message.content,
|
||||
latency_ms=message.latency_ms,
|
||||
created_at=message.created_at,
|
||||
)
|
||||
self.session.add(db)
|
||||
self.session.commit()
|
||||
self.session.refresh(db)
|
||||
message.id = db.id
|
||||
return message
|
||||
|
||||
def list_by_session(self, session_id: str) -> list[ExplorationMessage]:
|
||||
statement = (
|
||||
select(ExplorationMessageDB)
|
||||
.where(ExplorationMessageDB.session_id == session_id)
|
||||
.order_by(ExplorationMessageDB.created_at)
|
||||
)
|
||||
return [
|
||||
ExplorationMessage(
|
||||
id=r.id,
|
||||
session_id=r.session_id,
|
||||
round_index=r.round_index,
|
||||
role=r.role,
|
||||
content=r.content,
|
||||
latency_ms=r.latency_ms,
|
||||
created_at=r.created_at,
|
||||
)
|
||||
for r in self.session.exec(statement).all()
|
||||
]
|
||||
|
||||
@ -11,7 +11,7 @@ from fastapi import APIRouter, Body, Depends, HTTPException, Response
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.analysis import resolve_analysis_model
|
||||
from agenteval.evaluation.analysis import AnalysisError, validate_analysis_request
|
||||
from agenteval.evaluation.campaign_lifecycle import CampaignCreateError, CampaignLifecycleError
|
||||
from agenteval.evaluation.campaign_lifecycle import cancel_campaign as cancel_campaign_lifecycle
|
||||
from agenteval.evaluation.campaign_lifecycle import create_campaign as create_campaign_lifecycle
|
||||
@ -20,7 +20,7 @@ from agenteval.evaluation.campaign_runner import campaign_runtime
|
||||
from agenteval.evaluation.comparison import ComparisonError, validate_comparison_request
|
||||
from agenteval.evaluation.intelligence_jobs import enqueue_campaign_analysis, enqueue_campaign_comparison
|
||||
from agenteval.evaluation.report_render import render_campaign_markdown
|
||||
from agenteval.models import CampaignPlanEntry, CampaignStatus, ExplorationBudgetConfig, ExplorationSeeds
|
||||
from agenteval.models import CampaignPlanEntry, ExplorationBudgetConfig, ExplorationSeeds
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
from agenteval.web.deps import get_db
|
||||
|
||||
@ -125,13 +125,10 @@ async def trigger_campaign_analysis(campaign_id: str, session: Session = Depends
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if not campaign:
|
||||
raise HTTPException(status_code=404, detail="campaign not found")
|
||||
if campaign.status in (CampaignStatus.PLANNED, CampaignStatus.RUNNING):
|
||||
raise HTTPException(status_code=400, detail="活动完成后才能生成智能分析")
|
||||
if resolve_analysis_model(campaign, session) is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="未配置分析模型:请在模型配置中心将某个 chat 配置设为「分析默认」,或为该活动指定分析模型",
|
||||
)
|
||||
try:
|
||||
validate_analysis_request(session, campaign)
|
||||
except AnalysisError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc))
|
||||
enqueue_campaign_analysis(campaign_id, triggered_by="manual")
|
||||
return {"status": "generating"}
|
||||
|
||||
|
||||
@ -19,11 +19,8 @@ from agenteval.exploration.errors import (
|
||||
ExplorationNotFoundError,
|
||||
)
|
||||
from agenteval.exploration.models import ExplorationTrigger
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationMessageRepository,
|
||||
ExplorationSessionRepository,
|
||||
)
|
||||
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
from agenteval.web.deps import get_db
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
41
docs/adr/0012-intelligence-job-settlement-unification.md
Normal file
41
docs/adr/0012-intelligence-job-settlement-unification.md
Normal file
@ -0,0 +1,41 @@
|
||||
# ADR-0012: 活动智能作业结算统一——execute 深接缝与 DB 权威
|
||||
|
||||
**状态**: 已接受
|
||||
**日期**: 2026-08-24
|
||||
**决策者**: 架构团队
|
||||
**相关**: ADR-0004(跨 run 聚合口径)、ADR-0006(拒绝通用条件写模块)、ADR-0011(终态纪律)、CONTEXT.md(智能分析、周期对比)
|
||||
|
||||
## Context
|
||||
|
||||
智能分析与周期对比两个 executor 各自手写了一遍完整的作业结算序列:幂等建行 → `claim_queued` 认领 → `upsert(generating)` → 领域工作 → 异常截断落 failed → completed 落账 → session 关闭。两遍编排约 200 行,逐句相似又逐处微妙不同(失败行带不带 baseline、截断长度、校验时机),任何一处结算语义的修改都要在两处同步,漂移风险随每次改动累积。
|
||||
|
||||
同时,自动触发的跳过守卫(正式线 + 分析模型可解析)在 `campaign_runner._auto_start_analysis` 与 executor 的自动对比链各写了一遍;分析触发的校验(活动终态 + 模型配置)在路由与分析 executor 各写了一遍——而周期对比已有 `validate_comparison_request` 共享入口的先例。
|
||||
|
||||
## Decision
|
||||
|
||||
1. **DB 行是作业权威,`TaskRegistry` 只保存进程句柄。** 排队/执行/成败一切状态以 `CampaignAnalysisDB` / `CampaignPeriodComparisonDB` 行为准;registry 只负责当前进程内的任务句柄、幂等启动与关停。重启后凭 DB 行恢复(`recover_campaign_intelligence_jobs`),不凭 registry。
|
||||
|
||||
2. **结算顺序:先认领再 generating。** `claim_queued` 的条件更新是耐久幂等权威——竞争者可以观察到同一 queued 行,但只有一个能把它推到 generating;抢不到的一方静默退出,绝不重复跑领域工作。`execute` 内部的固定次序:建行(ensure_queued)→ 认领 → 校验 → generating → work → 结算。
|
||||
|
||||
3. **`execute(job_kind, campaign_id, …)` 是结算的唯一实现。** 幂等建行、认领、异常归一(error 截断 500 字符一处)、failed/completed 落账全部收进 `intelligence_jobs.execute`;分析与对比退化为三个小 adapter:`ensure_queued`(建行/静默放弃)、`validate`(认领后校验,失败抛携带落账字段的 `JobValidationError`)、`work_fn`(纯领域工作,返回结果 + 附加落账字段)。删除任一侧 adapter 的编排序列与截断复制。
|
||||
|
||||
4. **校验共享入口对齐对比先例。** 新增 `validate_analysis_request`(活动终态 → 模型),路由捕获映射 400、executor 捕获落 failed 行,两处措辞与顺序不再漂移;与既有 `validate_comparison_request` 同构。自动触发跳过守卫收敛为 `auto_intelligence_eligible` 单一判断点(活动完成自动分析 + 分析完成自动对比共用)。
|
||||
|
||||
5. **截断与恢复上限语义。** work 异常落账统一截断 500 字符;重启恢复 queued 行上限 3 次(`MAX_QUEUED_RECOVERY_ATTEMPTS`),超限置 failed 并注明原因——与 ADR-0011 终态纪律一致:自愈有上限,超限收敛到终态且可见。
|
||||
|
||||
6. **投影归位。** `load_comparison_view`(GET /comparison 响应形状)移回 `campaign_read_model`;`comparison.py` 只留领域工作(指纹、基线解析、机械 diff、校验、叙述)。零调用的 `build_comparison_payload` 删除。
|
||||
|
||||
### app.py 关停顺序现状(暂不动)
|
||||
|
||||
关停依序:`scheduler_runtime` → `campaign_runtime` → `run_registry` → 智能作业 → `judge_registry`。顺序有意(上游先停,避免停掉的调度再派生新任务),现状有 lifespan 测试覆盖智能评估侧;在出现顺序相关的真实故障前不归一为通用机制(与 ADR-0006 拒绝通用化的立场一致)。
|
||||
|
||||
### Considered Options(已拒绝)
|
||||
|
||||
- **把结算序列抽成通用 async-job 框架**:两个 adapter 的校验时机、落账字段、静默放弃语义各有领域差异,通用框架的接口会比两个小 adapter 更复杂(浅模块)。等出现第三、第四种耐久作业再议。
|
||||
- **CAS 共享原语收编认领逻辑**:本轮决策不收敛(三处条件写语义有差异),见 v3 重构计划「刻意不做」。
|
||||
|
||||
## Consequences
|
||||
|
||||
- 结算语义修改只需动 `execute` 一处;characterization 测试(认领竞争、重复触发、截断、恢复上限)锁定契约。
|
||||
- executor 的校验面变宽:非终态活动直接调 executor 会落 failed 行(此前会继续跑)——与路由校验对齐后的刻意收敛。
|
||||
- 失败行的模型缺失措辞统一为路由侧完整文案(含「或为该活动指定分析模型」引导)。
|
||||
@ -69,13 +69,15 @@ def seeded_db(db_session, monkeypatch):
|
||||
@pytest.fixture()
|
||||
def analysis_spy(monkeypatch):
|
||||
"""Spy the analysis seam: resolvable model, recorded enqueue calls."""
|
||||
from agenteval.evaluation import analysis as analysis_module
|
||||
|
||||
calls: list[tuple[str, str]] = []
|
||||
monkeypatch.setattr(
|
||||
campaign_runner,
|
||||
"enqueue_campaign_analysis",
|
||||
lambda cid, *, triggered_by: calls.append((cid, triggered_by)),
|
||||
)
|
||||
monkeypatch.setattr(campaign_runner, "resolve_analysis_model", lambda campaign, session: object())
|
||||
monkeypatch.setattr(analysis_module, "resolve_analysis_model", lambda campaign, session: object())
|
||||
return calls
|
||||
|
||||
|
||||
@ -141,7 +143,9 @@ async def test_cancelled_campaign_does_not_enqueue(seeded_db, runtime, analysis_
|
||||
|
||||
|
||||
async def test_missing_analysis_model_skips_silently(seeded_db, runtime, monkeypatch, analysis_spy):
|
||||
monkeypatch.setattr(campaign_runner, "resolve_analysis_model", lambda campaign, session: None)
|
||||
from agenteval.evaluation import analysis as analysis_module
|
||||
|
||||
monkeypatch.setattr(analysis_module, "resolve_analysis_model", lambda campaign, session: None)
|
||||
campaign = _make_campaign(seeded_db)
|
||||
assert runtime.start(campaign.id)
|
||||
await _await_terminal(seeded_db, campaign.id)
|
||||
|
||||
@ -307,7 +307,7 @@ async def test_campaign_report_markdown_export(client, seeded_db):
|
||||
|
||||
def _seed_exploration_session(seeded_db, campaign_id: str) -> None:
|
||||
from agenteval.exploration.models import ExplorationSession
|
||||
from agenteval.storage.repository import ExplorationSessionRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
|
||||
repo = ExplorationSessionRepository(seeded_db)
|
||||
session_obj = repo.create(
|
||||
|
||||
@ -11,7 +11,8 @@ from datetime import timedelta
|
||||
import pytest
|
||||
from agenteval.models import Campaign, ChannelType, EvalTarget, PlatformType, TargetStatus
|
||||
from agenteval.storage.db import ExplorationSessionDB, utc_now
|
||||
from agenteval.storage.repository import CampaignRepository, ExplorationSessionRepository, TargetRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
from agenteval.storage.repository import CampaignRepository, TargetRepository
|
||||
from agenteval.web.app import app
|
||||
from httpx import ASGITransport, AsyncClient
|
||||
|
||||
|
||||
@ -150,7 +150,7 @@ async def test_patrol_watermark_advances_and_reports_increments(client, seeded_d
|
||||
|
||||
async def test_patrol_budget_reflects_existing_sessions(client, seeded_db):
|
||||
from agenteval.exploration.models import ExplorationSession
|
||||
from agenteval.storage.repository import ExplorationSessionRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
|
||||
ExplorationSessionRepository(seeded_db).create(ExplorationSession(
|
||||
campaign_id="c-prod", target_id="t-1", goal="查询账单", persona={"name": "x"},
|
||||
|
||||
@ -23,9 +23,9 @@ from agenteval.models import (
|
||||
TargetStatus,
|
||||
)
|
||||
from agenteval.storage.db import utc_now
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationSessionRepository,
|
||||
ScenarioRepository,
|
||||
TargetRepository,
|
||||
)
|
||||
|
||||
@ -10,7 +10,7 @@ from agenteval.evaluation.analysis import (
|
||||
resolve_analysis_model,
|
||||
)
|
||||
from agenteval.evaluation.intelligence_jobs import execute_campaign_analysis_job
|
||||
from agenteval.models import Campaign, CampaignPlanEntry, EvalRun, RunStatus
|
||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus, EvalRun, RunStatus
|
||||
from agenteval.storage.db import CampaignAnalysisDB, EvalResultDB, ModelConfigDB, TurnDB
|
||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||
from agenteval.storage.repository import CampaignRepository, RunRepository
|
||||
@ -41,6 +41,7 @@ def _campaign(**overrides) -> Campaign:
|
||||
"target_id": "t-1",
|
||||
"window_seconds": 86400,
|
||||
"time_scale": 1.0,
|
||||
"status": CampaignStatus.COMPLETED,
|
||||
"plan": [CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=2)],
|
||||
}
|
||||
data.update(overrides)
|
||||
|
||||
@ -4,7 +4,8 @@ import pytest
|
||||
from agenteval.evaluation.campaign_lifecycle import CampaignLifecycleError, cancel_campaign
|
||||
from agenteval.exploration.models import ExplorationSession, ExplorationSessionStatus
|
||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus
|
||||
from agenteval.storage.repository import CampaignRepository, ExplorationSessionRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
from sqlalchemy import event
|
||||
|
||||
|
||||
|
||||
@ -4,7 +4,8 @@ import pytest
|
||||
from agenteval.evaluation.campaign_lifecycle import CampaignLifecycleError, complete_campaign
|
||||
from agenteval.exploration.models import ExplorationSession, ExplorationSessionStatus
|
||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus
|
||||
from agenteval.storage.repository import CampaignRepository, ExplorationSessionRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
from agenteval.storage.repository import CampaignRepository
|
||||
from sqlalchemy import event
|
||||
|
||||
|
||||
|
||||
@ -12,8 +12,8 @@ from agenteval.models import (
|
||||
EvalTarget,
|
||||
Scenario,
|
||||
)
|
||||
from agenteval.storage.async_job_repository import CampaignAnalysisRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignRepository,
|
||||
ScenarioRepository,
|
||||
TargetRepository,
|
||||
@ -86,7 +86,7 @@ def test_view_with_comparison(db_session):
|
||||
comparisons.upsert(baseline.id, status="completed", result={"summary": "baseline"})
|
||||
comparisons.upsert(campaign.id, status="completed", result={"summary": "current"})
|
||||
# Create a comparison row
|
||||
from agenteval.storage.repository import CampaignPeriodComparisonRepository
|
||||
from agenteval.storage.async_job_repository import CampaignPeriodComparisonRepository
|
||||
|
||||
cmp_repo = CampaignPeriodComparisonRepository(db_session)
|
||||
cmp_repo.upsert(
|
||||
|
||||
@ -9,9 +9,9 @@ from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from agenteval.evaluation.campaign_read_model import load_comparison_view
|
||||
from agenteval.evaluation.comparison import (
|
||||
ComparisonError,
|
||||
load_comparison_view,
|
||||
validate_comparison_request,
|
||||
)
|
||||
from agenteval.models import (
|
||||
@ -23,9 +23,8 @@ from agenteval.models import (
|
||||
EvalTarget,
|
||||
Scenario,
|
||||
)
|
||||
from agenteval.storage.async_job_repository import CampaignAnalysisRepository, CampaignPeriodComparisonRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignPeriodComparisonRepository,
|
||||
CampaignRepository,
|
||||
ScenarioRepository,
|
||||
TargetRepository,
|
||||
|
||||
@ -23,10 +23,9 @@ from agenteval.exploration.models import (
|
||||
from agenteval.models import Campaign, CampaignPlanEntry
|
||||
from agenteval.storage.db import ModelConfigDB
|
||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||
from agenteval.storage.exploration_repository import ExplorationMessageRepository, ExplorationSessionRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationMessageRepository,
|
||||
ExplorationSessionRepository,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@ -22,9 +22,9 @@ from agenteval.models import (
|
||||
ExplorationSeeds,
|
||||
RunStatus,
|
||||
)
|
||||
from agenteval.storage.exploration_repository import ExplorationSessionRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
ExplorationSessionRepository,
|
||||
RunRepository,
|
||||
TargetRepository,
|
||||
)
|
||||
|
||||
330
tests/unit/test_intelligence_jobs_settlement_characterization.py
Normal file
330
tests/unit/test_intelligence_jobs_settlement_characterization.py
Normal file
@ -0,0 +1,330 @@
|
||||
"""Characterization tests: intelligence-job settlement contract (Phase 2.8).
|
||||
|
||||
锁定两个 executor(智能分析 / 周期对比)的结算契约边界,供
|
||||
``execute(job_kind, work_fn)`` 接缝收敛(Phase 2.9–2.11)对照:
|
||||
|
||||
1. 认领竞争:两个执行流观察同一 queued 行,只有一个真正跑领域工作。
|
||||
2. 重复触发:completed 行被再次触发时重置回 queued 并重跑;generating
|
||||
中的行不被回退(后者已有 test_reenqueue_does_not_move_generating_*)。
|
||||
3. 失败截断:work 异常落账时 error 截断到 500 字符(一处语义)。
|
||||
4. 恢复上限:``recover_campaign_intelligence_jobs`` 把超过恢复次数上限
|
||||
的 queued 行落 failed,不再重启。
|
||||
|
||||
不改动任何生产代码。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import timedelta
|
||||
|
||||
import pytest
|
||||
from agenteval.evaluation import intelligence_jobs
|
||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus, EvalRun, RunStatus, RunSummary
|
||||
from agenteval.storage.db import (
|
||||
CampaignAnalysisDB,
|
||||
ModelConfigDB,
|
||||
utc_now,
|
||||
)
|
||||
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||
from agenteval.storage.async_job_repository import CampaignAnalysisRepository, CampaignPeriodComparisonRepository
|
||||
from agenteval.storage.repository import (
|
||||
CampaignRepository,
|
||||
RunRepository,
|
||||
)
|
||||
from sqlmodel import Session, SQLModel, create_engine
|
||||
|
||||
T0 = utc_now().replace(tzinfo=None) - timedelta(hours=1)
|
||||
|
||||
STAGE1 = json.dumps({"narrative": "n", "problems": []})
|
||||
STAGE2 = json.dumps({"overall": "o", "problems": [], "suggestions": []})
|
||||
NARRATION = json.dumps({"trend": "stable", "summary": "s"})
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db_session(tmp_path):
|
||||
from agenteval.storage.db import ( # noqa: F401
|
||||
CampaignDB,
|
||||
)
|
||||
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'settlement.db'}",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
SQLModel.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session.close()
|
||||
engine.dispose()
|
||||
|
||||
|
||||
class FakeChatClient:
|
||||
def __init__(self, *responses):
|
||||
self._responses = list(responses)
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, messages):
|
||||
self.calls += 1
|
||||
if not self._responses:
|
||||
raise AssertionError("unexpected extra LLM call")
|
||||
item = self._responses.pop(0)
|
||||
if isinstance(item, Exception):
|
||||
raise item
|
||||
return item
|
||||
|
||||
|
||||
def _campaign(
|
||||
campaign_id: str = "camp-1",
|
||||
*,
|
||||
time_scale: float = 1.0,
|
||||
status: CampaignStatus = CampaignStatus.COMPLETED,
|
||||
completed_at=T0 + timedelta(hours=1),
|
||||
) -> Campaign:
|
||||
return Campaign(
|
||||
id=campaign_id,
|
||||
name=f"campaign-{campaign_id}",
|
||||
target_id="t-1",
|
||||
window_seconds=86400,
|
||||
time_scale=time_scale,
|
||||
plan=[CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=2)],
|
||||
status=status,
|
||||
completed_at=completed_at,
|
||||
)
|
||||
|
||||
|
||||
def _seed_config(session) -> None:
|
||||
ModelConfigRepository(session).create(
|
||||
ModelConfigDB(
|
||||
id="mc-default",
|
||||
name="cfg",
|
||||
provider="openai_compatible",
|
||||
capability="chat",
|
||||
endpoint_url="https://models.example.com/v1/chat/completions",
|
||||
model_name="m",
|
||||
is_analysis_default=True,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _seed_failed_run(session, campaign_id: str = "camp-1") -> None:
|
||||
from agenteval.storage.db import EvalResultDB, TurnDB
|
||||
|
||||
RunRepository(session).create(
|
||||
EvalRun(
|
||||
id=f"run-{campaign_id}",
|
||||
target_id="t-1",
|
||||
scenario_id="s-1",
|
||||
campaign_id=campaign_id,
|
||||
status=RunStatus.COMPLETED,
|
||||
)
|
||||
)
|
||||
turn = TurnDB(id=f"run-{campaign_id}-turn-0", run_id=f"run-{campaign_id}", case_id="c0", round_index=0)
|
||||
turn.set_sent_message({"msgBody": {"content": "用户消息"}})
|
||||
turn.set_reply({"msgBody": {"content": "答"}})
|
||||
session.add(turn)
|
||||
session.add(
|
||||
EvalResultDB(
|
||||
run_id=f"run-{campaign_id}",
|
||||
case_id="c0",
|
||||
turn_id=turn.id,
|
||||
rule_type="llm_score",
|
||||
passed=False,
|
||||
reason="不合格",
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
|
||||
class TestClaimRace:
|
||||
"""两个执行流竞争同一 queued 行:只有一个跑领域工作。"""
|
||||
|
||||
async def test_concurrent_analysis_jobs_run_work_once(self, db_session):
|
||||
_seed_config(db_session)
|
||||
CampaignRepository(db_session).create(_campaign())
|
||||
_seed_failed_run(db_session)
|
||||
CampaignAnalysisRepository(db_session).enqueue("camp-1", triggered_by="auto")
|
||||
|
||||
client = FakeChatClient(STAGE1, STAGE2)
|
||||
await asyncio.gather(
|
||||
intelligence_jobs.execute_campaign_analysis_job(
|
||||
"camp-1",
|
||||
triggered_by="auto",
|
||||
chat_client=client,
|
||||
session_factory=lambda: db_session,
|
||||
),
|
||||
intelligence_jobs.execute_campaign_analysis_job(
|
||||
"camp-1",
|
||||
triggered_by="auto",
|
||||
chat_client=client,
|
||||
session_factory=lambda: db_session,
|
||||
),
|
||||
)
|
||||
|
||||
assert client.calls == 2 # 阶段一 1 场景 + 阶段二 1 次 = 恰好一次完整工作
|
||||
row = CampaignAnalysisRepository(db_session).get_by_campaign("camp-1")
|
||||
assert row.status == "completed"
|
||||
|
||||
async def test_generating_row_blocks_second_executor(self, db_session):
|
||||
"""已在 generating 的行再次进入 executor:不认领、不跑工作。"""
|
||||
_seed_config(db_session)
|
||||
CampaignRepository(db_session).create(_campaign())
|
||||
CampaignAnalysisRepository(db_session).upsert("camp-1", status="generating", triggered_by="auto")
|
||||
|
||||
client = FakeChatClient(STAGE1, STAGE2)
|
||||
await intelligence_jobs.execute_campaign_analysis_job(
|
||||
"camp-1",
|
||||
triggered_by="auto",
|
||||
chat_client=client,
|
||||
session_factory=lambda: db_session,
|
||||
)
|
||||
|
||||
assert client.calls == 0
|
||||
assert CampaignAnalysisRepository(db_session).get_by_campaign("camp-1").status == "generating"
|
||||
|
||||
|
||||
class TestRetriggerSemantics:
|
||||
"""completed 行再触发 → 重置 queued 并重跑;queued 重复入队保持原样。"""
|
||||
|
||||
def test_retrigger_of_completed_campaign_resets_and_relaunches(self, db_session, monkeypatch):
|
||||
"""重复触发幂等观察:completed 行经 enqueue 重置回 queued 并重启 worker。"""
|
||||
repo = CampaignAnalysisRepository(db_session)
|
||||
repo.upsert("camp-1", status="completed", result={"overall": "旧结论"}, triggered_by="auto")
|
||||
launched = []
|
||||
monkeypatch.setattr(
|
||||
intelligence_jobs,
|
||||
"_launch_analysis",
|
||||
lambda cid, *, triggered_by: launched.append((cid, triggered_by)),
|
||||
)
|
||||
|
||||
intelligence_jobs.enqueue_campaign_analysis("camp-1", triggered_by="manual", session=db_session)
|
||||
|
||||
row = repo.get_by_campaign("camp-1")
|
||||
assert row.status == "queued"
|
||||
assert row.result is None
|
||||
assert launched == [("camp-1", "manual")]
|
||||
|
||||
def test_reenqueue_completed_row_resets_to_queued(self, db_session):
|
||||
repo = CampaignAnalysisRepository(db_session)
|
||||
repo.upsert("camp-1", status="completed", result={"overall": "旧"}, triggered_by="auto")
|
||||
|
||||
row = repo.enqueue("camp-1", triggered_by="manual")
|
||||
|
||||
assert row.status == "queued"
|
||||
assert row.result is None
|
||||
assert row.triggered_by == "manual"
|
||||
|
||||
def test_reenqueue_completed_comparison_resets_and_replaces_baseline(self, db_session):
|
||||
repo = CampaignPeriodComparisonRepository(db_session)
|
||||
repo.upsert(
|
||||
"camp-cur",
|
||||
status="completed",
|
||||
baseline_campaign_id="camp-old",
|
||||
result={"trend": "stable"},
|
||||
triggered_by="auto",
|
||||
)
|
||||
|
||||
row = repo.enqueue("camp-cur", baseline_campaign_id="camp-new", triggered_by="manual")
|
||||
|
||||
assert row.status == "queued"
|
||||
assert row.baseline_campaign_id == "camp-new"
|
||||
assert row.result is None
|
||||
|
||||
|
||||
class TestFailureTruncation:
|
||||
"""work 异常落账:error 截断到 500 字符。"""
|
||||
|
||||
async def test_analysis_error_truncated_to_500(self, db_session):
|
||||
_seed_config(db_session)
|
||||
CampaignRepository(db_session).create(_campaign())
|
||||
_seed_failed_run(db_session)
|
||||
|
||||
client = FakeChatClient(RuntimeError("长" * 800))
|
||||
await intelligence_jobs.execute_campaign_analysis_job(
|
||||
"camp-1",
|
||||
triggered_by="auto",
|
||||
chat_client=client,
|
||||
session_factory=lambda: db_session,
|
||||
)
|
||||
|
||||
row = CampaignAnalysisRepository(db_session).get_by_campaign("camp-1")
|
||||
assert row.status == "failed"
|
||||
assert len(row.error) == 500
|
||||
|
||||
async def test_comparison_error_truncated_to_500(self, db_session):
|
||||
_seed_config(db_session)
|
||||
baseline, current = _campaign("camp-base"), _campaign("camp-cur")
|
||||
CampaignRepository(db_session).create(baseline)
|
||||
CampaignRepository(db_session).create(current)
|
||||
for cid in ("camp-base", "camp-cur"):
|
||||
analysis = CampaignAnalysisDB(campaign_id=cid, status="completed")
|
||||
analysis.set_result({"overall": cid})
|
||||
db_session.add(analysis)
|
||||
RunRepository(db_session).create(
|
||||
EvalRun(
|
||||
id=f"run-{cid}",
|
||||
target_id="t-1",
|
||||
scenario_id="s-1",
|
||||
campaign_id=cid,
|
||||
status=RunStatus.COMPLETED,
|
||||
started_at=utc_now(),
|
||||
summary=RunSummary(total_cases=2, pass_rate=0.5, avg_latency_ms=700),
|
||||
)
|
||||
)
|
||||
db_session.commit()
|
||||
|
||||
client = FakeChatClient(RuntimeError("长" * 800))
|
||||
await intelligence_jobs.execute_campaign_comparison_job(
|
||||
"camp-cur",
|
||||
triggered_by="manual",
|
||||
baseline_campaign_id="camp-base",
|
||||
chat_client=client,
|
||||
session_factory=lambda: db_session,
|
||||
)
|
||||
|
||||
row = CampaignPeriodComparisonRepository(db_session).get_by_campaign("camp-cur")
|
||||
assert row.status == "failed"
|
||||
assert len(row.error) == 500
|
||||
|
||||
|
||||
class TestRecoveryCap:
|
||||
"""recover_campaign_intelligence_jobs:超限落 failed,中断行计入返回。"""
|
||||
|
||||
def test_exhausted_queued_job_fails_instead_of_relaunch(self, db_session, monkeypatch):
|
||||
repo = CampaignAnalysisRepository(db_session)
|
||||
repo.enqueue("camp-1", triggered_by="auto")
|
||||
launched = []
|
||||
monkeypatch.setattr(
|
||||
intelligence_jobs,
|
||||
"_launch_analysis",
|
||||
lambda cid, *, triggered_by: launched.append(cid),
|
||||
)
|
||||
|
||||
for _ in range(intelligence_jobs.MAX_QUEUED_RECOVERY_ATTEMPTS):
|
||||
intelligence_jobs.recover_campaign_intelligence_jobs(db_session)
|
||||
assert launched == ["camp-1"] * intelligence_jobs.MAX_QUEUED_RECOVERY_ATTEMPTS
|
||||
|
||||
interrupted, relaunched = intelligence_jobs.recover_campaign_intelligence_jobs(db_session)
|
||||
assert relaunched == 0
|
||||
assert launched == ["camp-1"] * intelligence_jobs.MAX_QUEUED_RECOVERY_ATTEMPTS
|
||||
row = repo.get_by_campaign("camp-1")
|
||||
assert row.status == "failed"
|
||||
assert "上限" in row.error
|
||||
|
||||
def test_recovery_reports_interrupted_generating_rows(self, db_session, monkeypatch):
|
||||
CampaignAnalysisRepository(db_session).upsert("camp-a", status="generating", triggered_by="auto")
|
||||
CampaignPeriodComparisonRepository(db_session).upsert(
|
||||
"camp-c",
|
||||
status="generating",
|
||||
baseline_campaign_id="camp-base",
|
||||
triggered_by="auto",
|
||||
)
|
||||
monkeypatch.setattr(intelligence_jobs, "_launch_analysis", lambda *a, **kw: None)
|
||||
monkeypatch.setattr(intelligence_jobs, "_launch_comparison", lambda *a, **kw: None)
|
||||
|
||||
interrupted, relaunched = intelligence_jobs.recover_campaign_intelligence_jobs(db_session)
|
||||
|
||||
assert interrupted == 2
|
||||
assert relaunched == 0
|
||||
assert CampaignAnalysisRepository(db_session).get_by_campaign("camp-a").status == "failed"
|
||||
assert CampaignPeriodComparisonRepository(db_session).get_by_campaign("camp-c").status == "failed"
|
||||
@ -10,7 +10,7 @@ import asyncio
|
||||
import pytest
|
||||
from agenteval.evaluation import intelligence_jobs
|
||||
from agenteval.exploration import judge
|
||||
from agenteval.storage.repository import (
|
||||
from agenteval.storage.async_job_repository import (
|
||||
AsyncJobClaimStatus,
|
||||
CampaignAnalysisRepository,
|
||||
CampaignPeriodComparisonRepository,
|
||||
|
||||
Loading…
Reference in New Issue
Block a user